diff --git a/cosmos_framework/configs/base/defaults/model_config.py b/cosmos_framework/configs/base/defaults/model_config.py index c82fee5f..05236b70 100644 --- a/cosmos_framework/configs/base/defaults/model_config.py +++ b/cosmos_framework/configs/base/defaults/model_config.py @@ -130,6 +130,11 @@ class OmniMoTModelConfig: """ tokenizer: LazyDict = None + load_vision_tokenizer: bool = True + """Instantiate the vision tokenizer. + + Reasoner-only inference disables this to avoid loading the generation VAE. + """ net: LazyDict = None ema: EMAConfig = EMAConfig() diff --git a/cosmos_framework/inference/args.py b/cosmos_framework/inference/args.py index 5d50bec1..755a6b26 100644 --- a/cosmos_framework/inference/args.py +++ b/cosmos_framework/inference/args.py @@ -6,7 +6,7 @@ import os from functools import cache from pathlib import Path -from typing import TYPE_CHECKING, Annotated, Any, ClassVar, Literal, Self, cast, override +from typing import TYPE_CHECKING, Annotated, Any, ClassVar, Literal, Self, Sequence, cast, override import pydantic import pynvml @@ -197,6 +197,11 @@ def is_sound_condition(self) -> bool: SOUND_CONDITION_MODEL_MODES: frozenset[ModelMode] = frozenset({ModelMode.AUDIO_IMAGE2VIDEO}) +def is_reasoner_only(sample_overrides: Sequence["OmniSampleOverrides"]) -> bool: + """Return whether every requested sample uses the reasoner-only path.""" + return bool(sample_overrides) and all(sample.sample_meta.model_mode.is_reasoner for sample in sample_overrides) + + class VisionMode(StrEnum): IMAGE = "image" VIDEO = "video" diff --git a/cosmos_framework/inference/args_test.py b/cosmos_framework/inference/args_test.py index 2497eb9d..defc8b7c 100644 --- a/cosmos_framework/inference/args_test.py +++ b/cosmos_framework/inference/args_test.py @@ -19,6 +19,7 @@ OmniSetupOverrides, SoundDataOverrides, _get_nvml_device_memory_info, + is_reasoner_only, ) from cosmos_framework.inference.common.config import structure_config @@ -31,6 +32,31 @@ _GB200_MEMORY_BYTES = 192 * 1024**3 +def test_reasoner_only_detection() -> None: + reasoner = OmniSampleOverrides(model_mode=ModelMode.REASONER) + generator = OmniSampleOverrides(model_mode=ModelMode.TEXT2VIDEO) + + assert is_reasoner_only([reasoner]) + assert is_reasoner_only([reasoner, reasoner]) + assert not is_reasoner_only([reasoner, generator]) + assert not is_reasoner_only([]) + + +def test_reasoner_only_override_disables_vision_tokenizer_in_model_config(tmp_path: Path) -> None: + setup_args = OmniSetupOverrides( + checkpoint_path=DEFAULT_CHECKPOINT_NAME, + output_dir=tmp_path / "outputs", + ).build_setup(world_size=1, local_world_size=1, device_memory_bytes=_H100_MEMORY_BYTES) + + model_dict = structure_config(setup_args.load_model_config_dict(), omegaconf.DictConfig) + assert model_dict.config.load_vision_tokenizer is True + + setup_args.experiment_overrides.append("model.config.load_vision_tokenizer=false") + + model_dict = structure_config(setup_args.load_model_config_dict(), omegaconf.DictConfig) + assert model_dict.config.load_vision_tokenizer is False + + def test_build_parallelism(monkeypatch: pytest.MonkeyPatch): parallelism_args = OmniSetupOverrides( checkpoint_path=DEFAULT_CHECKPOINT_NAME, diff --git a/cosmos_framework/inference/inference.py b/cosmos_framework/inference/inference.py index ecd319f2..6e89047c 100644 --- a/cosmos_framework/inference/inference.py +++ b/cosmos_framework/inference/inference.py @@ -291,8 +291,13 @@ def _compute_num_tokens_for_sample(sample_args: OmniSampleArgs, model: OmniMoTMo w, h = sample_args.vision_size T = sample_args.num_frames - spatial_cf = cast(int, model.tokenizer_vision_gen.spatial_compression_factor) - temporal_cf = cast(int, model.tokenizer_vision_gen.temporal_compression_factor) + vision_tokenizer = model.tokenizer_vision_gen + if vision_tokenizer is None: + spatial_cf = cast(int, model.config.tokenizer.spatial_compression_factor) + temporal_cf = cast(int, model.config.tokenizer.temporal_compression_factor) + else: + spatial_cf = vision_tokenizer.spatial_compression_factor + temporal_cf = vision_tokenizer.temporal_compression_factor patch_spatial: int = model.config.diffusion_expert_config.patch_spatial vae_spatial_downsample = spatial_cf * patch_spatial @@ -1289,7 +1294,10 @@ def _create(cls, setup_args: SetupArgs, **kwargs: Any) -> Self: log.debug(f"Sampler overridden to: {sampler_override}") vae_decode_stream: torch.cuda.Stream | None = None - if setup_args.use_separate_pipeline_vision_decode_gpu: + vision_tokenizer = model.tokenizer_vision_gen + if setup_args.use_separate_pipeline_vision_decode_gpu and vision_tokenizer is None: + log.info("Separate vision decode GPU setup skipped because the generation vision tokenizer is not loaded") + elif setup_args.use_separate_pipeline_vision_decode_gpu: # The CP/CFGP ranks are partitioned into replica-local groups of size # cp_size * cfgp_size. Only the first rank in each group owns separate-VAE # decode work. For example, with cp_size=2 and cfgp_size=1, ranks [0,1] @@ -1309,7 +1317,7 @@ def _create(cls, setup_args: SetupArgs, **kwargs: Any) -> Self: vae_device = torch.device("cuda", vae_device_index) inference_device = torch.device("cuda", torch.cuda.current_device()) vae_decode_stream = torch.cuda.Stream(device=vae_device) - vae = model.tokenizer_vision_gen.model + vae = vision_tokenizer.model vae.device = str(vae_device) vae.model = vae.model.to(device=vae_device) vae.scale = tree_map_only(torch.Tensor, lambda tensor: tensor.to(device=vae_device), vae.scale) diff --git a/cosmos_framework/inference/inference_test.py b/cosmos_framework/inference/inference_test.py index eb0dd094..37e0b96b 100644 --- a/cosmos_framework/inference/inference_test.py +++ b/cosmos_framework/inference/inference_test.py @@ -31,6 +31,60 @@ def test_finalize_data_batch_does_not_mutate_reusable_video_list() -> None: assert "is_preprocessed" not in source_batch +def test_compute_num_tokens_uses_config_when_vision_tokenizer_is_not_loaded() -> None: + from cosmos_framework.inference.inference import _compute_num_tokens_for_sample + + model = SimpleNamespace( + tokenizer_vision_gen=None, + config=SimpleNamespace( + tokenizer=SimpleNamespace( + spatial_compression_factor=16, + temporal_compression_factor=4, + ), + diffusion_expert_config=SimpleNamespace(patch_spatial=2), + ), + ) + sample_args = SimpleNamespace(vision_size=(256, 128), num_frames=9) + + assert _compute_num_tokens_for_sample(sample_args, model) == 96 + + +def test_separate_vision_decode_gpu_is_ignored_without_vision_tokenizer( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Path, +) -> None: + from cosmos_framework.inference import inference + from cosmos_framework.inference.args import DEFAULT_CHECKPOINT_NAME, OmniSetupOverrides + from cosmos_framework.inference.common.args import ConfigFileType + + setup_args = OmniSetupOverrides( + checkpoint_path=DEFAULT_CHECKPOINT_NAME, + output_dir=tmp_path / "outputs", + guardrails=False, + use_separate_pipeline_vision_decode_gpu=True, + ).build_setup(world_size=1, local_world_size=1, device_memory_bytes=80 * 1024**3) + setup_args.config_file_type = ConfigFileType.MODULE + + model = SimpleNamespace( + config=SimpleNamespace( + rectified_flow_inference_config=SimpleNamespace(scheduler_type=setup_args.sampler), + ), + tokenizer_vision_gen=None, + ) + load_model = Mock(return_value=SimpleNamespace(model=model)) + device_count = Mock(side_effect=AssertionError("separate VAE setup should not inspect CUDA devices")) + monkeypatch.setattr(inference, "_download_on_rank0", lambda _download: tmp_path) + monkeypatch.setattr(inference.Cosmos3OmniModel, "from_pretrained_dcp", load_model) + monkeypatch.setattr(inference.torch.cuda, "device_count", device_count) + + pipe = inference.OmniInference._create(setup_args, guardrails=None, _timer=None) + + assert pipe.model is model + assert pipe.vae_decode_stream is None + load_model.assert_called_once() + device_count.assert_not_called() + + def _make_v2v_sample_args(**overrides: Any) -> SimpleNamespace: """v2v ``OmniSampleArgs`` stand-in for ``get_sample_data`` tests.""" from cosmos_framework.inference.args import ModelMode, NegativeMetadataMode diff --git a/cosmos_framework/model/generator/omni_mot_model.py b/cosmos_framework/model/generator/omni_mot_model.py index c2bcaf5f..8dfd9b36 100644 --- a/cosmos_framework/model/generator/omni_mot_model.py +++ b/cosmos_framework/model/generator/omni_mot_model.py @@ -165,13 +165,19 @@ def set_up_tokenizers(self) -> None: self.llm_special_tokens = special_tokens self.llm_special_tokens["eos_token_id"] = vlm_tokenizer.eos_token_id - # 2. Vision tokenizer (images/videos) for generation. - self.tokenizer_vision_gen: VideoTokenizerInterface = lazy_instantiate(self.config.tokenizer) - assert self.tokenizer_vision_gen.latent_ch == self.config.state_ch, ( - f"vision tokenizer latent_ch {self.tokenizer_vision_gen.latent_ch} != state_shape {self.config.state_ch}" - ) - if hasattr(self.tokenizer_vision_gen, "reset_dtype"): - self.tokenizer_vision_gen.reset_dtype() + # 2. Vision tokenizer (images/videos) for generation. Reasoner-only + # inference does not encode or decode generation latents, so it can + # leave the VAE unloaded. + self.tokenizer_vision_gen: VideoTokenizerInterface | None = None + if self.config.load_vision_tokenizer: + self.tokenizer_vision_gen = lazy_instantiate(self.config.tokenizer) + assert self.tokenizer_vision_gen.latent_ch == self.config.state_ch, ( + f"vision tokenizer latent_ch {self.tokenizer_vision_gen.latent_ch} != state_shape {self.config.state_ch}" + ) + if hasattr(self.tokenizer_vision_gen, "reset_dtype"): + self.tokenizer_vision_gen.reset_dtype() + else: + log.info("Vision tokenizer initialization skipped") # 3. Sound/audio tokenizer (optional) if self.config.sound_gen: @@ -218,7 +224,11 @@ def build_net(self, dtype: torch.dtype, *, lora_enabled: bool | None = None) -> timestep_scale=1.0 / float(num_train_timesteps) * self.config.diffusion_expert_config.timestep_range, action_dim=self.config.max_action_dim, num_embodiment_domains=self.config.num_embodiment_domains, - temporal_compression_factor_vision=self.tokenizer_vision_gen.temporal_compression_factor, + temporal_compression_factor_vision=( + self.tokenizer_vision_gen.temporal_compression_factor + if self.tokenizer_vision_gen is not None + else self.config.tokenizer.temporal_compression_factor + ), natten_parameter_list=self.config.natten_parameter_list, video_temporal_causal=self.config.video_temporal_causal, # Sound generation parameters diff --git a/cosmos_framework/model/generator/omni_mot_model_test.py b/cosmos_framework/model/generator/omni_mot_model_test.py new file mode 100644 index 00000000..440ed417 --- /dev/null +++ b/cosmos_framework/model/generator/omni_mot_model_test.py @@ -0,0 +1,69 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +from types import SimpleNamespace +from unittest.mock import Mock + +import pytest + + +def test_reasoner_only_setup_skips_vision_tokenizer(monkeypatch: pytest.MonkeyPatch) -> None: + from cosmos_framework.model.generator import omni_mot_model + + vlm_tokenizer = SimpleNamespace(eos_token_id=42) + vlm_processor = SimpleNamespace(tokenizer=vlm_tokenizer) + vlm_config = SimpleNamespace(tokenizer="vlm-tokenizer-config") + vision_config = SimpleNamespace(temporal_compression_factor=4) + config = SimpleNamespace( + load_vision_tokenizer=False, + sound_gen=False, + tokenizer=vision_config, + vlm_config=vlm_config, + ) + instantiated = [] + + def _instantiate(candidate): + instantiated.append(candidate) + return vlm_processor + + monkeypatch.setattr(omni_mot_model, "lazy_instantiate", _instantiate) + monkeypatch.setattr(omni_mot_model, "add_special_tokens", lambda tokenizer: (tokenizer, {})) + + model = SimpleNamespace(config=config) + omni_mot_model.OmniMoTModel.set_up_tokenizers(model) + + assert instantiated == [vlm_config.tokenizer] + assert model.tokenizer_vision_gen is None + assert model.tokenizer_sound_gen is None + + +def test_default_setup_loads_vision_tokenizer(monkeypatch: pytest.MonkeyPatch) -> None: + from cosmos_framework.model.generator import omni_mot_model + + vlm_tokenizer = SimpleNamespace(eos_token_id=42) + vlm_processor = SimpleNamespace(tokenizer=vlm_tokenizer) + vision_tokenizer = SimpleNamespace(latent_ch=48, reset_dtype=Mock()) + vlm_config = SimpleNamespace(tokenizer="vlm-tokenizer-config") + vision_config = SimpleNamespace(temporal_compression_factor=4) + config = SimpleNamespace( + load_vision_tokenizer=True, + sound_gen=False, + state_ch=48, + tokenizer=vision_config, + vlm_config=vlm_config, + ) + + def _instantiate(candidate): + if candidate == vlm_config.tokenizer: + return vlm_processor + assert candidate is vision_config + return vision_tokenizer + + monkeypatch.setattr(omni_mot_model, "lazy_instantiate", _instantiate) + monkeypatch.setattr(omni_mot_model, "add_special_tokens", lambda tokenizer: (tokenizer, {})) + + model = SimpleNamespace(config=config) + omni_mot_model.OmniMoTModel.set_up_tokenizers(model) + + assert model.tokenizer_vision_gen is vision_tokenizer + vision_tokenizer.reset_dtype.assert_called_once_with() diff --git a/cosmos_framework/scripts/inference.py b/cosmos_framework/scripts/inference.py index 63c5cf60..d363c47c 100644 --- a/cosmos_framework/scripts/inference.py +++ b/cosmos_framework/scripts/inference.py @@ -12,7 +12,7 @@ import pydantic import tyro -from cosmos_framework.inference.args import OmniSetupOverrides +from cosmos_framework.inference.args import OmniSetupOverrides, is_reasoner_only from cosmos_framework.inference.common.args import SampleOutputs, SetupOverrides, tyro_cli from cosmos_framework.inference.common.init import init_output_dir from cosmos_framework.utils import log @@ -44,6 +44,9 @@ def inference(args: InferenceArgs): args.input_files, overrides=setup_args.sample_overrides ) log.info(f"Loaded {len(sample_overrides_list)} samples") + if is_reasoner_only(sample_overrides_list): + setup_args.experiment_overrides.append("model.config.load_vision_tokenizer=false") + log.info("Reasoner-only inputs detected; generation vision tokenizer will not be loaded") for sample_overrides in sample_overrides_list: assert sample_overrides.name sample_overrides.output_dir = setup_args.output_dir / sample_overrides.name