Skip to content

Expose batched CFG through the inference config - #279

Open
sunil-srinivasa wants to merge 1 commit into
NVIDIA:mainfrom
sunil-srinivasa:ssrinivasa/expose-batched-cfg
Open

sunil-srinivasa wants to merge 1 commit into
NVIDIA:mainfrom
sunil-srinivasa:ssrinivasa/expose-batched-cfg

Conversation

@sunil-srinivasa

Copy link
Copy Markdown
Contributor

What

Adds a --use-batched-cfg inference flag that runs the two classifier-free guidance branches as a single forward of size 2N instead of two sequential N-forwards.

Why

The model side is already implemented and shipped. OmniMoTModel.generate_samples_from_batch takes a use_batched_cfg argument, and when it is set the model builds a doubled (cond + uncond) context, reuses a doubled pack template and a shared text K/V cache, and runs both guidance branches in one forward.

Nothing ever sets it. OmniInference omits the argument from its generate_samples_from_batch call, and there is no corresponding field on SetupArgs, so the model-side default of False is the only value the path has ever taken. The implementation is unreachable from the CLI and cannot even be benchmarked.

This PR adds the missing plumbing — one config field, one argument at the call site — so the existing implementation becomes usable.

Behaviour

Default is False, matching the model-side default, so no existing run changes behaviour. This is purely additive.

Batching the two branches does not reduce the FLOPs that CFG costs. It only helps while a single N-forward still leaves the GPU underutilized — small images, short clips, low step counts. Long or high-resolution workloads already saturate the device, where it gives nothing back and raises peak activation memory. The flag is therefore opt-in and documented as something to measure per workload rather than switch on globally.

CFG parallelism takes precedence

CFG parallelism splits the same conditional/unconditional pair across two GPUs, so the two features are mutually exclusive — and the model already enforces this internally by disabling batched CFG when cfgp_enabled.

The config-level collision is easy to hit by accident: the latency preset raises cfgp_size on its own, so a user who passes --use-batched-cfg on multiple GPUs can land on it without having touched --cfgp-size at all. This resolves it in favour of cfg-parallelism at argument-build time and logs a warning, so the effective configuration is visible up front rather than silently corrected deeper in the stack.

Changes

File Change
cosmos_framework/inference/common/args.py Add use_batched_cfg to SetupArgs / SetupOverrides (default False)
cosmos_framework/inference/args.py Resolve the cfg-parallelism collision in _build_parallelism
cosmos_framework/inference/inference.py Pass the flag into generate_samples_from_batch
cosmos_framework/inference/args_test.py Cover the default, the opt-in, and the cfgp precedence
docs/inference.md, CHANGELOG.md Document the flag

58 insertions, no deletions.

Usage

python -m cosmos_framework.scripts.inference \
    --parallelism-preset=latency \
    -i "inputs/omni/t2i.json" \
    -o outputs/omni_edge \
    --checkpoint-path Cosmos3-Edge \
    --use-batched-cfg \
    --seed=0

To A/B it, hold guidance, seed and step count fixed and vary only this flag, using the existing benchmark harness (--benchmark --warmup=2 --num-iterations=10 --no-diffusion-cache). Because batching is a pure scheduling change, outputs should be numerically near-identical between the two arms, which doubles as a correctness check.

Testing

  • cosmos_framework/inference/args_test.py and cosmos_framework/inference/common/args_test.py: 32 passed.
  • Full cosmos_framework/inference/ collection: 184 tests collect with no import errors, confirming the new required field on SetupArgs breaks no construction site (every one goes through _build from the overrides).
  • ruff==0.12.7 check and format produce output identical to pristine main on the touched files — no new findings.

The model already implements batched classifier-free guidance: pass
use_batched_cfg=True to generate_samples_from_batch and it runs the
conditional and unconditional branches as one forward of size 2N
instead of two sequential N-forwards, reusing a doubled pack template
and a shared text K/V cache.

Nothing reaches it. OmniInference omits the argument entirely, so the
model-side default of False is the only value the path has ever seen
and the implementation is unreachable from the CLI.

Add the missing `use_batched_cfg` setup argument and pass it through.
Default stays False, so no existing run changes behaviour.

CFG parallelism splits the same cond/uncond pair across GPUs, so the
two are mutually exclusive; the `latency` preset raises cfgp_size on
its own, which means a user can land on that collision without having
touched cfgp at all. Resolve it in favour of cfgp with a warning.

Signed-off-by: Sunil Srinivasa <ssrinivasa@nvidia.com>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant