Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions cosmos_framework/configs/base/defaults/model_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()

Expand Down
7 changes: 6 additions & 1 deletion cosmos_framework/inference/args.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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"
Expand Down
26 changes: 26 additions & 0 deletions cosmos_framework/inference/args_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@
OmniSetupOverrides,
SoundDataOverrides,
_get_nvml_device_memory_info,
is_reasoner_only,
)
from cosmos_framework.inference.common.config import structure_config

Expand All @@ -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,
Expand Down
16 changes: 12 additions & 4 deletions cosmos_framework/inference/inference.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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]
Expand All @@ -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)
Expand Down
54 changes: 54 additions & 0 deletions cosmos_framework/inference/inference_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
26 changes: 18 additions & 8 deletions cosmos_framework/model/generator/omni_mot_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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
Expand Down
69 changes: 69 additions & 0 deletions cosmos_framework/model/generator/omni_mot_model_test.py
Original file line number Diff line number Diff line change
@@ -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()
5 changes: 4 additions & 1 deletion cosmos_framework/scripts/inference.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down