Repository navigation
Expose batched CFG through the inference config - #279
Open
sunil-srinivasa wants to merge 1 commit into
Open
sunil-srinivasa wants to merge 1 commit into
sunil-srinivasa wants to merge 1 commit into
Conversation
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>
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
What
Adds a
--use-batched-cfginference 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_batchtakes ause_batched_cfgargument, 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.
OmniInferenceomits the argument from itsgenerate_samples_from_batchcall, and there is no corresponding field onSetupArgs, so the model-side default ofFalseis 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
latencypreset raisescfgp_sizeon its own, so a user who passes--use-batched-cfgon multiple GPUs can land on it without having touched--cfgp-sizeat 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
cosmos_framework/inference/common/args.pyuse_batched_cfgtoSetupArgs/SetupOverrides(defaultFalse)cosmos_framework/inference/args.py_build_parallelismcosmos_framework/inference/inference.pygenerate_samples_from_batchcosmos_framework/inference/args_test.pydocs/inference.md,CHANGELOG.md58 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=0To 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.pyandcosmos_framework/inference/common/args_test.py: 32 passed.cosmos_framework/inference/collection: 184 tests collect with no import errors, confirming the new required field onSetupArgsbreaks no construction site (every one goes through_buildfrom the overrides).ruff==0.12.7check and format produce output identical to pristinemainon the touched files — no new findings.