diff --git a/.file_mapping.json b/.file_mapping.json index b96ddff3..68f42c2e 100644 --- a/.file_mapping.json +++ b/.file_mapping.json @@ -1,7 +1,7 @@ { - "_source_commit": "461431bae40d4a8ccf27d812fc9d1757fe8d3b96-dirty", - "_dest_commit": "0460be81f16883aa380e716dc6f58c1189481172", - "_generated_at": "2026-09-23T04:47:38Z", + "_source_commit": "358182a3a006975c25fc079fac9ff80e704077e2-dirty", + "_dest_commit": "cf5d68c00d97ccd2480a2320ed652b92dec63102", + "_generated_at": "2026-09-25T05:52:16Z", "files": { "imaginaire/__init__.py": "cosmos_framework/__init__.py", "imaginaire/attention/__init__.py": "cosmos_framework/model/attention/__init__.py", @@ -293,6 +293,8 @@ "projects/cosmos3/cosmos3/datasets/augmentors/cropping.py": "cosmos_framework/data/generator/augmentors/cropping.py", "projects/cosmos3/cosmos3/datasets/augmentors/duration_fps_text_timestamps.py": "cosmos_framework/data/generator/augmentors/duration_fps_text_timestamps.py", "projects/cosmos3/cosmos3/datasets/augmentors/duration_fps_text_timestamps_test.py": "cosmos_framework/data/generator/augmentors/duration_fps_text_timestamps_test.py", + "projects/cosmos3/cosmos3/datasets/augmentors/first_caption_paragraph.py": "cosmos_framework/data/generator/augmentors/first_caption_paragraph.py", + "projects/cosmos3/cosmos3/datasets/augmentors/first_caption_paragraph_test.py": "cosmos_framework/data/generator/augmentors/first_caption_paragraph_test.py", "projects/cosmos3/cosmos3/datasets/augmentors/idle_frames_text_info.py": "cosmos_framework/data/generator/augmentors/idle_frames_text_info.py", "projects/cosmos3/cosmos3/datasets/augmentors/image_editing_transform.py": "cosmos_framework/data/generator/augmentors/image_editing_transform.py", "projects/cosmos3/cosmos3/datasets/augmentors/image_editing_transform_test.py": "cosmos_framework/data/generator/augmentors/image_editing_transform_test.py", @@ -470,6 +472,8 @@ "projects/cosmos3/cosmos3/models/utils/load_balancing_stats.py": "cosmos_framework/model/generator/utils/load_balancing_stats.py", "projects/cosmos3/cosmos3/models/utils/memory.py": "cosmos_framework/model/generator/utils/memory.py", "projects/cosmos3/cosmos3/models/utils/moe_utils.py": "cosmos_framework/model/generator/utils/moe_utils.py", + "projects/cosmos3/cosmos3/models/utils/rig_view_embedding.py": "cosmos_framework/model/generator/utils/rig_view_embedding.py", + "projects/cosmos3/cosmos3/models/utils/rig_view_embedding_test.py": "cosmos_framework/model/generator/utils/rig_view_embedding_test.py", "projects/cosmos3/cosmos3/models/utils/safetensors_loader.py": "cosmos_framework/model/generator/utils/safetensors_loader.py", "projects/cosmos3/cosmos3/models/utils/safetensors_loader_indexed_test.py": "cosmos_framework/model/generator/utils/safetensors_loader_indexed_test.py", "projects/cosmos3/cosmos3/models/utils/safetensors_loader_test.py": "cosmos_framework/model/generator/utils/safetensors_loader_test.py", @@ -569,6 +573,7 @@ "projects/cosmos3/cosmos3/utils/reasoner/pretrained_models_downloader_test.py": "cosmos_framework/utils/generator/reasoner/pretrained_models_downloader_test.py", "projects/cosmos3/cosmos3/utils/reasoner/true_packing.py": "cosmos_framework/utils/generator/reasoner/true_packing.py", "projects/cosmos3/cosmos3/utils/source_video_timing.py": "cosmos_framework/utils/generator/source_video_timing.py", + "projects/cosmos3/cosmos3/utils/spatial_patch.py": "cosmos_framework/utils/generator/spatial_patch.py", "projects/cosmos3/cosmos3/utils/video_frame_sampling.py": "cosmos_framework/utils/generator/video_frame_sampling.py", "projects/cosmos3/cosmos3/utils/video_preprocess.py": "cosmos_framework/utils/generator/video_preprocess.py", "projects/cosmos3/cosmos3/utils/video_source_metadata.py": "cosmos_framework/utils/generator/video_source_metadata.py", @@ -580,6 +585,8 @@ "projects/cosmos3/interactive/models/mot/context_parallel_test.py": "cosmos_framework/model/generator/mot/causal_context_parallel_test.py", "projects/cosmos3/interactive/models/mot/cosmos3_vfm_network.py": "cosmos_framework/model/generator/mot/causal_cosmos3_vfm_network.py", "projects/cosmos3/interactive/models/mot/flex_attention.py": "cosmos_framework/model/generator/mot/causal_flex_attention.py", + "projects/cosmos3/interactive/models/mot/maskless_attention.py": "cosmos_framework/model/generator/mot/maskless_attention.py", + "projects/cosmos3/interactive/models/mot/merge_attention.py": "cosmos_framework/model/generator/mot/merge_attention.py", "projects/cosmos3/interactive/models/mot/post_saturation/__init__.py": "cosmos_framework/model/generator/mot/post_saturation/__init__.py", "projects/cosmos3/interactive/models/mot/post_saturation/cuda_graph.py": "cosmos_framework/model/generator/mot/post_saturation/cuda_graph.py", "projects/cosmos3/interactive/models/mot/post_saturation/installer.py": "cosmos_framework/model/generator/mot/post_saturation/installer.py", diff --git a/cosmos_framework/callbacks/every_n_draw_sample.py b/cosmos_framework/callbacks/every_n_draw_sample.py index bd086cf7..1c876a34 100644 --- a/cosmos_framework/callbacks/every_n_draw_sample.py +++ b/cosmos_framework/callbacks/every_n_draw_sample.py @@ -45,9 +45,9 @@ class WandbAnimation(NamedTuple): # What a clip is previewed as in W&B: one still grid, one panel per sampled -# frame, or an animation of every frame. +# frame, an animation of every frame, or direct interactive point-cloud media. WandbClipPreview = Literal["grid", "frames", "animation"] -WandbMedia = str | dict[str, str | WandbAnimation] | WandbAnimation +WandbMedia = str | WandbAnimation | list[wandb.Object3D] | dict[str, "WandbMedia"] def resize_image(image: torch.Tensor, size: int = 1024) -> torch.Tensor: @@ -292,10 +292,10 @@ def _decode_transfer_pixel_row( ) -> torch.Tensor | None: # [V,C,F,H,W] or None """Decode one latent into a host-side pixel row without display conversion. - ``decode`` defaults to the model's main VAE; a joint camera + LiDAR sample passes - ``model.decode_lidar`` for range clips because the two streams have separate VAEs. - LiDAR callers pass ``decode_per_view=False``: V0 and V1 are both 1x temporal on a - single range view, not 4x camera-major WAN clips. + ``decode`` defaults to the model's main VAE; a joint camera + LiDAR sample + passes ``model.decode_lidar`` for those clips because the streams have + separate VAEs. LiDAR callers pass ``decode_per_view=False``: the rangemap + is 1x temporal on a single view, not 4x camera-major WAN clips. """ decode = decode if decode is not None else model.decode if decode_per_view is None: @@ -565,6 +565,10 @@ def _add_wandb_media( ) -> None: if media is None: return + if isinstance(media, list): + if media: + info[key_prefix] = media + return if isinstance(media, WandbAnimation): info[key_prefix] = wandb.Video(media.path, caption=f"{caption} | {media.num_frames} frames") return @@ -747,6 +751,13 @@ class EveryNDrawSample(EveryN): use_negative_prompt (bool, optional): whether to use negative prompt. Defaults to False. fps (int, optional): frames per second when saving the video. Defaults to 16. wandb_log_image_size (int, optional): max side length for W&B JPEG/GIF previews. Defaults to 1024. + log_lidar_point_clouds (bool, optional): also log separate real/generated LiDAR point-cloud panels. Defaults to True. + point_cloud_max_frames (int | None, optional): preview scans per row, starting at the first generated scan. + None keeps all scans, including labelled conditioning inputs. Defaults to 3. + point_cloud_max_points (int, optional): points per cloud, at most 300000. Defaults to 300000. + point_cloud_color_by (str, optional): BEV height colors or grayscale intensity. Defaults to "height". + point_cloud_include_conditioned_frames (bool, optional): sample across the full clip, including labelled + conditioning inputs, when bounding the preview scan count. Defaults to False. """ def __init__( @@ -1005,6 +1016,9 @@ def _generate_transfer_latents( if plans and not getattr(plans[0], "has_lidar", False): generation_batch.pop("lidar", None) generation_batch.pop("num_lidar_items_per_sample", None) + if plans and not getattr(plans[0], "has_radar", False): + generation_batch.pop("radar", None) + generation_batch.pop("num_radar_items_per_sample", None) # Preserve an explicit caller-provided prompt. Otherwise restore the exact prompt recorded # by the training tokenizer before train-sample inference re-tokenizes the raw caption. text_system_prompt = generation_batch.pop(TEXT_SYSTEM_PROMPT_KEY, None) diff --git a/cosmos_framework/configs/base/defaults/model_config.py b/cosmos_framework/configs/base/defaults/model_config.py index 777f8a59..0be08be2 100644 --- a/cosmos_framework/configs/base/defaults/model_config.py +++ b/cosmos_framework/configs/base/defaults/model_config.py @@ -46,10 +46,25 @@ class DiffusionExpertConfig: # Enabled by default enable_sound_modality_embedding: bool = True + # Zero disables physical rig embeddings. When enabled, N >= 2 reserves camera IDs + # 0..N-2 and LiDAR ID N-1, even for camera-only inputs. N is the sensor vocabulary + # size for the model, not the number of sensors selected in a sample. + # Examples for the MADS rig (N=12, camera IDs 0..10, LiDAR ID 11): + # - Camera only: selecting camera 8 uses row 8; row 11 is unused. + # - LiDAR only: uses row 11; no vision_view_ids are needed. + # - Camera + LiDAR: selected cameras use their physical IDs and LiDAR uses row 11. + num_view_embeddings: int = 0 + patch_spatial: int = 2 + # None preserves the shared camera/LiDAR/radar patch size used by existing checkpoints. + # An int selects square patches; (height, width) selects rectangular patches. + # Any is required here because OmegaConf rejects unions containing tuples. + lidar_patch_spatial_hw: Any = None + radar_patch_spatial_hw: Any = None max_vae_latent_side_after_patchify: int = ( 52 # Max h/w of the VAE latent after patchification; 52 -> up to ~1664px square (52*32). Was 20 (=640px). ) + # Vision/action/sound position information is always provided through # Qwen3VL-style 3D mRoPE attention IDs. enable_fps_modulation: bool = False @@ -72,6 +87,7 @@ class DiffusionExpertConfig: class RectifiedFlowTrainingConfig: shift: Any = 5 # Training time shift. If dict, maps resolution (str) to shift value (int) shift_image: Any | None = None # Image-specific shift; None inherits shift + shift_lidar: int | None = None # LiDAR-only batches; None preserves the shared vision schedule use_dynamic_shift: bool = False # Whether to use dynamic shifting train_time_image_distribution: str = "logitnormal" # Training time distribution for images train_time_video_distribution: str = "logitnormal" # Training time distribution for videos @@ -82,6 +98,7 @@ class RectifiedFlowTrainingConfig: image_loss_scale: float | None = None # If set, overrides loss_scale for images sound_loss_scale: float | None = None # If set, overrides loss_scale for sound lidar_loss_scale: float | None = None # If set, overrides loss_scale for lidar + radar_loss_scale: float | None = None # If set, overrides loss_scale for radar use_discrete_rf: bool = False # Whether to use discrete formulation of rectified flow # user: please adjust this value according to loss_scale to balance the action loss with the video loss. @@ -110,6 +127,12 @@ class RectifiedFlowTrainingConfig: # False by default to preserve legacy loss magnitudes; enable for AR/DF training. normalize_loss_by_active: bool = False + # None preserves the legacy strategy-dependent item mean: teacher forcing + # excludes fully conditioned items, other strategies include them. Set + # False to match the bidirectional teacher's denominator, including zero-loss + # control items; this is independent of per-item active-token normalization. + exclude_fully_conditioned_items: bool | None = None + # Sample-level (vs rank-level) loss averaging for the vision modality. # # By default the vision loss on each rank is a mean over that rank's samples, and @@ -194,6 +217,20 @@ class OmniMoTModelConfig: seconds, which is what puts the two sensors' latents on one mRoPE time axis. """ + radar_state_ch: int | None = None + """Radar VAE latent channel count, i.e. the width of the network's radar heads. + + The radar VAE is as wide as the LiDAR one (128), but the two are sized from their own + fields because nothing ties the two sensors' widths together. + """ + + radar_fps: float | None = None + """Cycle rate in Hz of radar items, the counterpart of ``lidar_fps``. + + Radar cycles at 20 Hz against LiDAR's 10 Hz, and its VAE does not compress time, so + this is what places a radar scan on the mRoPE time axis the other streams share. + """ + net: LazyDict = None ema: EMAConfig = EMAConfig() diff --git a/cosmos_framework/configs/base/defaults/multiview_attention.py b/cosmos_framework/configs/base/defaults/multiview_attention.py index 89b71d2f..740f5012 100644 --- a/cosmos_framework/configs/base/defaults/multiview_attention.py +++ b/cosmos_framework/configs/base/defaults/multiview_attention.py @@ -184,6 +184,13 @@ class MultiviewAttentionMaskConfig: # ._multiview_pair_predicate. lidar_attends_captions: bool = True + # The same switch for the radar stream of a joint camera + radar pack, read the same way. + # Radar has a stronger prior claim to needing it than LiDAR does: the captions describe what + # the cameras see, and a BEV occupancy grid shares even less of that vocabulary than a range + # image does. Independent of ``lidar_attends_captions`` so a three-sensor pack can drop text + # for one stream and keep it for the other. + radar_attends_captions: bool = True + @attrs.define(slots=False) class MultiviewAttentionConfig: diff --git a/cosmos_framework/configs/base/defaults/replay_attention.py b/cosmos_framework/configs/base/defaults/replay_attention.py index 0d4065f4..fb894e9d 100644 --- a/cosmos_framework/configs/base/defaults/replay_attention.py +++ b/cosmos_framework/configs/base/defaults/replay_attention.py @@ -12,7 +12,7 @@ AttentionScope, ) -TeacherForcingKVImplementation = Literal["multiview_flex_kv", "singleview_threeway_kv"] +TeacherForcingKVImplementation = Literal["multiview_flex_kv", "multiview_maskless_kv", "singleview_threeway_kv"] TEACHER_FORCING_KV_IMPLEMENTATIONS = get_args(TeacherForcingKVImplementation) TeacherForcingControlVisibility = Literal["global", "causal", "current"] diff --git a/cosmos_framework/data/generator/augmentors/first_caption_paragraph.py b/cosmos_framework/data/generator/augmentors/first_caption_paragraph.py new file mode 100644 index 00000000..e34e7850 --- /dev/null +++ b/cosmos_framework/data/generator/augmentors/first_caption_paragraph.py @@ -0,0 +1,34 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +"""Keep the opening temporal paragraph for a short clip at the caption chunk's start.""" + +from typing import Any + +from cosmos_framework.data.imaginaire.webdataset.augmentors.augmentor import Augmentor +from cosmos_framework.utils import log +from cosmos_framework.data.generator.multiview.caption_format import first_caption_paragraph + + +class FirstCaptionParagraph(Augmentor): + """Retain the first timestamped paragraph of a plain-text Tier-1 caption. + + Timestamp boundaries also handle consecutive paragraphs separated by only + one newline. Untimestamped captions keep their first blank-line paragraph. + Run before duration/resolution metadata is appended and before CFG dropout. + """ + + def __init__(self) -> None: + super().__init__(input_keys=["ai_caption"], output_keys=["ai_caption"]) + + def __call__(self, data: dict[str, Any]) -> dict[str, Any] | None: + caption = data["ai_caption"] + if not isinstance(caption, str): + log.warning( + f"FirstCaptionParagraph: dropping non-text caption ({type(caption).__name__}); " + f"sample_key={data.get('__key__')}, url={data.get('__url__')}", + rank0_only=False, + ) + return None + data["ai_caption"] = first_caption_paragraph(caption) + return data diff --git a/cosmos_framework/data/generator/augmentors/first_caption_paragraph_test.py b/cosmos_framework/data/generator/augmentors/first_caption_paragraph_test.py new file mode 100644 index 00000000..500e2a0b --- /dev/null +++ b/cosmos_framework/data/generator/augmentors/first_caption_paragraph_test.py @@ -0,0 +1,35 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +"""Temporal-caption trimming must not retain events outside the short training window.""" + +import pytest + +from cosmos_framework.data.generator.augmentors.first_caption_paragraph import FirstCaptionParagraph + +pytestmark = [pytest.mark.L0, pytest.mark.CPU] + + +@pytest.mark.parametrize( + ("caption", "expected"), + [ + ("[0-2s] The car approaches.\n\n[2-5s] It turns.", "[0-2s] The car approaches."), + ("[0-1.8s] The car approaches.\n[1.8-5s] It turns.", "[0-1.8s] The car approaches."), + ("[0s – 2.5s] The car approaches, [2.5s – 5s] It turns.", "[0s – 2.5s] The car approaches"), + ("Overview.\n\n[0-3s] The car\napproaches.\n\n[3-5s] It turns.", "[0-3s] The car\napproaches."), + (" A car approaches.\n\nIt turns later. ", "A car approaches."), + (" [0-2s] A car approaches. ", "[0-2s] A car approaches."), + ("", ""), + ], +) +def test_keeps_only_the_opening_caption_paragraph(caption: str, expected: str) -> None: + data = {"ai_caption": caption, "__key__": "clip"} + result = FirstCaptionParagraph()(data) + assert result is not None + assert result["ai_caption"] == expected + assert result["__key__"] == "clip" + + +@pytest.mark.parametrize("caption", [{"scene": "not a Tier-1 paragraph"}, None, ["caption"]]) +def test_drops_nontext_captions(caption: object) -> None: + assert FirstCaptionParagraph()({"ai_caption": caption, "__key__": "clip"}) is None diff --git a/cosmos_framework/data/generator/augmentors/pkl_to_media.py b/cosmos_framework/data/generator/augmentors/pkl_to_media.py index 4b0fda79..bc29014e 100644 --- a/cosmos_framework/data/generator/augmentors/pkl_to_media.py +++ b/cosmos_framework/data/generator/augmentors/pkl_to_media.py @@ -225,7 +225,9 @@ def _bytes_to_video_frames(self, video_bytes: bytes, identifier: str = "video") data=video_bytes, **self.video_decoder_params, ) - result["videos"] = tensor_to_pil_images(result["videos"]) # 3,T,H,W -> list of PIL images + result["videos"] = tensor_to_pil_images( + result["videos"], channels_first=True + ) # 3,T,H,W -> list of PIL images if result is not None: return result else: diff --git a/cosmos_framework/data/generator/augmentors/reasoner/bytes_to_media.py b/cosmos_framework/data/generator/augmentors/reasoner/bytes_to_media.py index 5d9d6fdf..f3f4c913 100644 --- a/cosmos_framework/data/generator/augmentors/reasoner/bytes_to_media.py +++ b/cosmos_framework/data/generator/augmentors/reasoner/bytes_to_media.py @@ -257,7 +257,9 @@ def _bytes_to_video_frames( return None if self.video_timestamp_mode == "source_pts": validate_source_video_timing(result.get(SOURCE_VIDEO_TIMING_KEY), result["videos"].shape[1]) - result["videos"] = tensor_to_pil_images(result["videos"]) # 3,T,H,W -> list of PIL images + result["videos"] = tensor_to_pil_images( + result["videos"], channels_first=True + ) # 3,T,H,W -> list of PIL images return result except Exception as e: if self.video_timestamp_mode == "source_pts": diff --git a/cosmos_framework/data/generator/joint_dataloader.py b/cosmos_framework/data/generator/joint_dataloader.py index 0f9ca647..7d1cb445 100644 --- a/cosmos_framework/data/generator/joint_dataloader.py +++ b/cosmos_framework/data/generator/joint_dataloader.py @@ -33,6 +33,7 @@ normalize_uniae_chunk_frames, ) from cosmos_framework.utils.generator.data_utils import read_positive_int_metadata +from cosmos_framework.utils.generator.spatial_patch import normalize_spatial_patch_hw _BATCH_TIMING_KEYS = { "_worker_batch_time", @@ -174,6 +175,10 @@ def custom_collate_fn(batch: list[dict[str, Any]] | dict[str, Any]) -> dict[str, # Like "video": a per-sample list of range clips, which default_collate would try to # stack even though the two sensors' clips differ in length and resolution. "lidar", + # The BEV pair is [control, target]. Outside this set it goes to default_collate, + # which transposes it into a list of stacked tensors rather than failing; the + # per-sample split then yields one 5-D tensor where a list of clips is expected. + "radar", DROP_SAMPLE_KEY, DROP_SAMPLE_REASON_KEY, *_ACTION_SAMPLER_METADATA_KEYS, @@ -450,6 +455,8 @@ def __init__( max_samples_per_batch: int | None, lidar_spatial_compression: Sequence[int] | None = None, lidar_temporal_compression_factor: int | None = None, + radar_spatial_compression: Sequence[int] | None = None, + radar_temporal_compression_factor: int | None = None, sound_latent_fps: float = 0, audio_sample_rate: int = 48000, prewarm: bool = True, @@ -461,6 +468,8 @@ def __init__( lazy_initialize_child_iterators: bool = False, iteration_time_budget: IterationTimeBudgetConfig | None = None, forkserver_preload_modules: list[str] | None = None, + lidar_patch_spatial_hw: int | tuple[int, int] | None = None, + radar_patch_spatial_hw: int | tuple[int, int] | None = None, ) -> None: """ Initialize the JointDataLoader with multiple datasets. @@ -479,12 +488,20 @@ def __init__( tokenizer_spatial_compression_factor: The spatial compression factor of the tokenizer. tokenizer_temporal_compression_factor: The temporal compression factor of the tokenizer. patch_spatial: Spatial pathification factor. + lidar_patch_spatial_hw: LiDAR patch side or (height, width); None inherits patch_spatial. + radar_patch_spatial_hw: Radar patch side or (height, width); None inherits patch_spatial. max_samples_per_batch: Max number of samples per packed batch (alternative to max_sequence_length). lidar_spatial_compression: ``(height, width)`` compression of the LiDAR VAE. Required only for streams whose samples carry a ``lidar`` key, whose clips are costed with the LiDAR VAE rather than the camera's — the two compress time differently (4x versus 1x), and an item costed with the wrong factor silently over-packs the batch. lidar_temporal_compression_factor: Temporal compression of the LiDAR VAE. + radar_spatial_compression: ``(height, width)`` compression of the radar VAE. Required only + for streams whose samples carry a ``radar`` key. Radar cannot borrow the LiDAR pair + even on a run that carries both sensors: the two VAEs read different grids (a square + BEV versus a range image) and the recipes pick their tokenizer versions + independently, so one pair of factors cannot price both streams. + radar_temporal_compression_factor: Temporal compression of the radar VAE. sound_latent_fps: Sound tokenizer latent rate in Hz (e.g. 25). If 0, sound tokens are not counted. audio_sample_rate: Audio sample rate in Hz (e.g. 48000). Used with sound_latent_fps to estimate sound token count. @@ -542,7 +559,32 @@ def __init__( raise ValueError( f"lidar_temporal_compression_factor must be positive, got {self.lidar_temporal_compression_factor}" ) + self.radar_spatial_compression = ( + tuple(int(factor) for factor in radar_spatial_compression) + if radar_spatial_compression is not None + else None + ) + if self.radar_spatial_compression is not None and ( + len(self.radar_spatial_compression) != 2 or any(factor <= 0 for factor in self.radar_spatial_compression) + ): + raise ValueError( + "radar_spatial_compression must contain two positive factors " + f"(height, width), got {self.radar_spatial_compression}" + ) + self.radar_temporal_compression_factor = ( + int(radar_temporal_compression_factor) if radar_temporal_compression_factor is not None else None + ) + if self.radar_temporal_compression_factor is not None and self.radar_temporal_compression_factor <= 0: + raise ValueError( + f"radar_temporal_compression_factor must be positive, got {self.radar_temporal_compression_factor}" + ) self.patch_spatial = patch_spatial + self.lidar_patch_spatial_hw: tuple[int, int] = normalize_spatial_patch_hw( + patch_spatial if lidar_patch_spatial_hw is None else lidar_patch_spatial_hw + ) + self.radar_patch_spatial_hw: tuple[int, int] = normalize_spatial_patch_hw( + patch_spatial if radar_patch_spatial_hw is None else radar_patch_spatial_hw + ) self.max_sequence_length = max_sequence_length self.max_samples_per_batch = max_samples_per_batch self.sound_latent_fps = sound_latent_fps @@ -668,12 +710,39 @@ def _num_lidar_tokens(self, data_batch: Mapping[str, Any]) -> int: num_tokens = 0 for clip in clips: _, T, H, W = clip.shape - patch_h = math.ceil(H // spatial_h / self.patch_spatial) - patch_w = math.ceil(W // spatial_w / self.patch_spatial) + patch_h = math.ceil(H // spatial_h / self.lidar_patch_spatial_hw[0]) + patch_w = math.ceil(W // spatial_w / self.lidar_patch_spatial_hw[1]) latent_t = 1 + (T - 1) // self.lidar_temporal_compression_factor num_tokens += patch_h * patch_w * latent_t return num_tokens + def _num_radar_tokens(self, data_batch: Mapping[str, Any]) -> int: + """Cost the sample's radar BEV clips with the radar VAE's own compression. + + Same arithmetic as the sweeps over a square BEV grid rather than a range image: the + radar VAE does not compress time either, so one scan is one latent frame and the + camera's 4x would undercount a clip fourfold. Radar is the more expensive of the two + sensors per second of clip -- it cycles at ~20 Hz against LiDAR's 10 and each scan + carries more latents -- so an underpriced clip over-packs the batch by more here. + """ + clips = data_batch.get("radar") + if not clips: + return 0 + if self.radar_spatial_compression is None or self.radar_temporal_compression_factor is None: + raise ValueError( + "This batch carries a radar stream, but the loader has no radar compression factors. " + "Set radar_spatial_compression and radar_temporal_compression_factor." + ) + spatial_h, spatial_w = self.radar_spatial_compression + num_tokens = 0 + for clip in clips: + _, T, H, W = clip.shape + patch_h = math.ceil(H // spatial_h / self.radar_patch_spatial_hw[0]) + patch_w = math.ceil(W // spatial_w / self.radar_patch_spatial_hw[1]) + latent_t = 1 + (T - 1) // self.radar_temporal_compression_factor + num_tokens += patch_h * patch_w * latent_t + return num_tokens + def _normalize_uniae_chunk_frames( self, uniae_chunk_frames: int | Mapping[str, int] | None ) -> int | dict[str, int] | None: @@ -904,6 +973,11 @@ def _compute_token_split_per_sample(self, data_batch: dict) -> tuple[int, int]: # gen_tokens so the iteration-time cost model sees the sweeps. gen_tokens += self._num_lidar_tokens(data_batch) + # Radar part: a third VAE and a third pair of factors, charged the same way. A joint + # camera + radar sample carries no sweeps and a joint camera + LiDAR one no scans, so + # each of the two calls prices nothing on the other's recipe. + gen_tokens += self._num_radar_tokens(data_batch) + # Action part: each action time step is 1 token. # Action tensor shape is (T_action, D) per sample; stored as a single-element list. if "action" in data_batch: @@ -1123,6 +1197,8 @@ def __init__( max_samples_per_batch: int | None = None, lidar_spatial_compression: Sequence[int] | None = None, lidar_temporal_compression_factor: int | None = None, + radar_spatial_compression: Sequence[int] | None = None, + radar_temporal_compression_factor: int | None = None, sound_latent_fps: float = 0, audio_sample_rate: int = 48000, seed: int | None = 42, @@ -1138,6 +1214,8 @@ def __init__( iteration_time_budget: IterationTimeBudgetConfig | None = None, token_mix_control: TokenMixControlConfig | None = None, forkserver_preload_modules: list[str] | None = None, + lidar_patch_spatial_hw: int | tuple[int, int] | None = None, + radar_patch_spatial_hw: int | tuple[int, int] | None = None, ) -> None: if async_batch_building_timeout_s <= 0: raise ValueError(f"async_batch_building_timeout_s must be positive, got {async_batch_building_timeout_s}.") @@ -1178,6 +1256,8 @@ def __init__( max_samples_per_batch, lidar_spatial_compression=lidar_spatial_compression, lidar_temporal_compression_factor=lidar_temporal_compression_factor, + radar_spatial_compression=radar_spatial_compression, + radar_temporal_compression_factor=radar_temporal_compression_factor, sound_latent_fps=sound_latent_fps, audio_sample_rate=audio_sample_rate, prewarm=prewarm, @@ -1189,6 +1269,8 @@ def __init__( lazy_initialize_child_iterators=lazy_initialize_child_iterators, iteration_time_budget=iteration_time_budget, forkserver_preload_modules=forkserver_preload_modules, + lidar_patch_spatial_hw=lidar_patch_spatial_hw, + radar_patch_spatial_hw=radar_patch_spatial_hw, ) self.seed = seed diff --git a/cosmos_framework/data/generator/multiview/caption_chunks.py b/cosmos_framework/data/generator/multiview/caption_chunks.py index 214f4297..15bbe5e8 100644 --- a/cosmos_framework/data/generator/multiview/caption_chunks.py +++ b/cosmos_framework/data/generator/multiview/caption_chunks.py @@ -10,6 +10,7 @@ from cosmos_framework.data.generator.multiview.camera_attributes import MADS_CAMERA_ATTRIBUTES from cosmos_framework.data.generator.multiview.caption_format import ( DEFAULT_CAPTION_PREFIXES, + first_caption_paragraph, format_multiview_caption, format_separate_view_captions, format_view_caption, @@ -209,6 +210,7 @@ def load_multiview_caption_chunks_per_view( views: Sequence[Any], *, first_only: bool = False, + first_caption_paragraph_only: bool = False, ) -> list[MultiviewCaptionChunk]: """Read every camera's caption file and assemble one prompt per chunk. @@ -217,6 +219,9 @@ def load_multiview_caption_chunks_per_view( ``separate_view_text_tokenization``. Which one is packed is decided at generation time from the checkpoint's own recorded layout, not here. + ``first_caption_paragraph_only`` matches short-window LiDAR training. It trims + the raw text before adding camera headers; the default preserves full captions. + Chunkwise rollout generates all cameras together and slices the model output by one shared frame count, so the cameras have to agree on the chunk boundaries: a caption file written against a different frame span would silently describe frames other than the generated ones. @@ -258,6 +263,8 @@ def load_multiview_caption_chunks_per_view( # current-camera headers as the training dataloader. Which one is packed is decided at # generation time from the checkpoint's own training config. view_captions = [chunk.prompt for chunk in view_chunks] + if first_caption_paragraph_only: + view_captions = [first_caption_paragraph(caption) for caption in view_captions] labeled = label_view_captions(view_captions, camera_keys=camera_keys) separate_view_captions = format_separate_view_captions( view_captions, diff --git a/cosmos_framework/data/generator/multiview/caption_format.py b/cosmos_framework/data/generator/multiview/caption_format.py index 6544c789..3d13d310 100644 --- a/cosmos_framework/data/generator/multiview/caption_format.py +++ b/cosmos_framework/data/generator/multiview/caption_format.py @@ -6,6 +6,7 @@ from __future__ import annotations import json +import re from collections.abc import Mapping, Sequence from typing import Any, Final @@ -31,6 +32,19 @@ } +_TIMESTAMP: re.Pattern[str] = re.compile(r"\[\s*\d+(?:\.\d+)?\s*s?\s*[-–—]\s*\d+(?:\.\d+)?\s*s?\s*\]", re.IGNORECASE) + + +def first_caption_paragraph(caption: str) -> str: + """Keep the opening timestamped paragraph, or the first blank-line paragraph.""" + caption = caption.strip() + timestamps = list(_TIMESTAMP.finditer(caption)) + if timestamps: + end = timestamps[1].start() if len(timestamps) > 1 else len(caption) + caption = caption[timestamps[0].start() : end].rstrip(" ,;\t\r\n") + return re.split(r"\r?\n\s*\r?\n", caption, maxsplit=1)[0].strip() + + def _camera_identity(camera_name: str, camera_attributes: Mapping[str, str | int]) -> str: """Render a camera's role and lens type as a natural-language identity.""" required_attributes = ("camera_role", "camera_type") diff --git a/cosmos_framework/data/generator/processors/qwen35_caption_video_test.py b/cosmos_framework/data/generator/processors/qwen35_caption_video_test.py index e4e8837c..058c2275 100644 --- a/cosmos_framework/data/generator/processors/qwen35_caption_video_test.py +++ b/cosmos_framework/data/generator/processors/qwen35_caption_video_test.py @@ -157,7 +157,7 @@ def test_qwen35_framewise_decoder_preserves_exact_source_timestamps(processor: Q video_temporal_mode="framewise", ) assert decoded is not None - frames = tensor_to_pil_images(decoded["videos"]) + frames = tensor_to_pil_images(decoded["videos"], channels_first=True) source_indices = decoded["source_frames_indices"] assert len(frames) == len(source_indices) diff --git a/cosmos_framework/data/generator/sequence_packing/packers.py b/cosmos_framework/data/generator/sequence_packing/packers.py index ab985abb..2f128287 100644 --- a/cosmos_framework/data/generator/sequence_packing/packers.py +++ b/cosmos_framework/data/generator/sequence_packing/packers.py @@ -169,6 +169,9 @@ def pack_input_sequence( action_dim: int = 32, initial_mrope_temporal_offset: int | float | list[int | float] = 0, lidar_temporal_compression_factor: int | None = None, + radar_temporal_compression_factor: int | None = None, + lidar_patch_spatial_hw: int | tuple[int, int] | None = None, + radar_patch_spatial_hw: int | tuple[int, int] | None = None, ) -> PackedSequence: """ Pack a sequence of input strings and VAE latents into a packed tensor format. @@ -181,6 +184,7 @@ def pack_input_sequence( gen_data_clean: GenerationDataClean containing vision, LiDAR, action, and sound tensors. - x0_tokens_vision: Vision tensors for samples where has_vision=True - x0_tokens_lidar: LiDAR tensors for samples where has_lidar=True + - x0_tokens_radar: Radar tensors for samples where has_radar=True - x0_tokens_action: Action tensors for samples where has_action=True - x0_tokens_sound: Sound tensors (list of [C, T]) for samples where has_sound=True input_timesteps: Diffusion timesteps for each sample. Shape (B,) or (B, 1) for @@ -189,7 +193,9 @@ def pack_input_sequence( sample as a float (numel==1) or Tensor(T_max,) for per-frame indexing. special_tokens: Dictionary containing special token IDs (eos_token_id, start_of_generation, end_of_generation) max_num_tokens: Maximum number of tokens in the packed sequence - latent_patch_size: Patch size used by the network to pack latents + latent_patch_size: Patch size used by the network to pack camera latents. + lidar_patch_spatial_hw: LiDAR patch side or (height, width); None inherits latent_patch_size. + radar_patch_spatial_hw: Radar patch side or (height, width); None inherits latent_patch_size. skip_text_tokens: If True, skip packing text tokens include_end_of_generation_token: If True, append end-of-generation token unified_3d_mrope_reset_spatial_ids: If True (default), spatial (H, W) indices @@ -281,6 +287,11 @@ def pack_input_sequence( raise ValueError("A sequence plan sets has_lidar, but gen_data_clean.x0_tokens_lidar is None.") if video_temporal_causal: raise NotImplementedError("Temporal-causal packing is not wired for the LiDAR stream yet.") + if any(plan.has_radar for plan in sequence_plans): + if gen_data_clean.x0_tokens_radar is None: + raise ValueError("A sequence plan sets has_radar, but gen_data_clean.x0_tokens_radar is None.") + if video_temporal_causal: + raise NotImplementedError("Temporal-causal packing is not wired for the radar stream yet.") use_float_mrope_positions = enable_fps_modulation or explicit_vision_temporal_positions_active @@ -300,6 +311,7 @@ def pack_input_sequence( idx_text = 0 idx_vision = 0 idx_lidar = 0 + idx_radar = 0 idx_action = 0 idx_sound = 0 null_action_flags: list[bool] = [] # collected from TC path; asserted consistent after the loop @@ -331,6 +343,7 @@ def pack_input_sequence( has_generation_for_sample = ( sequence_plan.has_vision or sequence_plan.has_lidar + or sequence_plan.has_radar or sequence_plan.has_action or sequence_plan.has_sound ) @@ -466,6 +479,7 @@ def pack_input_sequence( sample_len += vision_split_len action_split_len = 0 # Already absorbed into vision_split_len lidar_split_len = 0 # Temporal-causal packing rejects LiDAR above + radar_split_len = 0 # Temporal-causal packing rejects radar above else: # Standard path: vision and action packed separately @@ -702,7 +716,9 @@ def pack_input_sequence( input_lidar_tokens=input_lidar_tokens, condition_frame_indexes_lidar=item_condition_frames, input_timestep=input_timestep, - latent_patch_size=latent_patch_size, + latent_patch_size=( + latent_patch_size if lidar_patch_spatial_hw is None else lidar_patch_spatial_hw + ), lidar_fps=sample_lidar_fps, enable_fps_modulation=enable_fps_modulation, base_fps=base_fps, @@ -716,6 +732,60 @@ def pack_input_sequence( else: lidar_split_len = 0 + # Pack radar tokens if has_radar=True. They follow this sample's LiDAR items, so + # the packed stream reads [camera items | LiDAR items | radar items] and the + # multiview mask still sees one item list per sample. + if sequence_plan.has_radar: + num_radar = ( + gen_data_clean.num_radar_items_per_sample[sample_idx] + if gen_data_clean.num_radar_items_per_sample is not None + else 1 + ) + if radar_temporal_compression_factor is None: + raise ValueError("radar_temporal_compression_factor must be set when has_radar=True") + + sample_radar_fps = _get_optional_fps(gen_data_clean.fps_radar, sample_idx) + if sample_radar_fps is None: + raise ValueError("sample_radar_fps must be set when has_radar=True") + + # Radar was cut from the same window as the other sensors, so each radar item + # starts where the vision items started and only extends the sample's clock if + # its scans outlast them. + streams_end_offset = seq_builder.mrope_temporal_offset + + radar_split_len = 0 + for item_idx in range(num_radar): + input_radar_tokens = gen_data_clean.x0_tokens_radar[idx_radar] # [1,C,T,H,W] + idx_radar += 1 + + item_condition_frames = resolve_item_condition_frames( + sequence_plan.condition_frame_indexes_radar, + item_idx=item_idx, + num_items=num_radar, + latent_t=input_radar_tokens.shape[2], + ) + + seq_builder.set_mrope_temporal_offset(vision_start_temporal_offset) + radar_split_len += seq_builder.pack_radar_tokens( + input_radar_tokens=input_radar_tokens, + condition_frame_indexes_radar=item_condition_frames, + input_timestep=input_timestep, + latent_patch_size=( + latent_patch_size if radar_patch_spatial_hw is None else radar_patch_spatial_hw + ), + radar_fps=sample_radar_fps, + enable_fps_modulation=enable_fps_modulation, + base_fps=base_fps, + temporal_compression_factor=radar_temporal_compression_factor, + base_temporal_compression_factor=temporal_compression_factor, + ) + streams_end_offset = max(streams_end_offset, seq_builder.mrope_temporal_offset) + + seq_builder.set_mrope_temporal_offset(streams_end_offset) + sample_len += radar_split_len + else: + radar_split_len = 0 + # Pack action tokens if has_action=True if sequence_plan.has_action: input_action_tokens = gen_data_clean.x0_tokens_action[idx_action] @@ -760,7 +830,11 @@ def pack_input_sequence( # Add end-of-generation token if needed eov_len = 0 has_any_generation = ( - sequence_plan.has_vision or sequence_plan.has_lidar or sequence_plan.has_action or sequence_plan.has_sound + sequence_plan.has_vision + or sequence_plan.has_lidar + or sequence_plan.has_radar + or sequence_plan.has_action + or sequence_plan.has_sound ) if include_end_of_generation_token and has_any_generation: eov_len = seq_builder.append_end_of_generation_token( @@ -769,7 +843,9 @@ def pack_input_sequence( ) sample_len += eov_len - combined_split_len = vision_split_len + lidar_split_len + action_split_len + sound_split_len + eov_len + combined_split_len = ( + vision_split_len + lidar_split_len + radar_split_len + action_split_len + sound_split_len + eov_len + ) seq_builder.finish_sample(combined_split_len, sample_len) # Assert consistent null_action_supertokens across all TC samples, then set once diff --git a/cosmos_framework/data/generator/sequence_packing/sequence.py b/cosmos_framework/data/generator/sequence_packing/sequence.py index be297c14..aafad7bf 100644 --- a/cosmos_framework/data/generator/sequence_packing/sequence.py +++ b/cosmos_framework/data/generator/sequence_packing/sequence.py @@ -21,6 +21,7 @@ prepare_sequence_pack_metadata, to_device_nonblocking, ) +from cosmos_framework.utils.generator.spatial_patch import normalize_spatial_patch_hw if TYPE_CHECKING: from cosmos_framework.model.generator.utils.data_and_condition import GenerationDataClean @@ -121,6 +122,7 @@ class PackedSequenceBuilder: # Generation modality construction state vision: ModalityDataBuilder | None = None lidar: ModalityDataBuilder | None = None + radar: ModalityDataBuilder | None = None action: ModalityDataBuilder | None = None sound: ModalityDataBuilder | None = None @@ -144,6 +146,16 @@ def ensure_lidar(self) -> ModalityDataBuilder: self.lidar = ModalityDataBuilder() return self.lidar + def ensure_radar(self) -> ModalityDataBuilder: + """Return the radar builder, creating it on first use. + + Returns: + Radar ``ModalityDataBuilder`` for subsequent append operations. + """ + if self.radar is None: + self.radar = ModalityDataBuilder() + return self.radar + def ensure_action(self) -> ModalityDataBuilder: """Return the action builder, creating it on first use. @@ -412,7 +424,7 @@ def pack_lidar_tokens( input_lidar_tokens: torch.Tensor, condition_frame_indexes_lidar: list[int], input_timestep: float | torch.Tensor, - latent_patch_size: int, + latent_patch_size: int | tuple[int, int], lidar_fps: float, enable_fps_modulation: bool, base_fps: float, @@ -431,7 +443,7 @@ def pack_lidar_tokens( input_lidar_tokens: LiDAR latent tokens (C, T, H, W). condition_frame_indexes_lidar: Indexes of conditioning sweeps. input_timestep: Diffusion timestep, as in ``pack_vision_tokens``. - latent_patch_size: Patch size for latent patchification. + latent_patch_size: Patch side or (height, width). lidar_fps: Sweep rate of the LiDAR data. Used when enable_fps_modulation=True. enable_fps_modulation: If True, scale temporal position IDs based on FPS. base_fps: Base FPS for normalization. @@ -458,6 +470,57 @@ def pack_lidar_tokens( temporal_position_period=None, ) + def pack_radar_tokens( + self, + input_radar_tokens: torch.Tensor, + condition_frame_indexes_radar: list[int], + input_timestep: float | torch.Tensor, + latent_patch_size: int, + radar_fps: float, + enable_fps_modulation: bool, + base_fps: float, + temporal_compression_factor: int, + base_temporal_compression_factor: int, + ) -> int: + """Pack radar BEV tokens into the sequence. + + A radar BEV clip is a grid latent like a range clip, so it goes through the same + routine as LiDAR; only the clock differs. Radar cycles at ~20 Hz against LiDAR's + 10 Hz and the camera's 30 fps, and its VAE does not compress time, so ``radar_fps`` + with ``temporal_compression_factor=1`` is what places a scan on the mRoPE axis + shared with the other streams. + + Args: + input_radar_tokens: Radar latent tokens (C, T, H, W). + condition_frame_indexes_radar: Indexes of conditioning scans. + input_timestep: Diffusion timestep, as in ``pack_vision_tokens``. + latent_patch_size: Patch size for latent patchification. + radar_fps: Cycle rate of the radar data. Used when enable_fps_modulation=True. + enable_fps_modulation: If True, scale temporal position IDs based on FPS. + base_fps: Base FPS for normalization. + temporal_compression_factor: Temporal compression factor of the radar VAE, + which is 1. + base_temporal_compression_factor: Temporal compression factor defining the mRoPE time + unit shared with the vision stream. + + Returns: + Radar split length. + """ + return self._pack_grid_tokens( + self.ensure_radar(), + input_tokens=input_radar_tokens, + condition_frame_indexes=condition_frame_indexes_radar, + input_timestep=input_timestep, + latent_patch_size=latent_patch_size, + fps=radar_fps, + enable_fps_modulation=enable_fps_modulation, + base_fps=base_fps, + temporal_compression_factor=temporal_compression_factor, + base_temporal_compression_factor=base_temporal_compression_factor, + temporal_positions=None, + temporal_position_period=None, + ) + def _pack_grid_tokens( self, modality: ModalityDataBuilder, @@ -465,7 +528,7 @@ def _pack_grid_tokens( input_tokens: torch.Tensor, condition_frame_indexes: list[int], input_timestep: float | torch.Tensor, - latent_patch_size: int, + latent_patch_size: int | tuple[int, int], fps: float | None, enable_fps_modulation: bool, base_fps: float, @@ -485,11 +548,10 @@ def _pack_grid_tokens( """ # Compute position IDs for image patches _, _, latent_t, latent_h, latent_w = input_tokens.shape - if latent_patch_size < 1: - raise ValueError(f"latent_patch_size must be >= 1, got {latent_patch_size}") + patch_height, patch_width = normalize_spatial_patch_hw(latent_patch_size) # Use ceil to support latent dims not divisible by patch size (padding handled in network) - patch_h = math.ceil(latent_h / latent_patch_size) - patch_w = math.ceil(latent_w / latent_patch_size) + patch_h = math.ceil(latent_h / patch_height) + patch_w = math.ceil(latent_w / patch_width) modality.token_shapes.append((latent_t, patch_h, patch_w)) modality.tokens.append(input_tokens) payload_index = len(modality.tokens) - 1 @@ -1027,6 +1089,7 @@ def finalize( vision = self._finalize_modality(self.vision) lidar = self._finalize_modality(self.lidar) + radar = self._finalize_modality(self.radar) action_domain_id = None if self.action is not None: if gen_data_clean.action_domain_id is not None: @@ -1108,6 +1171,7 @@ def finalize( # Generation modalities vision=vision, lidar=lidar, + radar=radar, action=action, sound=sound, # Temporal causal @@ -1121,8 +1185,11 @@ def finalize( # Vision item layout (multi-item samples, multiview cameras) num_vision_items_per_sample=gen_data_clean.num_vision_items_per_sample, num_views_per_vision_item=gen_data_clean.num_views_per_vision_item, + vision_view_ids=gen_data_clean.vision_view_ids, # LiDAR item layout num_lidar_items_per_sample=gen_data_clean.num_lidar_items_per_sample, + # Radar item layout + num_radar_items_per_sample=gen_data_clean.num_radar_items_per_sample, ) @@ -1156,6 +1223,7 @@ class PackedSequence: temporal-causal vision supertoken. vision: Finalized vision modality data, or ``None`` if no vision is present. lidar: Finalized LiDAR modality data, or ``None`` if no LiDAR is present. + radar: Finalized radar modality data, or ``None`` if no radar is present. action: Finalized action modality data, or ``None`` if no action is present. sound: Finalized sound modality data, or ``None`` if no sound is present. vision_item_split_lens: Per-sample per-vision-item token counts for multi-control @@ -1167,6 +1235,8 @@ class PackedSequence: vision item, or ``None`` when per-camera VAE encoding is disabled. num_lidar_items_per_sample: Number of LiDAR items owned by each sample, or ``None`` when the batch carries no LiDAR. + num_radar_items_per_sample: Number of radar items owned by each sample, or ``None`` + when the batch carries no radar. """ # Sequence structure @@ -1202,6 +1272,7 @@ class PackedSequence: # Generation modalities - NAMED FIELDS for type safety vision: ModalityData | None = None lidar: ModalityData | None = None + radar: ModalityData | None = None action: ModalityData | None = None sound: ModalityData | None = None @@ -1233,6 +1304,10 @@ class PackedSequence: # build the multiview FlexAttention mask. num_vision_items_per_sample: list[int] | None = None num_views_per_vision_item: list[int] | None = None + # Physical camera IDs, one [V] tensor per flattened vision item in camera-major order, + # copied from GenerationDataClean. Controls and targets retain the same IDs without + # renumbering selected cameras. None when disabled or no RGB is present; excludes LiDAR. + vision_view_ids: list[torch.Tensor] | None = None # LiDAR items owned by each sample, grouping the flattened LiDAR items the way # num_vision_items_per_sample groups the vision ones. None when the batch has no LiDAR. @@ -1240,6 +1315,10 @@ class PackedSequence: # FlexAttention mask as its vision items followed by its LiDAR items. num_lidar_items_per_sample: list[int] | None = None + # Same bookkeeping for the radar stream, whose items follow the sample's LiDAR items. + # None when the batch carries no radar. + num_radar_items_per_sample: list[int] | None = None + def __post_init__(self) -> None: self._sequence_pack_metadata: SequencePackMetadata | None = None assert isinstance(self.text_ids, torch.Tensor), "PackedSequence.text_ids must be finalized" @@ -1251,7 +1330,7 @@ def __post_init__(self) -> None: assert isinstance(self.ce_loss_indexes, torch.Tensor), "PackedSequence.ce_loss_indexes must be finalized" if self.ce_loss_weights is not None: assert isinstance(self.ce_loss_weights, torch.Tensor), "PackedSequence.ce_loss_weights must be finalized" - for modality in [self.vision, self.lidar, self.action, self.sound]: + for modality in [self.vision, self.lidar, self.radar, self.action, self.sound]: assert modality is None or isinstance(modality, ModalityData), ( "PackedSequence modality fields must be finalized ModalityData" ) @@ -1275,10 +1354,14 @@ def to_cuda(self) -> None: self.ce_loss_indexes = to_device_nonblocking(self.ce_loss_indexes, "cuda") if isinstance(self.ce_loss_weights, torch.Tensor): self.ce_loss_weights = to_device_nonblocking(self.ce_loss_weights, "cuda") + if self.vision_view_ids is not None: + self.vision_view_ids = [to_device_nonblocking(ids, "cuda") for ids in self.vision_view_ids] # list[[V]] if self.vision is not None: self.vision.to_cuda() if self.lidar is not None: self.lidar.to_cuda() + if self.radar is not None: + self.radar.to_cuda() if self.action is not None: self.action.to_cuda() if self.sound is not None: @@ -1355,6 +1438,10 @@ class SequencePlan: condition_frame_indexes_lidar: Indexes of latent LiDAR sweeps that are clean/conditioning, read the same way as ``condition_frame_indexes_vision`` and applying to each LiDAR item individually. + has_radar: Whether radar BEV latents are present for this sample. + condition_frame_indexes_radar: Indexes of latent radar scans that are clean/conditioning, + read the same way as ``condition_frame_indexes_lidar`` and applying to each radar + item individually. has_action: Whether action input is present for robotics/embodied AI tasks. Defaults to False. condition_frame_indexes_action: Indexes of action steps that are clean/conditioning. @@ -1391,6 +1478,13 @@ class SequencePlan: has_lidar: bool = False condition_frame_indexes_lidar: list[int] = field(default_factory=list) + # -- radar modality -- + # Radar items start at the same instant as the sample's vision items, exactly as LiDAR + # does. Radar cycles at ~20 Hz rather than 10 Hz, but that rate lives in ``radar_fps`` + # on the packing call, so nothing about the plan differs from LiDAR's. + has_radar: bool = False + condition_frame_indexes_radar: list[int] = field(default_factory=list) + # -- action modality -- has_action: bool = False condition_frame_indexes_action: list[int] = field(default_factory=list) @@ -1406,11 +1500,13 @@ def as_dict(self) -> dict: "text_view_ids": self.text_view_ids, "has_vision": self.has_vision, "has_lidar": self.has_lidar, + "has_radar": self.has_radar, "has_action": self.has_action, "has_sound": self.has_sound, "condition_frame_indexes_vision": self.condition_frame_indexes_vision, "condition_view_indexes_vision": self.condition_view_indexes_vision, "condition_frame_indexes_lidar": self.condition_frame_indexes_lidar, + "condition_frame_indexes_radar": self.condition_frame_indexes_radar, "condition_frame_indexes_action": self.condition_frame_indexes_action, "condition_frame_indexes_sound": self.condition_frame_indexes_sound, "share_vision_temporal_positions": self.share_vision_temporal_positions, @@ -1448,6 +1544,7 @@ def build_sequence_plans_from_data_batch( assert "action" not in data_batch or data_batch["action"] is None, "Action data SHOULD have sequence_plans!" assert "sound" not in data_batch or data_batch["sound"] is None, "Sound data SHOULD have sequence_plans!" assert "lidar" not in data_batch or data_batch["lidar"] is None, "LiDAR data SHOULD have sequence_plans!" + assert "radar" not in data_batch or data_batch["radar"] is None, "Radar data SHOULD have sequence_plans!" # Determine batch size from available tensors batch_size = 0 diff --git a/cosmos_framework/model/generator/joint_transfer_ar.py b/cosmos_framework/model/generator/joint_transfer_ar.py index 97741af1..830c5747 100644 --- a/cosmos_framework/model/generator/joint_transfer_ar.py +++ b/cosmos_framework/model/generator/joint_transfer_ar.py @@ -9,6 +9,7 @@ import torch +from cosmos_framework.model.generator.mot.context_parallel_utils import context_parallel_broadcast_tensor_list from cosmos_framework.model.generator.utils.data_and_condition import GenerationDataClean from cosmos_framework.data.generator.sequence_packing import PackedSequence, SequencePlan from cosmos_framework.model.generator.utils.kv_cache import TeacherForcingMemoryState @@ -165,7 +166,7 @@ def sample_joint_transfer_ar( raise ValueError("Joint AR requires exactly one RGB+LiDAR sequence") if plans[0].has_action or plans[0].has_sound: raise ValueError("Joint AR supports only RGB and LiDAR targets") - if not host._uses_multiview_flex_kv() or host.config.compile.enabled: + if not host._uses_multiview_replay_kv() or host.config.compile.enabled: raise ValueError("Joint AR requires eager multiview Flex teacher forcing") if host.parallel_dims is not None and host.parallel_dims.cfgp_enabled: raise ValueError("Joint AR currently uses serial CFG branches; set cfg_parallel_shard_degree=1") @@ -192,6 +193,14 @@ def sample_joint_transfer_ar( if seed < 0 or num_steps < 1 or not math.isfinite(guidance) or not math.isfinite(shift) or shift <= 0: raise ValueError("Joint AR requires a nonnegative seed and finite valid sampling settings") + # Rank-local VAE encoding can differ even with identical pixels and seeds. + # CP shards must read one canonical set of controls and conditioned targets + # before the initial pack and every subsequent clean/noisy replay. + # Keep the existing Flex inference input behavior unchanged. + if host._get_teacher_forcing_kv_implementation() == "multiview_maskless_kv": + context_parallel_broadcast_tensor_list(data.x0_tokens_vision, host.parallel_dims) # each [1,Cv,V*Tv,Hv,Wv] + context_parallel_broadcast_tensor_list(data.x0_tokens_lidar, host.parallel_dims) # each [1,Cl,Tl,Hl,Wl] + initial_pack = host._pack_input_sequence(plans, conditional_text, data, torch.zeros(1)) # timestep: [1] if initial_pack.vision is None or initial_pack.lidar is None: raise ValueError("The joint inference pack omitted a required modality") diff --git a/cosmos_framework/model/generator/mot/causal_attention.py b/cosmos_framework/model/generator/mot/causal_attention.py index b8d06a73..4801dbbc 100644 --- a/cosmos_framework/model/generator/mot/causal_attention.py +++ b/cosmos_framework/model/generator/mot/causal_attention.py @@ -37,6 +37,12 @@ get_gen_seq, ) from cosmos_framework.configs.base.defaults.replay_attention import TeacherForcingReplayPolicyConfig +from cosmos_framework.model.generator.mot.maskless_attention import ( + ReplayMasklessPlan, + cat_replay_kv, + replay_maskless_attention, +) +from cosmos_framework.model.generator.mot.merge_attention import merge_attentions_ac_safe from cosmos_framework.model.generator.utils.kv_cache import ( ARMemoryValue, FlexARMemoryValue, @@ -410,12 +416,13 @@ def two_way_flex_attention_with_memory( packed_value_states: SequencePack, *, packed_key_states_normalized: SequencePack | None, - flex_block_mask: BlockMask, - flex_backend: FlexBackend, + flex_block_mask: BlockMask | None, + flex_backend: FlexBackend | None, flex_memory_k: torch.Tensor | None, flex_memory_v: torch.Tensor | None, + maskless_plan: ReplayMasklessPlan | None = None, ) -> SequencePack: - """Run two-way attention with an optional key-only Flex K/V suffix.""" + """Run two-way replay attention with an optional key-only K/V suffix.""" if (flex_memory_k is None) != (flex_memory_v is None): raise ValueError("flex_memory_k and flex_memory_v must be provided together.") packed_key_normalized = ( @@ -435,6 +442,12 @@ def two_way_flex_attention_with_memory( causal_q_offsets, max_causal_len = caption_offsets # [N_captions+1], int causal_k_offsets = causal_q_offsets # [N_captions+1] + if maskless_plan is not None and torch.compiler.is_compiling(): + # NATTEN validates these bounds with Python branches. State the caption + # packing contract explicitly when CP stream lengths are unbacked. + torch._check(max_causal_len <= causal_q.shape[0]) + torch._check(max_causal_len <= causal_k.shape[0]) + use_dont_care_mask = causal_q_offsets is causal_k_offsets causal_res = attention( causal_q.unsqueeze(0), # [1,N_und,H,D] @@ -459,17 +472,30 @@ def two_way_flex_attention_with_memory( value_parts = [und_v, gen_v] if flex_memory_k is not None: assert flex_memory_v is not None + if maskless_plan is not None and torch.compiler.is_compiling(): + torch._check(flex_memory_k.shape[1] == flex_memory_v.shape[1]) key_parts.append(flex_memory_k.squeeze(0)) # [N_memory,H,D] value_parts.append(flex_memory_v.squeeze(0)) # [N_memory,H,D] - flex_keys = torch.cat(key_parts).unsqueeze(0) # [1,N_und+N_gen+N_memory,H,D] - flex_values = torch.cat(value_parts).unsqueeze(0) # [1,N_und+N_gen+N_memory,H,D] - full_res = flex_attention( - full_q.unsqueeze(0), # [1,N_gen,H,D] - flex_keys, - flex_values, - flex_block_mask, - flex_backend, - ) # [1,N_gen,H,D] + concatenate = cat_replay_kv if maskless_plan is not None and flex_memory_k is not None else torch.cat + flex_keys = concatenate(key_parts).unsqueeze(0) # [1,N_und+N_gen+N_memory,H,D] + flex_values = concatenate(value_parts).unsqueeze(0) # [1,N_und+N_gen+N_memory,H,D] + if maskless_plan is not None: + full_res = replay_maskless_attention( + full_q.unsqueeze(0), + flex_keys, + flex_values, + maskless_plan, # [1,N_gen,H,D], [1,KV,H,D], [1,KV,H,D] + ) # [1,N_gen,H,D] + else: + if flex_block_mask is None or flex_backend is None: + raise ValueError("Two-way Flex replay requires its mask and backend.") + full_res = flex_attention( + full_q.unsqueeze(0), # [1,N_gen,H,D] + flex_keys, + flex_values, + flex_block_mask, + flex_backend, + ) # [1,N_gen,H,D] full_out = full_res.squeeze(0).flatten(-2, -1) # [N_gen,H*D] return from_mode_splits(causal_out, full_out, packed_query_states) @@ -1661,102 +1687,6 @@ def three_way_attention_with_kv_cache( return from_mode_splits(text_out, video_out, packed_query_states) -class _ACSafeMergeAttentionsFn(torch.autograd.Function): - """AC-compatible drop-in for NATTEN's ``MergeAttentionsAutogradFn``. - - NATTEN's backward indexes ``ctx.saved_tensors`` in multiple slices - (``[:2]``, ``[2 : N+2]``, ``[N+2:]``), and each indexing access fires - the non-reentrant ``torch.utils.checkpoint`` unpack hook for *every* - saved tensor. The hook only permits one unpack per saved tensor, so - activation checkpointing + NATTEN merge_attentions raises - ``CheckpointError: Unpack is being triggered for a tensor that was - already unpacked once`` (see the replayed-LSE + AC=full long-video - TF path). - - This version preserves the same forward math and the same - storage-patching backward contract (see :class:`MergeAttentionsBridge` - docstring for the full description), but reads ``ctx.saved_tensors`` - exactly once. Numerics match ``naive_merge_attentions`` (iterative - pairwise LSE rescale) — which is what NATTEN's kernel implements up - to reduction order. - """ - - @staticmethod - def forward( - ctx, - num_components: int, - *tensors: torch.Tensor, - ) -> tuple[torch.Tensor, torch.Tensor]: - outputs = tensors[:num_components] - lses = tensors[num_components:] - output_dtype = outputs[0].dtype - normalized_lses = [lse.squeeze(-1) if lse.ndim == 4 else lse for lse in lses] - - merged_lse = normalized_lses[0] - merged_out = outputs[0] - for i in range(1, num_components): - new_lse = torch.logaddexp(merged_lse, normalized_lses[i]) - w_old = torch.exp(merged_lse - new_lse).unsqueeze(-1) - w_new = torch.exp(normalized_lses[i] - new_lse).unsqueeze(-1) - merged_out = w_old * merged_out + w_new * outputs[i] - merged_lse = new_lse - merged_out = merged_out.to(output_dtype) - - ctx.save_for_backward(merged_out, merged_lse, *outputs, *lses) - ctx.num_components = num_components - return merged_out, merged_lse - - @staticmethod - def backward( - ctx, - grad_merged_out: torch.Tensor, - grad_merged_lse: torch.Tensor, - ) -> tuple[torch.Tensor | None, ...]: - # Single access — avoid retriggering the AC unpack hook for any saved tensor. - saved = ctx.saved_tensors - merged_out = saved[0] - merged_lse = saved[1] - num = ctx.num_components - outputs = saved[2 : 2 + num] - lses = saved[2 + num : 2 + 2 * num] - - # Patch each component's storage with the merged O / LSE. The - # upstream attention kernel's backward will read these as its - # saved O / LSE and compute gradients as if it had produced the - # merged output. The original LSE shape is preserved (the - # forward squeezes a trailing singleton, so we re-broadcast). - for o in outputs: - o.data.copy_(merged_out.data) - for l in lses: - if l.ndim == merged_lse.ndim + 1 and l.shape[-1] == 1: - l.data.copy_(merged_lse.data.unsqueeze(-1)) - else: - l.data.copy_(merged_lse.data) - - # Same upstream-grad contract as NATTEN: dL/dO_i = dL/dO_merged - # for every component; dL/dLSE_i is forwarded unchanged for - # parity (i4 attention treats LSE as non-differentiable, so this - # gradient is silently dropped at the kernel boundary). - grads = (None,) + (grad_merged_out,) * num + (grad_merged_lse,) * num - return grads - - -def merge_attentions_ac_safe( - outputs: list[torch.Tensor], - lse_tensors: list[torch.Tensor], -) -> tuple[torch.Tensor, torch.Tensor]: - """AC-safe drop-in for ``cosmos_framework.model.attention.merge_attentions``. - - Use at call sites that live inside an activation-checkpointed module - boundary. Matches NATTEN's storage-patching backward contract so - upstream i4 attention kernels (whose LSE is not differentiable) - still receive correct gradients via their own saved O / LSE - backward formulas. - """ - assert len(outputs) == len(lse_tensors) >= 2 - return _ACSafeMergeAttentionsFn.apply(len(outputs), *outputs, *lse_tensors) - - def naive_merge_attentions( outputs: list[torch.Tensor], lse_tensors: list[torch.Tensor], @@ -2159,14 +2089,24 @@ def dispatch_attention_with_memory( - ``ARMemoryValue`` with ``frame_idx == 0`` → interactive no-memory dispatch - ``None`` → interactive no-memory dispatch """ + maskless_plan = getattr(attention_mask, "replay_maskless_plan", None) + if maskless_plan is not None and not isinstance(maskless_plan, ReplayMasklessPlan): + raise TypeError("Replay attention requires a ReplayMasklessPlan.") if ( - isinstance(memory_value, (TFReplayCleanMemoryValue, TFNoisyMemoryValue, FlexARMemoryValue)) + ( + isinstance(memory_value, (TFReplayCleanMemoryValue, TFNoisyMemoryValue, FlexARMemoryValue)) + or maskless_plan is not None + ) and isinstance(attention_mask, SplitInfo) and not attention_mask.is_three_way - and attention_mask.flex_block_mask is not None + and (attention_mask.flex_block_mask is not None or maskless_plan is not None) ): - if attention_mask.flex_backend is None: - raise ValueError("Two-way Flex memory attention requires a FlexBackend.") + if ( + maskless_plan is not None + and memory_value is not None + and not isinstance(memory_value, (TFReplayCleanMemoryValue, TFNoisyMemoryValue, FlexARMemoryValue)) + ): + raise TypeError("Maskless multiview replay received an incompatible memory layout.") output = two_way_flex_attention_with_memory( packed_query_states, packed_key_states, @@ -2174,6 +2114,7 @@ def dispatch_attention_with_memory( packed_key_states_normalized=packed_key_states_normalized, flex_block_mask=attention_mask.flex_block_mask, flex_backend=attention_mask.flex_backend, + maskless_plan=maskless_plan, flex_memory_k=( memory_value.cached_clean_gen_k if isinstance(memory_value, TFNoisyMemoryValue) diff --git a/cosmos_framework/model/generator/mot/causal_cosmos3_vfm_network.py b/cosmos_framework/model/generator/mot/causal_cosmos3_vfm_network.py index cc646fca..141cd041 100644 --- a/cosmos_framework/model/generator/mot/causal_cosmos3_vfm_network.py +++ b/cosmos_framework/model/generator/mot/causal_cosmos3_vfm_network.py @@ -1,7 +1,7 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: OpenMDW-1.1 -"""Interactive network adapter for teacher-forcing multiview FlexAttention.""" +"""Interactive network adapter for multiview replay attention.""" from __future__ import annotations @@ -9,6 +9,7 @@ from dataclasses import replace from typing import Any +import attrs import torch from torch.utils.hooks import RemovableHandle @@ -34,6 +35,7 @@ build_teacher_forcing_block_mask, build_teacher_forcing_multiview_flex_metadata, ) +from cosmos_framework.model.generator.mot.maskless_attention import build_replay_maskless_plan def build_interactive_multiview_mask_items( @@ -86,12 +88,14 @@ class InteractiveCosmos3VFMNetwork(Cosmos3VFMNetwork): _teacher_forcing_mask_hook: RemovableHandle teacher_forcing_replay_policy: TeacherForcingReplayPolicyConfig | None teacher_forcing_frames_per_chunk: int + teacher_forcing_maskless: bool = False def __init__(self, language_model: torch.nn.Module, config: Any) -> None: super().__init__(language_model=language_model, config=config) self._active_packed_seq = None self.teacher_forcing_replay_policy = None self.teacher_forcing_frames_per_chunk = 1 + self.teacher_forcing_maskless = False self._teacher_forcing_mask_hook = self.language_model.register_forward_pre_hook( self._replace_teacher_forcing_mask, with_kwargs=True, @@ -115,8 +119,16 @@ def _replace_teacher_forcing_mask( teacher_forcing_replay_policy = self.teacher_forcing_replay_policy if not isinstance(teacher_forcing_replay_policy, TeacherForcingReplayPolicyConfig): raise TypeError("Interactive teacher forcing requires a TeacherForcingReplayPolicyConfig.") - if self.flex_backend is None: + if self.flex_backend is None and not self.teacher_forcing_maskless: raise ValueError("Interactive teacher forcing requires a resolved FlexAttention backend.") + # Metadata is independent of spatial scope. Use the same-view metadata + # builder for maskless joint inputs; the Flex-only cross-sensor frame-index + # restriction does not apply to the maskless midpoint buckets. + metadata_policy = ( + attrs.evolve(teacher_forcing_replay_policy, multiview_attention_scope="same_view") + if self.teacher_forcing_maskless + else teacher_forcing_replay_policy + ) attention_meta = kwargs.get("attention_mask") if not isinstance(attention_meta, SplitInfo) or attention_meta.is_three_way: raise ValueError("Interactive multiview teacher forcing requires two-way SplitInfo metadata.") @@ -152,13 +164,14 @@ def _replace_teacher_forcing_mask( f"Multiview transfer AR current_role must be one of {MULTIVIEW_TRANSFER_AR_CURRENT_ROLES}, " f"got {current_role!r}." ) + sensor_items = build_interactive_multiview_mask_items( + packed_seq, + lidar_attends_captions=self.config.multiview_attention_config.mask.lidar_attends_captions, + ) flex_metadata = build_multiview_transfer_ar_flex_metadata( seq_len=global_gen_seq_len, full_q_offsets=full_q_offsets, - sensor_mask_items=build_interactive_multiview_mask_items( - packed_seq, - lidar_attends_captions=self.config.multiview_attention_config.mask.lidar_attends_captions, - ), + sensor_mask_items=sensor_items, caption_mask_items=_multiview_caption_mask_items(packed_seq), device=full_only_seq.device, num_und=global_und_seq_len, @@ -167,7 +180,7 @@ def _replace_teacher_forcing_mask( frames_per_view=frames_per_view, frames_per_chunk=frames_per_chunk, current_role=current_role, - teacher_forcing_replay_policy=teacher_forcing_replay_policy, + teacher_forcing_replay_policy=metadata_policy, memory_layout=memory_layout, ) elif teacher_forcing_pass is not None: @@ -189,21 +202,22 @@ def _replace_teacher_forcing_mask( original_masks = getattr(packed_seq, "teacher_forcing_original_condition_masks_vision", None) if original_masks is None: raise ValueError("Flex teacher forcing requires the original sensor condition masks.") + sensor_items = build_interactive_multiview_mask_items( + packed_seq, + lidar_attends_captions=self.config.multiview_attention_config.mask.lidar_attends_captions, + condition_masks=original_masks, + ) flex_metadata = build_teacher_forcing_multiview_flex_metadata( seq_len=global_gen_seq_len, full_q_offsets=full_q_offsets, - sensor_mask_items=build_interactive_multiview_mask_items( - packed_seq, - lidar_attends_captions=self.config.multiview_attention_config.mask.lidar_attends_captions, - condition_masks=original_masks, - ), + sensor_mask_items=sensor_items, caption_mask_items=_multiview_caption_mask_items(packed_seq), device=full_only_seq.device, num_und=global_und_seq_len, causal_offsets=causal_offsets, frames_per_chunk=self.teacher_forcing_frames_per_chunk, pass_kind=teacher_forcing_pass, - teacher_forcing_replay_policy=teacher_forcing_replay_policy, + teacher_forcing_replay_policy=metadata_policy, materialized_target_frame_ranges=materialized_target_frame_ranges, clean_memory_seq_len=( int(getattr(packed_seq, "teacher_forcing_selected_clean_target_padded_capacity", 0)) @@ -214,6 +228,14 @@ def _replace_teacher_forcing_mask( else: return None + if self.teacher_forcing_maskless: + flex_metadata = replace(flex_metadata, teacher_forcing_replay_policy=teacher_forcing_replay_policy) + setattr(attention_meta, "replay_maskless_plan", build_replay_maskless_plan(flex_metadata, sensor_items)) + # The marked replay dispatch must consume this plan, never the base + # bidirectional plan that was constructed before the pre-hook. + attention_meta.multiview_maskless = None + return args, kwargs + assert self.flex_backend is not None attention_meta.flex_block_mask = build_teacher_forcing_block_mask( flex_metadata, full_only_seq.device, @@ -233,6 +255,15 @@ def forward( previous_packed_seq = self._active_packed_seq self._active_packed_seq = packed_seq try: + if ( + self.teacher_forcing_maskless + and video_temporal_causal is not False + and ( + getattr(packed_seq, "teacher_forcing_pass", None) is not None + or getattr(packed_seq, "multiview_transfer_ar_metadata", None) is not None + ) + ): + kwargs["bounded_cp_output_gather"] = True return super().forward( packed_seq=packed_seq, memory=memory, diff --git a/cosmos_framework/model/generator/mot/causal_flex_attention.py b/cosmos_framework/model/generator/mot/causal_flex_attention.py index c7a07a1b..4d404ccd 100644 --- a/cosmos_framework/model/generator/mot/causal_flex_attention.py +++ b/cosmos_framework/model/generator/mot/causal_flex_attention.py @@ -920,8 +920,11 @@ def pair_allowed( clean_step_allowed = (q_is_current & (kv_step < q_step)) | ((~q_is_current) & clean_pass_causal) # [Q,KV] target_to_clean = q_is_target & (kv_role == _ROLE_CLEAN_TARGET) & clean_step_allowed & in_scope # [Q,KV] target_to_control = q_is_target & (kv_role == _ROLE_CONTROL) & same_view & control_step_allowed # [Q,KV] + # Observed sensor targets follow the configured target view/time scope. + # Initial LiDAR can read simultaneous RGB conditions; control ownership + # and the replay-chunk guard remain unchanged. target_to_condition = ( - q_is_target & (kv_role == _ROLE_TARGET_CONDITION) & same_view & (kv_step <= q_step) + q_is_target & (kv_role == _ROLE_TARGET_CONDITION) & in_scope & (kv_step <= q_step) ) # [Q,KV] padding_to_padding = (q_role == _ROLE_PADDING) & (kv_role == _ROLE_PADDING) # [Q,KV] allowed = ( diff --git a/cosmos_framework/model/generator/mot/context_parallel_gradient_test.py b/cosmos_framework/model/generator/mot/context_parallel_gradient_test.py index 04d1d033..d1f208fd 100644 --- a/cosmos_framework/model/generator/mot/context_parallel_gradient_test.py +++ b/cosmos_framework/model/generator/mot/context_parallel_gradient_test.py @@ -4,6 +4,7 @@ """Check CP output gathering against an unsharded, globally averaged objective.""" import copy +import os from datetime import timedelta from pathlib import Path @@ -14,6 +15,7 @@ from torch import nn from torch.nn.parallel import DistributedDataParallel +from cosmos_framework.model.generator.mot import context_parallel_utils from cosmos_framework.model.generator.mot.context_parallel_utils import ( get_context_parallel_last_hidden_state, get_context_parallel_sharded_sequence, @@ -124,3 +126,137 @@ def test_context_parallel_output_parameter_gradients( nprocs=4, join=True, ) + + +def _check_interleaved_output( + dims: ParallelDims, + split_lens: list[int], + *, + correct_cp_gradients: bool, + device: torch.device, + dtype: torch.dtype = torch.float32, + check_memory: bool = False, + bounded_memory: bool = True, +) -> None: + """Check a two-sample permutation, independent rank cotangents, and optional allocation bounds.""" + width = 64 + length = sum(split_lens) + original = ( + torch.arange(length * width, device=device, dtype=torch.float32).reshape(length, width).to(dtype) + ) # [N,D] + positions = torch.arange(length, device=device) # [N] + und_indices = torch.cat( + (positions[: split_lens[0]], positions[sum(split_lens[:2]) : sum(split_lens[:3])]) + ) # [N_und] + gen_indices = torch.cat( + (positions[split_lens[0] : sum(split_lens[:2])], positions[sum(split_lens[:3]) :]) + ) # [N_gen] + packed = sequence_pack_from_packed_sequence( + packed_sequence=original, + attn_modes=["causal", "full", "causal", "full"], + split_lens=split_lens, + sample_lens=[sum(split_lens[:2]), sum(split_lens[2:])], + packed_und_token_indexes=und_indices, + packed_gen_token_indexes=gen_indices, + full_seq_alignment=dims.cp_size, + causal_seq_alignment=dims.cp_size, + ) + local_pack, _ = get_context_parallel_sharded_sequence(packed, positions, dims) + und = local_pack["causal_seq"].detach().clone().requires_grad_(True) # [N_und_local,D] + gen = local_pack["full_only_seq"].detach().clone().requires_grad_(True) # [N_gen_local,D] + outputs = from_mode_splits(und, gen, local_pack) + # A noncontiguous, rank-dependent cotangent catches wrong reduction, rank and token order. + cotangent = torch.arange(length * width, device=device, dtype=torch.float32).remainder_(17) # [N*D] + cotangent = cotangent.reshape(width, length).T.to(dtype) # [N,D] + cotangent.mul_(dims.cp_rank + 1) # [N,D] + if check_memory: + torch.cuda.synchronize() + baseline = torch.cuda.memory_allocated() + torch.cuda.reset_peak_memory_stats() + actual = get_context_parallel_last_hidden_state( + outputs, dims, correct_cp_gradients=correct_cp_gradients, bounded_memory=bounded_memory + ) # [N,D] + if check_memory: + torch.cuda.synchronize() + # One full result plus bounded communication scratch; a second full result violates this. + forward_extra = torch.cuda.max_memory_allocated() - baseline + assert ( + forward_extra <= actual.numel() * actual.element_size() + 3 * context_parallel_utils._CP_OUTPUT_GATHER_BYTES + ) + torch.testing.assert_close(actual, original, atol=0, rtol=0) + if check_memory: + torch.cuda.synchronize() + baseline = torch.cuda.memory_allocated() + torch.cuda.reset_peak_memory_stats() + gradients = torch.autograd.grad(actual, (und, gen), cotangent) # tuple[[N_und_local,D],[N_gen_local,D]] + if check_memory: + torch.cuda.synchronize() + backward_extra = torch.cuda.max_memory_allocated() - baseline + gradient_bytes = sum(value.numel() * value.element_size() for value in gradients) + assert backward_extra <= gradient_bytes + 3 * context_parallel_utils._CP_OUTPUT_GATHER_BYTES + factor = dims.cp_size * (dims.cp_size + 1) // 2 if correct_cp_gradients else dims.cp_rank + 1 + expected_full = cotangent / (dims.cp_rank + 1) * factor # [N,D] + for gradient, indices in zip(gradients, (und_indices, gen_indices), strict=True): + local_rows = gradient.shape[0] + offset = dims.cp_rank * local_rows + count = max(0, min(local_rows, indices.numel() - offset)) + expected = torch.zeros_like(gradient) # [N_local,D] + expected[:count] = expected_full[indices[offset : offset + count]] # [count,D] + torch.testing.assert_close(gradient, expected, atol=0, rtol=0) + + +def _compare_interleaved_gradients(rank: int, rendezvous: str) -> None: + torch.set_num_threads(1) + dist.init_process_group("gloo", init_method=rendezvous, rank=rank, world_size=4) + try: + dims = ParallelDims(world_size=4, dp_shard=4, dp_replicate=1, cp=2) + dims.build_meshes("cpu") + context_parallel_utils._CP_OUTPUT_GATHER_BYTES = 1024 + for bounded_memory in (False, True): + for correction in (False, True): + for lengths in ([3, 5, 2, 7], [0, 5, 0, 7], [3, 0, 2, 0]): + _check_interleaved_output( + dims, + lengths, + correct_cp_gradients=correction, + device=torch.device("cpu"), + bounded_memory=bounded_memory, + ) + finally: + dist.destroy_process_group() + + +@pytest.mark.L0 +@pytest.mark.CPU +def test_context_parallel_interleaved_output_gradients(tmp_path: Path) -> None: + mp.spawn(_compare_interleaved_gradients, args=((tmp_path / "rendezvous").as_uri(),), nprocs=4, join=True) + + +@pytest.mark.L1 +@pytest.mark.GPU +@pytest.mark.parametrize("correct_cp_gradients", [False, True]) +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16]) +def test_context_parallel_output_memory( + correct_cp_gradients: bool, dtype: torch.dtype, monkeypatch: pytest.MonkeyPatch +) -> None: + """Require bounded native CP output/reduction memory on a four-rank launch.""" + if not torch.cuda.is_available() or int(os.environ.get("WORLD_SIZE", "1")) != 4: + pytest.skip("Requires torchrun with four GPUs.") + torch.cuda.set_device(int(os.environ["LOCAL_RANK"])) + if not dist.is_initialized(): + dist.init_process_group("nccl") + dims = ParallelDims(world_size=4, dp_shard=4, dp_replicate=1, cp=4) + dims.build_meshes("cuda") + monkeypatch.setattr(context_parallel_utils, "_CP_OUTPUT_GATHER_BYTES", 128 * 1024) + # Warm collectives before measuring tensor allocations. Uneven streams retain padding coverage. + _check_interleaved_output( + dims, [3, 5, 2, 7], correct_cp_gradients=correct_cp_gradients, device=torch.device("cuda"), dtype=dtype + ) + _check_interleaved_output( + dims, + [4095, 8193, 4094, 8195], + correct_cp_gradients=correct_cp_gradients, + device=torch.device("cuda"), + dtype=dtype, + check_memory=True, + ) diff --git a/cosmos_framework/model/generator/mot/context_parallel_test.py b/cosmos_framework/model/generator/mot/context_parallel_test.py index dc5d01f3..76e33b5d 100644 --- a/cosmos_framework/model/generator/mot/context_parallel_test.py +++ b/cosmos_framework/model/generator/mot/context_parallel_test.py @@ -10,6 +10,9 @@ import pytest import torch import torch.distributed as dist +from torch._dynamo.decorators import mark_unbacked +from torch.distributed.device_mesh import init_device_mesh +from torch.utils.checkpoint import checkpoint from cosmos_framework.trainer import ContextParallelDataWindow, ImaginaireTrainer from cosmos_framework.utils import distributed @@ -28,6 +31,7 @@ from cosmos_framework.model.generator.utils.data_and_condition import GenerationDataClean from cosmos_framework.model.generator.utils.load_balancing_stats import LBLMetadata, compute_sample_lbl_stats +from cosmos_framework.model.generator.utils.memory import MemoryValue from cosmos_framework.data.generator.sequence_packing import ( PackedSequence, build_sequence_plans_from_data_batch, @@ -1153,6 +1157,89 @@ def test_sample_lbl_hsdp_weighting_matches_global_sample_mean() -> None: dist.barrier() +@pytest.mark.L0 +@pytest.mark.parametrize("gen_tokens,kv_heads", [(17, 16), (17, 2), (64, 2)]) +@pytest.mark.parametrize("full_ac", [False, True]) +def test_cp_memory_unbacked_cache_gradients( + monkeypatch: pytest.MonkeyPatch, gen_tokens: int, kv_heads: int, full_ac: bool +) -> None: + """Compile real CP cache extraction without specializing its padded stream lengths.""" + if not torch.cuda.is_available() or int(os.environ.get("WORLD_SIZE", "1")) != 4: + pytest.skip("Requires four distributed CUDA ranks.") + torch.cuda.set_device(int(os.environ["LOCAL_RANK"])) + if not dist.is_initialized(): + dist.init_process_group("nccl") + mesh = init_device_mesh("cuda", (4,)) + monkeypatch.setattr(torch.fx.experimental._config, "use_duck_shape", False) + torch.compiler.reset() + torch.manual_seed(47 + dist.get_rank()) + und_tokens = 11 + tokens = torch.randn(und_tokens + gen_tokens, 16, 8, device="cuda") # [U+G,16,8] + metadata, mask, _ = build_packed_sequence( + "two_way", + packed_sequence=tokens, + attn_modes=["causal", "full"], + split_lens=[und_tokens, gen_tokens], + sample_lens=[und_tokens + gen_tokens], + packed_und_token_indexes=cast(torch.LongTensor, torch.arange(und_tokens, device="cuda")), # [U] + packed_gen_token_indexes=cast( + torch.LongTensor, torch.arange(und_tokens, und_tokens + gen_tokens, device="cuda") + ), # [G] + num_heads=16, + head_dim=8, + num_layers=1, + cp_world_size=4, + full_seq_alignment=4, + causal_seq_alignment=4, + ) + assert isinstance(mask, SplitInfo) + memory = MemoryValue() + + def attention( + q: SequencePack, k: SequencePack, v: SequencePack, attention_mask: SplitInfo, **kwargs: Any + ) -> tuple[SequencePack, None]: + # Isolate the CP exchange and cache boundary from the attention kernel itself. + und = q["causal_seq"] + k["causal_seq"] + v["causal_seq"] # [U_padded,4,8] + gen = q["full_only_seq"] + k["full_only_seq"] + v["full_only_seq"] # [G_padded,4,8] + return from_mode_splits(und, gen, q), None + + def forward(und: torch.Tensor, gen: torch.Tensor) -> tuple[torch.Tensor, ...]: + # Inputs: [U_padded/4,16,8], [G_padded/4,16,8]. Outputs include both streams and real-token caches. + q = from_mode_splits(und, gen, metadata, is_sharded=True) + k = from_mode_splits( + (und * 2)[:, :kv_heads], (gen * 2)[:, :kv_heads], metadata, is_sharded=True + ) # [U_padded/4,H_KV,8], [G_padded/4,H_KV,8] + v = from_mode_splits( + (und * 3)[:, :kv_heads], (gen * 3)[:, :kv_heads], metadata, is_sharded=True + ) # [U_padded/4,H_KV,8], [G_padded/4,H_KV,8] + result, stored = context_parallel_attention( + mesh, q, k, v, mask, attention_function=attention, memory_value=memory + ) + assert stored is not None + return (result["causal_seq"], result["full_only_seq"], *stored) # CP streams and [1,S,H_KV_local,8] caches + + def wrapped(und: torch.Tensor, gen: torch.Tensor) -> tuple[torch.Tensor, ...]: + # Inputs and outputs have the same shapes as forward. + return checkpoint(forward, und, gen, use_reentrant=False) if full_ac else forward(und, gen) + + inputs = [ + torch.randn(metadata[key].shape[0] // 4, 16, 8, device="cuda", requires_grad=True) + for key in ("causal_seq", "full_only_seq") + ] # [U_padded/4,16,8], [G_padded/4,16,8] + reference = [value.detach().clone().requires_grad_() for value in inputs] # same input shapes + expected = forward(*reference) # CP streams and four cache tensors + torch.stack([value.square().mean() for value in expected]).sum().backward() # scalar loss + for value in inputs: + mark_unbacked(value, 0) + actual = torch.compile(wrapped, fullgraph=True, dynamic=True)(*inputs) # CP streams and four cache tensors + torch.stack([value.square().mean() for value in actual]).sum().backward() # scalar loss + for actual_value, expected_value in zip(actual, expected): # matching output shapes + torch.testing.assert_close(actual_value, expected_value, rtol=1e-5, atol=1e-5) + for actual_value, expected_value in zip(inputs, reference): # matching input shapes + torch.testing.assert_close(actual_value.grad, expected_value.grad, rtol=1e-5, atol=1e-5) + dist.barrier() + + def _multiview_maskless_cp_case( samples: list[tuple[int, int, int]], *, diff --git a/cosmos_framework/model/generator/mot/context_parallel_utils.py b/cosmos_framework/model/generator/mot/context_parallel_utils.py index 02eb7720..a7c6194e 100644 --- a/cosmos_framework/model/generator/mot/context_parallel_utils.py +++ b/cosmos_framework/model/generator/mot/context_parallel_utils.py @@ -31,6 +31,7 @@ loss or post-processing. """ +import math from typing import Any, Callable import torch @@ -54,6 +55,88 @@ ) from cosmos_framework.utils.generator.parallelism import ParallelDims +_CP_OUTPUT_GATHER_BYTES = 256 * 1024 * 1024 + + +class _GatherInterleavedOutput(torch.autograd.Function): + """Gather directly into token order with bounded forward/backward communication storage.""" + + @staticmethod + def forward( + ctx: Any, + und: torch.Tensor, # [N_und_local,D] + gen: torch.Tensor, # [N_gen_local,D] + und_indices: torch.Tensor, # [N_und] + gen_indices: torch.Tensor, # [N_gen] + group: dist.ProcessGroup, + correct_cp_gradients: bool, + ) -> torch.Tensor: # [N_und+N_gen,D] + world_size = dist.get_world_size(group) + row_bytes = math.prod(und.shape[1:]) * und.element_size() + chunk_rows = max(1, _CP_OUTPUT_GATHER_BYTES // max(1, world_size * row_bytes)) + ctx.group = group + ctx.correct_cp_gradients = correct_cp_gradients + ctx.local_lengths = (und.shape[0], gen.shape[0]) + ctx.chunk_rows = chunk_rows + # Only index metadata survives forward; neither gathered features nor stream inputs are needed. + und_indices = und_indices.to(torch.int64) # [N_und] + gen_indices = gen_indices.to(torch.int64) # [N_gen] + ctx.save_for_backward(und_indices, gen_indices) + output = und.new_empty((und_indices.numel() + gen_indices.numel(), *und.shape[1:])) # [N,D] + rows = min(chunk_rows, max(und.shape[0], gen.shape[0])) + workspace = und.new_empty((world_size * rows, *und.shape[1:])) # [CP*rows,D] + for stream, indices in ((und, und_indices), (gen, gen_indices)): + local_rows = stream.shape[0] + for start in range(0, local_rows, chunk_rows): + width = min(chunk_rows, local_rows - start) + gathered = workspace[: world_size * width] # [CP*width,D] + local = stream[start : start + width].contiguous() # [width,D] + dist.all_gather_into_tensor(gathered, local, group=group) # [CP*width,D] + for owner in range(world_size): + offset = owner * local_rows + start + count = max(0, min(width, indices.numel() - offset)) + output.index_copy_( + 0, indices[offset : offset + count], gathered[owner * width : owner * width + count] + ) # [N,D] + return output + + @staticmethod + def backward( + ctx: Any, + grad_output: torch.Tensor, # [N,D] + ) -> tuple[torch.Tensor, torch.Tensor, None, None, None, None]: + world_size = dist.get_world_size(ctx.group) + rank = dist.get_rank(ctx.group) + gradients: list[torch.Tensor] = [] + rows = min(ctx.chunk_rows, max(ctx.local_lengths)) if ctx.correct_cp_gradients else 0 + workspace = grad_output.new_empty((world_size * rows, *grad_output.shape[1:])) # [CP*rows,D] + for local_rows, indices in zip(ctx.local_lengths, ctx.saved_tensors, strict=True): + gradient = grad_output.new_zeros((local_rows, *grad_output.shape[1:])) # [N_local,D] + for start in range(0, local_rows, ctx.chunk_rows): + width = min(ctx.chunk_rows, local_rows - start) + if ctx.correct_cp_gradients: + gathered = workspace[: world_size * width].zero_() # [CP*width,D] + for owner in range(world_size): + offset = owner * local_rows + start + count = max(0, min(width, indices.numel() - offset)) + torch.index_select( + grad_output, + 0, + indices[offset : offset + count], + out=gathered[owner * width : owner * width + count], + ) # [count,D] + # Sum the replicated loss contributions before FSDP/DDP averages parameters. + dist.reduce_scatter_tensor(gradient[start : start + width], gathered, group=ctx.group) # [width,D] + else: + # Match the legacy Replicate-to-Shard backward without scaling head gradients. + offset = rank * local_rows + start + count = max(0, min(width, indices.numel() - offset)) + torch.index_select( + grad_output, 0, indices[offset : offset + count], out=gradient[start : start + count] + ) # [count,D] + gradients.append(gradient) + return gradients[0], gradients[1], None, None, None, None + def _pad_to_N(N: int, x: torch.Tensor) -> torch.Tensor: assert x.shape[0] <= N @@ -210,6 +293,7 @@ def get_context_parallel_last_hidden_state( parallel_dims: ParallelDims | None, *, correct_cp_gradients: bool = False, + bounded_memory: bool = False, ) -> torch.Tensor: # packed_outputs streams: [N_local,hidden_size], returns: [N,hidden_size] if parallel_dims is None or not parallel_dims.cp_enabled: return get_all_seq_unpadded(packed_outputs) @@ -219,16 +303,28 @@ def get_context_parallel_last_hidden_state( und_hidden_seq = get_und_seq(packed_outputs) # [text_shard_len,hidden_size] gen_hidden_seq = get_gen_seq(packed_outputs) # [gen_shard_len,hidden_size] - gathered_und_seq = all_gather_tensor( - und_hidden_seq, gather_dim=0, cp_mesh=parallel_dims.cp_mesh, correct_cp_gradients=correct_cp_gradients - ) # [text_len,hidden_size] - gathered_gen_seq = all_gather_tensor( - gen_hidden_seq, gather_dim=0, cp_mesh=parallel_dims.cp_mesh, correct_cp_gradients=correct_cp_gradients - ) # [gen_len,hidden_size] - - gathered_hidden_pack = from_mode_splits(gathered_und_seq, gathered_gen_seq, packed_outputs, is_sharded=False) - last_hidden_state = get_all_seq_unpadded(gathered_hidden_pack) - return last_hidden_state + if not bounded_memory: + gathered_und_seq = all_gather_tensor( + und_hidden_seq, gather_dim=0, cp_mesh=parallel_dims.cp_mesh, correct_cp_gradients=correct_cp_gradients + ) # [text_len,hidden_size] + gathered_gen_seq = all_gather_tensor( + gen_hidden_seq, gather_dim=0, cp_mesh=parallel_dims.cp_mesh, correct_cp_gradients=correct_cp_gradients + ) # [gen_len,hidden_size] + + gathered_hidden_pack = from_mode_splits(gathered_und_seq, gathered_gen_seq, packed_outputs, is_sharded=False) + last_hidden_state = get_all_seq_unpadded(gathered_hidden_pack) # [N,hidden_size] + return last_hidden_state + + # Gathering both full streams and then interleaving duplicates a full hidden state. + # Large teacher-forcing packs exceed memory here even after attention has completed. + return _GatherInterleavedOutput.apply( + und_hidden_seq, + gen_hidden_seq, + packed_outputs["_causal_indices"], + packed_outputs["_full_indices"], + parallel_dims.cp_mesh.get_group(), + correct_cp_gradients, + ) # [N,hidden_size] def all_to_all_tensor( @@ -473,18 +569,21 @@ def context_parallel_attention( if memory_value is not None: und_len = packed_key_states["_num_causal_tokens"] gen_len = packed_key_states["_num_full_tokens"] - # Concrete real-token counts off the metadata, sliced out of gathered streams whose lengths - # are unbacked -- so the same clamp question as the trim below, and the same answer: the - # gathered stream carries every real token plus whatever padding, never fewer. + # Metadata gives the real-token counts; the gathered streams also contain padding. + # With unbacked CP lengths (no concrete size hints), Python slicing may need to + # decide whether to truncate its stop index to the gathered length. This raised + # GuardOnDataDependentSymNode in the Torch 2.13 compiled CP path. + # The checks below ensure the real tokens fit. narrow specifies that exact length, + # returning the same view and preserving the slice's gradients for valid inputs. torch._check(und_len <= k_und_seq.shape[0]) torch._check(gen_len <= k_gen_seq.shape[0]) torch._check(und_len <= v_und_seq.shape[0]) torch._check(gen_len <= v_gen_seq.shape[0]) kv_to_store = ( - k_gen_seq[:gen_len].unsqueeze(0), - v_gen_seq[:gen_len].unsqueeze(0), - k_und_seq[:und_len].unsqueeze(0), - v_und_seq[:und_len].unsqueeze(0), + k_gen_seq.narrow(0, 0, gen_len).unsqueeze(0), # [1,N_gen,H_kv_local,head_dim] + v_gen_seq.narrow(0, 0, gen_len).unsqueeze(0), # [1,N_gen,H_kv_local,head_dim] + k_und_seq.narrow(0, 0, und_len).unsqueeze(0), # [1,N_und,H_kv_local,head_dim] + v_und_seq.narrow(0, 0, und_len).unsqueeze(0), # [1,N_und,H_kv_local,head_dim] ) q_und_seq_len = q_und_seq.shape[0] diff --git a/cosmos_framework/model/generator/mot/cosmos3_vfm_network.py b/cosmos_framework/model/generator/mot/cosmos3_vfm_network.py index fbc4be56..dd51d2ca 100644 --- a/cosmos_framework/model/generator/mot/cosmos3_vfm_network.py +++ b/cosmos_framework/model/generator/mot/cosmos3_vfm_network.py @@ -43,6 +43,7 @@ build_multiview_maskless_plan, ) from cosmos_framework.model.generator.utils.memory import MemoryState +from cosmos_framework.model.generator.utils.rig_view_embedding import add_view_embeddings from cosmos_framework.data.generator.sequence_packing import ModalityData, PackedSequence from cosmos_framework.data.generator.sequence_packing.natten import verify_natten_parameter_list from cosmos_framework.data.generator.sequence_packing.runtime import ( @@ -51,6 +52,7 @@ get_causal_seq, get_full_only_seq, ) +from cosmos_framework.utils.generator.spatial_patch import normalize_spatial_patch_hw class Cosmos3VFMNetworkConfig(PretrainedConfig): @@ -64,12 +66,14 @@ def __init__( latent_downsample_factor=8, latent_channel_size=16, lidar_latent_channel_size=None, + radar_latent_channel_size=None, max_latent_h=32, max_latent_w=32, max_latent_t=32, enable_fps_modulation=False, enable_vision_modality_embeddings: bool = False, enable_media_modality_embedding: bool = False, + num_view_embeddings: int = 0, enable_action_modality_embedding: bool = True, enable_sound_modality_embedding: bool = True, base_fps=24, @@ -78,6 +82,7 @@ def __init__( interpolate_pos=False, timestep_shift=1.0, timestep_scale=0.001, + timestep_range: float = 1.0, predict_text_tokens=False, joint_attn_implementation="two_way", multiview_attention_config: MultiviewAttentionConfig | None = None, @@ -93,21 +98,33 @@ def __init__( temporal_compression_factor_sound=1, sound_latent_fps: int = 25, enable_input_bias: bool = True, + lidar_patch_spatial_hw: int | tuple[int, int] | None = None, + radar_patch_spatial_hw: int | tuple[int, int] | None = None, **kwargs, ): self.vision_gen = vision_gen self.sound_gen = sound_gen self.vlm_config = vlm_config self.latent_patch_size = latent_patch_size + self.lidar_patch_spatial_hw: tuple[int, int] = normalize_spatial_patch_hw( + latent_patch_size if lidar_patch_spatial_hw is None else lidar_patch_spatial_hw + ) + self.radar_patch_spatial_hw: tuple[int, int] = normalize_spatial_patch_hw( + latent_patch_size if radar_patch_spatial_hw is None else radar_patch_spatial_hw + ) self.latent_downsample_factor = latent_downsample_factor self.latent_channel_size = latent_channel_size self.lidar_latent_channel_size = lidar_latent_channel_size + self.radar_latent_channel_size = radar_latent_channel_size self.max_latent_h = max_latent_h self.max_latent_w = max_latent_w self.max_latent_t = max_latent_t self.enable_fps_modulation = enable_fps_modulation self.enable_vision_modality_embeddings = enable_vision_modality_embeddings self.enable_media_modality_embedding = enable_media_modality_embedding + if num_view_embeddings < 0 or num_view_embeddings == 1: + raise ValueError("Rig embeddings need camera IDs and a final LiDAR ID") + self.num_view_embeddings: int = num_view_embeddings self.enable_action_modality_embedding = enable_action_modality_embedding self.enable_sound_modality_embedding = enable_sound_modality_embedding if self.enable_vision_modality_embeddings and self.enable_media_modality_embedding: @@ -120,6 +137,7 @@ def __init__( self.interpolate_pos = interpolate_pos self.timestep_shift = timestep_shift self.timestep_scale = timestep_scale + self.timestep_range = timestep_range self.predict_text_tokens = predict_text_tokens self.joint_attn_implementation = joint_attn_implementation # One object rather than five fields flattened out of it: the mask reads its scope and @@ -232,6 +250,8 @@ def __init__(self, language_model, config: Cosmos3VFMNetworkConfig): if config.vision_gen: self.latent_patch_size = config.latent_patch_size + self.lidar_patch_spatial_hw: tuple[int, int] = config.lidar_patch_spatial_hw + self.radar_patch_spatial_hw: tuple[int, int] = config.radar_patch_spatial_hw self.timestep_shift = config.timestep_shift self.timestep_scale = config.timestep_scale self.latent_downsample = config.latent_downsample_factor * config.latent_patch_size @@ -245,17 +265,35 @@ def __init__(self, language_model, config: Cosmos3VFMNetworkConfig): self.time_embedder = TimestepEmbedder(self.hidden_size, bias=_input_bias) self.vae2llm = nn.Linear(self.patch_latent_dim, self.hidden_size, bias=_input_bias) self.llm2vae = nn.Linear(self.hidden_size, self.patch_latent_dim) + if config.num_view_embeddings: + self.rig_view_embed: nn.Embedding = nn.Embedding(config.num_view_embeddings, self.hidden_size) # LiDAR is its own modality: a range clip enters and leaves the sequence through # its own pair of projections, the way action and sound do. Its VAE is wider than # the camera's (128 vs 48), and these two matrices are what that costs -- a patch - # count follows T, H and W, not channels, so the streams agree on everything else - # and share the grid packing, patchify and timestep machinery below. + # count follows T, H and W, not channels. Each stream can choose its patch size + # while sharing the grid packing, patchify and timestep machinery below. self.lidar_latent_channel = config.lidar_latent_channel_size if self.lidar_latent_channel is not None: - self.lidar_patch_latent_dim = self.latent_patch_size**2 * self.lidar_latent_channel + self.lidar_patch_latent_dim = ( + self.lidar_patch_spatial_hw[0] * self.lidar_patch_spatial_hw[1] * self.lidar_latent_channel + ) self.lidar2llm = nn.Linear(self.lidar_patch_latent_dim, self.hidden_size, bias=_input_bias) self.llm2lidar = nn.Linear(self.hidden_size, self.lidar_patch_latent_dim) + + # Radar BEV is a third sensor stream, standing to the sequence exactly as LiDAR + # does: its own VAE (128 channels, as wide as the LiDAR one) and so its own pair + # of projections. Its clock differs -- radar cycles at ~20 Hz against LiDAR's 10 + # and the camera's 30, and its VAE does not compress time -- but a rate reaches + # attention through the pack's ``seconds_per_frame``, not through a parameter. + # Its DiT patch follows the camera's unless a recipe sets ``radar_patch_spatial_hw``. + self.radar_latent_channel = config.radar_latent_channel_size + if self.radar_latent_channel is not None: + self.radar_patch_latent_dim = ( + self.radar_patch_spatial_hw[0] * self.radar_patch_spatial_hw[1] * self.radar_latent_channel + ) + self.radar2llm = nn.Linear(self.radar_patch_latent_dim, self.hidden_size, bias=_input_bias) + self.llm2radar = nn.Linear(self.hidden_size, self.radar_patch_latent_dim) if config.enable_vision_modality_embeddings: self.image_modality_embed = nn.Parameter(torch.zeros(self.hidden_size)) self.video_modality_embed = nn.Parameter(torch.zeros(self.hidden_size)) @@ -291,6 +329,8 @@ def init_weights(self, buffer_device: torch.device | None): self.time_embedder._init_weights(buffer_device=buffer_device) if self.config.vision_gen: + if self.config.num_view_embeddings: + torch.nn.init.zeros_(self.rig_view_embed.weight) # [V,D] std = 1.0 / math.sqrt(self.patch_latent_dim) torch.nn.init.trunc_normal_(self.vae2llm.weight, std=std, a=-3 * std, b=3 * std) if self.config.enable_input_bias: @@ -312,6 +352,18 @@ def init_weights(self, buffer_device: torch.device | None): torch.nn.init.trunc_normal_(self.llm2lidar.weight, std=std, a=-3 * std, b=3 * std) torch.nn.init.zeros_(self.llm2lidar.bias) + if self.radar_latent_channel is not None: + # As for LiDAR above, and named separately for the same reason: ``std`` below + # is still the llm2vae one the modality embeddings read. + radar_in_std = 1.0 / math.sqrt(self.radar_patch_latent_dim) + torch.nn.init.trunc_normal_( + self.radar2llm.weight, std=radar_in_std, a=-3 * radar_in_std, b=3 * radar_in_std + ) + if self.config.enable_input_bias: + torch.nn.init.zeros_(self.radar2llm.bias) + torch.nn.init.trunc_normal_(self.llm2radar.weight, std=std, a=-3 * std, b=3 * std) + torch.nn.init.zeros_(self.llm2radar.bias) + if self.config.enable_vision_modality_embeddings: torch.nn.init.trunc_normal_(self.image_modality_embed, std=std, a=-3 * std, b=3 * std) torch.nn.init.trunc_normal_(self.video_modality_embed, std=std, a=-3 * std, b=3 * std) @@ -431,13 +483,20 @@ def lidar_gen(self) -> bool: """Whether this network carries the LiDAR stream's own projections.""" return self.config.vision_gen and self.lidar_latent_channel is not None + @property + def radar_gen(self) -> bool: + """Whether this network carries the radar stream's own projections.""" + return self.config.vision_gen and self.radar_latent_channel is not None + def patchify_and_pack_latents( self, tokens_vision: torch.Tensor, token_shapes_vision: Sequence[tuple[int, ...]], latent_channel: int | None = None, + patch_size: int | tuple[int, int] | None = None, ) -> tuple[torch.Tensor, List[Tuple[int, int, int]]]: - p = self.latent_patch_size + patch = self.latent_patch_size if patch_size is None else patch_size + ph, pw = normalize_spatial_patch_hw(patch) # One channel count per call: the caller passes its stream's width, since patches of # different widths cannot pack into one tensor. latent_channel = self.latent_channel if latent_channel is None else latent_channel @@ -453,11 +512,11 @@ def patchify_and_pack_latents( _, t_actual, h_actual, w_actual = latent.shape original_latent_shapes.append((t_actual, h_actual, w_actual)) - # Compute padded dimensions (must be divisible by p) - h_padded = ((h_actual + p - 1) // p) * p - w_padded = ((w_actual + p - 1) // p) * p + # Compute padded dimensions (must be divisible by each patch side) + h_padded = ((h_actual + ph - 1) // ph) * ph + w_padded = ((w_actual + pw - 1) // pw) * pw - # Zero-pad if dimensions are not divisible by p + # Zero-pad if dimensions are not divisible by their patch sides if h_padded != h_actual or w_padded != w_actual: padded = torch.zeros( (latent_channel, t_actual, h_padded, w_padded), @@ -468,15 +527,15 @@ def patchify_and_pack_latents( latent = padded # [C,T,H_padded,W_padded] # Compute number of patches after padding - h_patches = h_padded // p - w_patches = w_padded // p + h_patches = h_padded // ph + w_patches = w_padded // pw # Patchify latent = latent.reshape( - latent_channel, t_actual, h_patches, p, w_patches, p - ) # [C,T,h_patches,p,w_patches,p] + latent_channel, t_actual, h_patches, ph, w_patches, pw + ) # [C,T,h_patches,ph,w_patches,pw] latent = torch.einsum("cthpwq->thwpqc", latent).reshape( - -1, p * p * latent_channel + -1, ph * pw * latent_channel ) # [T*h_patches*w_patches,patch_latent_dim] packed_latent.append(latent) @@ -491,8 +550,10 @@ def unpatchify_and_unpack_latents( noisy_frame_indexes_vision: list[torch.Tensor], original_latent_shapes: List[Tuple[int, int, int]] | None = None, latent_channel: int | None = None, + patch_size: int | tuple[int, int] | None = None, ) -> list[torch.Tensor]: - p = self.latent_patch_size + patch = self.latent_patch_size if patch_size is None else patch_size + ph, pw = normalize_spatial_patch_hw(patch) # One channel count per call, as in ``patchify_and_pack_latents``. latent_channel = self.latent_channel if latent_channel is None else latent_channel unpatchified_latents = [] @@ -504,13 +565,13 @@ def unpatchify_and_unpack_latents( if original_latent_shapes is not None: t_orig, h_orig, w_orig = original_latent_shapes[i] # Compute padded dimensions used during patchify - h_padded = ((h_orig + p - 1) // p) * p - w_padded = ((w_orig + p - 1) // p) * p - h_patches = h_padded // p - w_patches = w_padded // p + h_padded = ((h_orig + ph - 1) // ph) * ph + w_padded = ((w_orig + pw - 1) // pw) * pw + h_patches = h_padded // ph + w_patches = w_padded // pw else: # Fallback: use token shapes directly (assumes no padding was needed) - t_orig, h_orig, w_orig = t_c, h_c * p, w_c * p + t_orig, h_orig, w_orig = t_c, h_c * ph, w_c * pw h_patches, w_patches = h_c, w_c # noisy_frame_indexes_vision is a list of tensors, each with shape (T,), @@ -529,14 +590,16 @@ def unpatchify_and_unpack_latents( end_idx = start_idx + num_patches # Extract patches for this latent latent_patches = packed_mse_preds[start_idx:end_idx] # [num_patches,patch_latent_dim] - # Reshape back to [t_n, h_patches, w_patches, p, p, channels] + # Reshape back to [t_n, h_patches, w_patches, ph, pw, channels] latent_patches = latent_patches.reshape( - t_n, h_patches, w_patches, p, p, latent_channel - ) # [T_n,h_patches,w_patches,p,p,C] + t_n, h_patches, w_patches, ph, pw, latent_channel + ) # [T_n,h_patches,w_patches,ph,pw,C] # Invert the einsum operation: "thwpqc->cthpwq" - latent = torch.einsum("thwpqc->cthpwq", latent_patches) # [C,T_n,h_patches,p,w_patches,p] + latent = torch.einsum("thwpqc->cthpwq", latent_patches) # [C,T_n,h_patches,ph,w_patches,pw] # Reshape back to [channels, t_n, h_padded, w_padded] - latent = latent.reshape(latent_channel, t_n, h_patches * p, w_patches * p) # [C,T_n,H_padded,W_padded] + latent = latent.reshape( + latent_channel, t_n, h_patches * ph, w_patches * pw + ) # [C,T_n,H_padded,W_padded] # Crop to original dimensions (unpad the zeros) latent = latent[:, :, :h_orig, :w_orig] # [C,T_n,H_orig,W_orig] @@ -779,6 +842,7 @@ def _encode_vision( latent_channel=self.latent_channel, modality_embed=modality_embed, target_dtype=target_dtype, + view_ids=packed_seq.vision_view_ids, ) def _encode_lidar( @@ -790,8 +854,9 @@ def _encode_lidar( """Project LiDAR range-view tokens and fill into packed_sequence. Same treatment as the vision stream, through the LiDAR VAE's own width and its own - pair of projections. No modality embedding: the two streams already differ by their - projections and by where mRoPE puts them. + pair of projections. No separate modality embedding: the two streams already differ + by their projections and by where mRoPE puts them. An optional rig embedding uses + the final physical sensor ID for both LiDAR controls and targets. """ return self._encode_grid_stream( packed_seq, @@ -799,8 +864,33 @@ def _encode_lidar( packed_sequence, vae2llm=self.lidar2llm, latent_channel=self.lidar_latent_channel, + patch_size=self.lidar_patch_spatial_hw, + modality_embed=self.rig_view_embed.weight[-1] if self.config.num_view_embeddings else None, # [D] + target_dtype=target_dtype, + ) + + def _encode_radar( + self, + packed_seq: PackedSequence, + packed_sequence: torch.Tensor, + target_dtype: torch.dtype, + ) -> List[Tuple[int, int, int]] | None: + """Project radar BEV tokens and fill into packed_sequence. + + A BEV clip is a grid latent like a range clip, so it takes the same treatment as the + LiDAR stream, through the radar VAE's own width and its own pair of projections. No + modality embedding, for the reason LiDAR has none: the streams already differ by their + projections and by where mRoPE puts them. + """ + return self._encode_grid_stream( + packed_seq, + packed_seq.radar, + packed_sequence, + vae2llm=self.radar2llm, + latent_channel=self.radar_latent_channel, modality_embed=None, target_dtype=target_dtype, + patch_size=self.radar_patch_spatial_hw, ) def _encode_grid_stream( @@ -813,11 +903,14 @@ def _encode_grid_stream( latent_channel: int, modality_embed: torch.Tensor | None, target_dtype: torch.dtype, + patch_size: int | tuple[int, int] | None = None, + view_ids: list[torch.Tensor] | None = None, # one [V] tensor per RGB item ) -> List[Tuple[int, int, int]] | None: """Patchify, project and scatter one stream of VAE latent grids. - Shared by the vision and LiDAR streams so a second grid modality cannot drift from - the first on patchification, timestep embedding or where its tokens land. + Shared by the vision, LiDAR and radar streams so a second or third grid modality + cannot drift from the first on patchification, timestep embedding or where its + tokens land. Returns: Original latent shapes before padding, for unpadding during decode, or ``None`` @@ -836,9 +929,15 @@ def _encode_grid_stream( assert isinstance(modality.mse_loss_indexes, torch.Tensor) packed_patches, original_latent_shapes = self.patchify_and_pack_latents( - modality.tokens, modality.token_shapes, latent_channel=latent_channel + modality.tokens, modality.token_shapes, latent_channel=latent_channel, patch_size=patch_size ) # [total_patches,patch_latent_dim] packed_tokens = vae2llm(packed_patches.to(target_dtype)) # [total_patches,hidden_size] + if self.config.num_view_embeddings and modality is packed_seq.vision: + if view_ids is None: + raise ValueError("RGB tokens require physical view IDs when rig embeddings are enabled") + packed_tokens = add_view_embeddings( + packed_tokens, modality.token_shapes, view_ids, self.rig_view_embed + ) # [total_patches,hidden_size] if modality_embed is not None: packed_tokens = packed_tokens + modality_embed.view(1, -1) # [total_patches,hidden_size] @@ -903,10 +1002,32 @@ def _decode_lidar( llm2vae=self.llm2lidar, latent_channel=self.lidar_latent_channel, patch_latent_dim=self.lidar_patch_latent_dim, + patch_size=self.lidar_patch_spatial_hw, original_latent_shapes=original_latent_shapes, ) ) + def _decode_radar( + self, + packed_seq: PackedSequence, + last_hidden_state: torch.Tensor, + output_dict: dict, + original_latent_shapes: List[Tuple[int, int, int]] | None = None, + ) -> None: + """Decode radar tokens from hidden states and update output_dict.""" + output_dict.update( + preds_radar=self._decode_grid_stream( + packed_seq.radar, + last_hidden_state, + vae2llm=self.radar2llm, + llm2vae=self.llm2radar, + latent_channel=self.radar_latent_channel, + patch_latent_dim=self.radar_patch_latent_dim, + original_latent_shapes=original_latent_shapes, + patch_size=self.radar_patch_spatial_hw, + ) + ) + def _decode_grid_stream( self, modality: ModalityData | None, @@ -917,6 +1038,7 @@ def _decode_grid_stream( latent_channel: int, patch_latent_dim: int, original_latent_shapes: List[Tuple[int, int, int]] | None, + patch_size: int | tuple[int, int] | None = None, ) -> list[torch.Tensor]: """Read one stream's noisy patches back out of the hidden states. @@ -958,6 +1080,7 @@ def _decode_grid_stream( noisy_frame_indexes_vision=modality.noisy_frame_indexes, original_latent_shapes=original_latent_shapes, latent_channel=latent_channel, + patch_size=patch_size, ) def _encode_action( @@ -1204,11 +1327,11 @@ def _prepare_multiview_attention( if packed_seq.action is not None or packed_seq.sound is not None: raise ValueError( - "Multiview FlexAttention supports vision and LiDAR generation batches, not action or sound." + "Multiview FlexAttention supports vision, LiDAR and radar generation batches, not action or sound." ) - if packed_seq.vision is None and packed_seq.lidar is None: - raise ValueError("Multiview FlexAttention needs a vision or LiDAR generation stream.") + if packed_seq.vision is None and packed_seq.lidar is None and packed_seq.radar is None: + raise ValueError("Multiview FlexAttention needs a vision or LiDAR or radar generation stream.") # Before anything is built from the captions -- the mask's items, the folds' plan -- # because which layout the pack is in decides how every one of its tokens is keyed @@ -1218,6 +1341,7 @@ def _prepare_multiview_attention( sensor_mask_items = _multiview_sensor_mask_items( packed_seq, lidar_attends_captions=self.config.multiview_attention_config.mask.lidar_attends_captions, + radar_attends_captions=self.config.multiview_attention_config.mask.radar_attends_captions, ) caption_mask_items = _multiview_caption_mask_items(packed_seq) if caption_mask_items is not None and get_caption_seq_offsets(input_pack) is None: @@ -1290,6 +1414,7 @@ def forward( memory: MemoryState | None = None, video_temporal_causal: bool | None = None, correct_cp_gradients: bool = False, + bounded_cp_output_gather: bool = False, ) -> dict: """ Forward pass for Cosmos3VFMNetwork. @@ -1303,11 +1428,14 @@ def forward( video_temporal_causal: Per-call attention-mode override; ``None`` (default) uses the config-selected ``self.video_temporal_causal``. correct_cp_gradients: Sum CP output gradients before FSDP/DDP averaging. + bounded_cp_output_gather: Opt causal replay into bounded gathering; + bidirectional calls always retain the default gather. Returns: dict with keys: - "preds_vision": list[Tensor[C,T,H,W]], one per sample. - "preds_lidar": Velocity predictions for LiDAR tokens (if the LiDAR stream is configured). + - "preds_radar": Velocity predictions for radar tokens (if the radar stream is configured). - "preds_action": Velocity predictions for action tokens (if action_gen). - "preds_sound": Velocity predictions for sound tokens (if sound_gen). - "last_hidden_state": Last hidden state from the transformer. @@ -1323,6 +1451,7 @@ def forward( # encode vision tokens original_latent_shapes: List[Tuple[int, int, int]] | None = None original_latent_shapes_lidar: List[Tuple[int, int, int]] | None = None + original_latent_shapes_radar: List[Tuple[int, int, int]] | None = None if self.config.vision_gen: original_latent_shapes = self._encode_vision(packed_seq, packed_sequence, target_dtype) @@ -1330,6 +1459,10 @@ def forward( if self.lidar_gen: original_latent_shapes_lidar = self._encode_lidar(packed_seq, packed_sequence, target_dtype) + # encode radar tokens + if self.radar_gen: + original_latent_shapes_radar = self._encode_radar(packed_seq, packed_sequence, target_dtype) + # encode action tokens if self.config.action_gen: self._encode_action(packed_seq, packed_sequence, target_dtype) @@ -1355,6 +1488,8 @@ def forward( all_gen_indexes.append(packed_seq.vision.sequence_indexes) if packed_seq.lidar is not None and isinstance(packed_seq.lidar.sequence_indexes, torch.Tensor): all_gen_indexes.append(packed_seq.lidar.sequence_indexes) + if packed_seq.radar is not None and isinstance(packed_seq.radar.sequence_indexes, torch.Tensor): + all_gen_indexes.append(packed_seq.radar.sequence_indexes) if packed_seq.action is not None and isinstance(packed_seq.action.sequence_indexes, torch.Tensor): all_gen_indexes.append(packed_seq.action.sequence_indexes) if packed_seq.sound is not None and isinstance(packed_seq.sound.sequence_indexes, torch.Tensor): @@ -1494,6 +1629,7 @@ def forward( packed_outputs=packed_outputs, parallel_dims=sequence_shard_parallel_dims, correct_cp_gradients=correct_cp_gradients, + bounded_memory=bounded_cp_output_gather and use_video_temporal_causal, ) # [N_total,hidden_size] output_dict = dict() @@ -1505,6 +1641,10 @@ def forward( if self.lidar_gen: self._decode_lidar(packed_seq, last_hidden_state, output_dict, original_latent_shapes_lidar) + # decode radar tokens + if self.radar_gen: + self._decode_radar(packed_seq, last_hidden_state, output_dict, original_latent_shapes_radar) + # decode action tokens if self.config.action_gen: self._decode_action(packed_seq, last_hidden_state, output_dict) @@ -1616,12 +1756,12 @@ def _multiview_maskless_geometry( passed down as a description of the attention rather than tested as a condition. * at most one sensor item per stream per sample, beside its control item, and no action or - sound. A camera item, a range item, or one of each: a joint sample is served by - quantising both streams' capture times onto the camera's frame grid, so the two need not - share a frame index. A control item ahead of either is served too -- it joins its - target's view groups, and ``control_attends_sensor`` decides whether that group is one - varlen segment or two. A *third* item on one stream is an image-editing layout, which - this path does not serve. + sound. A camera item, a range item, a radar item, or one of each: a joint sample is + served by quantising every stream's capture times onto the camera's frame grid, so they + need not share a frame index. A control item ahead of any of them is served too -- it + joins its target's view groups, and ``control_attends_sensor`` decides whether that group + is one varlen segment or two. A *third* item on one stream is an image-editing layout, + which this path does not serve. * per-view captions are served, but only alongside the pack's per-caption boundaries: the gen->und pass then keys each *view's* GEN tokens against the caption written for that view -- and a range clip against every caption of its sample, since a sweep fuses the rig @@ -1636,10 +1776,10 @@ def _multiview_maskless_geometry( Args: packed_seq: the batch, which is what carries the item and caption structure. sensor_mask_items: the same items the mask is described with, in the same order -- the - packer's, its vision items then its LiDAR ones per sample. Taken rather than - rebuilt, for the reason ``caption_mask_items`` is: what each item is to its sample's - captions (:data:`CaptionAccess`) is one fact about the batch, and the two backends - working it out separately is how the folds came to ignore + packer's, its vision items then its LiDAR ones then its radar ones per sample. + Taken rather than rebuilt, for the reason ``caption_mask_items`` is: what each item + is to its sample's captions (:data:`CaptionAccess`) is one fact about the batch, and + the two backends working it out separately is how the folds came to ignore ``lidar_attends_captions`` while the mask honoured it. caption_mask_items: the batch's caption layout, or ``None`` where every sample packs a single caption. Taken rather than recomputed: the caller derives it for the mask @@ -1665,35 +1805,40 @@ def _multiview_maskless_geometry( ValueError: when this batch cannot be served by the folds. """ num_samples = len(packed_seq.sample_lens) - vision, lidar = packed_seq.vision, packed_seq.lidar - if vision is None and lidar is None: + vision, lidar, radar = packed_seq.vision, packed_seq.lidar, packed_seq.radar + if vision is None and lidar is None and radar is None: raise ValueError( - f"{_MASKLESS_REFUSAL}it carries neither a vision nor a LiDAR generation stream, so " - "there is no sensor grid to fold." + _MASKLESS_REFUSAL_TAIL + f"{_MASKLESS_REFUSAL}it carries neither a vision, a LiDAR nor a radar generation " + "stream, so there is no sensor grid to fold." + _MASKLESS_REFUSAL_TAIL ) - # Items per sample, per stream, in the order the packer lays a sample down: its vision items - # then its LiDAR ones. ``None`` means one item of that stream per sample, which is what the - # counts record for every batch that is not image-editing or transfer. + # Items per sample, per stream, in the order the packer lays a sample down: its vision items, + # then its LiDAR ones, then its radar ones. ``None`` means one item of that stream per sample, + # which is what the counts record for every batch that is not image-editing or transfer. vision_counts = (packed_seq.num_vision_items_per_sample or [1] * num_samples) if vision else [0] * num_samples lidar_counts = (packed_seq.num_lidar_items_per_sample or [1] * num_samples) if lidar else [0] * num_samples - if len(vision_counts) != num_samples or len(lidar_counts) != num_samples: + radar_counts = (packed_seq.num_radar_items_per_sample or [1] * num_samples) if radar else [0] * num_samples + if len(vision_counts) != num_samples or len(lidar_counts) != num_samples or len(radar_counts) != num_samples: raise ValueError( - f"{_MASKLESS_REFUSAL}it records {len(vision_counts)} vision and " - f"{len(lidar_counts)} LiDAR item counts for {num_samples} samples." + _MASKLESS_REFUSAL_TAIL + f"{_MASKLESS_REFUSAL}it records {len(vision_counts)} vision, " + f"{len(lidar_counts)} LiDAR and {len(radar_counts)} radar item counts for " + f"{num_samples} samples." + _MASKLESS_REFUSAL_TAIL ) # At most one sensor item per stream per sample beside its control item, and at least one - # item overall. A sample owning a camera item beside a range item is the joint case: the two - # sensors run at different rates, so the plan quantises both onto the camera's frame grid by - # capture time. The *second* item on a stream is that stream's control item, which the folds - # serve; a third is an image-editing layout, which this path does not -- ``control_weights`` - # catches most of those and this catches the rest. - if any(v > 2 or r > 2 or v + r < 1 for v, r in zip(vision_counts, lidar_counts)): + # item overall. A sample owning a camera item beside a range item or a radar item is the + # joint case: the sensors run at different rates, so the plan quantises each onto the + # camera's frame grid by capture time. The *second* item on a stream is that stream's + # control item, which the folds serve; a third is an image-editing layout, which this path + # does not -- ``control_weights`` catches most of those and this catches the rest. + if any( + cam > 2 or rng > 2 or bev > 2 or cam + rng + bev < 1 + for cam, rng, bev in zip(vision_counts, lidar_counts, radar_counts) + ): raise ValueError( f"{_MASKLESS_REFUSAL}its per-sample item counts are vision={list(vision_counts)}, " - f"lidar={list(lidar_counts)}. Each sample takes at most one item per stream beside " - "its control item, and at least one overall; more is an image-editing layout this " - "path does not serve." + _MASKLESS_REFUSAL_TAIL + f"lidar={list(lidar_counts)}, radar={list(radar_counts)}. Each sample takes at most " + "one item per stream beside its control item, and at least one overall; more is an " + "image-editing layout this path does not serve." + _MASKLESS_REFUSAL_TAIL ) views_per_vision_item = packed_seq.num_views_per_vision_item or [] @@ -1711,9 +1856,9 @@ def _multiview_maskless_geometry( items_per_sample: list[int] = [] is_control: list[bool] = [] view_axis: list[int] = [] - vision_cursor = lidar_cursor = 0 - for vision_count, lidar_count in zip(vision_counts, lidar_counts): - items_per_sample.append(vision_count + lidar_count) + vision_cursor = lidar_cursor = radar_cursor = 0 + for vision_count, lidar_count, radar_count in zip(vision_counts, lidar_counts, radar_counts): + items_per_sample.append(vision_count + lidar_count + radar_count) # The camera item first, which is the order the packer lays a sample down and the order # the plan anchors on: a joint sample quantises capture time onto its *first* item's # frame grid, so anchoring on the camera keeps its tokens on the frame indices a @@ -1742,14 +1887,26 @@ def _multiview_maskless_geometry( is_control.append(index < lidar_count - 1) view_axis.append(1) lidar_cursor += 1 + for index in range(radar_count): + assert radar is not None + # A BEV clip covers the scene around the rig rather than one of its cameras, so it + # is one "view", as a sweep is -- but on a third axis: a radar frame and a sweep + # are no more the same view than either is a camera's, and the two sensors cycle + # at 20 Hz and 10 Hz, so they do not even share an instant. + num_views.append(1) + token_shapes.append(tuple(radar.token_shapes[radar_cursor])) # type: ignore[arg-type] + rates.append(float(radar.seconds_per_frame[radar_cursor])) + is_control.append(index < radar_count - 1) + view_axis.append(2) + radar_cursor += 1 # The divisibility of each latent_t by its view count is checked by SensorMaskItem, which the # caller built for these same items before reaching here, and again by the plan builder. # Captions as (view_id, num_tokens) per sample, the two parallel lists the packer records. # None keeps the gen->und pass on its per-sample form, which is what a single caption wants. # Flattened in the same order this walked the items above -- per sample, its vision items - # then its LiDAR ones -- which is the order the packer lays a sample down and the order - # _multiview_sensor_mask_items builds in. The count check is what holds the two together. + # then its LiDAR ones then its radar ones -- which is the order the packer lays a sample down + # and the order _multiview_sensor_mask_items builds in. The count check holds the two together. caption_accesses = [item.caption_access for sample_items in sensor_mask_items for item in sample_items] if len(caption_accesses) != len(num_views): raise ValueError( @@ -1820,11 +1977,11 @@ def _multiview_caption_mask_items(packed_seq: PackedSequence) -> list[list[Capti def _multiview_sensor_mask_items( - packed_seq: PackedSequence, *, lidar_attends_captions: bool = True + packed_seq: PackedSequence, *, lidar_attends_captions: bool = True, radar_attends_captions: bool = True ) -> list[list[SensorMaskItem]]: - """Describe each sample to the multiview mask as its vision items, then its LiDAR items. + """Describe each sample to the multiview mask: its vision items, its LiDAR, then its radar. - The packer lays a sample out in exactly that order, so walking the two streams sample by + The packer lays a sample out in exactly that order, so walking the three streams sample by sample reproduces the packed order the mask assumes. LiDAR items take a view offset past the cameras. A range clip is not one of the rig's @@ -1840,7 +1997,14 @@ def _multiview_sensor_mask_items( default ``"camera"`` either way, which is also what says the per-view captions have to cover their views and not the sweep's. - Control items are marked per stream, not per sample: within each of the two streams, + Radar items take a view offset past the LiDAR one, for the same reason LiDAR takes one past + the cameras and one step further: a BEV clip is neither one of the rig's views nor the + sweep, and the three sensors run at 30 fps, 10 Hz and 20 Hz, so two of them sharing a view + id would have the view rules pair latents captured at different instants. They read the + captions as a sweep does -- ``"all_captions"``, or ``"no_captions"`` under + ``radar_attends_captions=False``. + + Control items are marked per stream, not per sample: within each of the three streams, every item but the last is a control item conditioning the one that follows it, which is the same convention the packer uses when it forces those items fully clean (``packers.py``). Doing it per stream is what keeps a camera item's position from @@ -1857,8 +2021,9 @@ def _multiview_sensor_mask_items( num_samples = len(packed_seq.sample_lens) vision = packed_seq.vision lidar = packed_seq.lidar - if vision is None and lidar is None: - raise ValueError("Multiview FlexAttention needs a vision or LiDAR generation stream.") + radar = packed_seq.radar + if vision is None and lidar is None and radar is None: + raise ValueError("Multiview FlexAttention needs a vision or LiDAR or radar generation stream.") # None means every sample owns exactly one vision item (standard T2V/I2V); # multi-item samples (image editing, transfer) carry explicit counts. A @@ -1870,15 +2035,15 @@ def _multiview_sensor_mask_items( vision_counts = packed_seq.num_vision_items_per_sample or [1] * num_samples views_per_vision_item = list(packed_seq.num_views_per_vision_item or []) if not views_per_vision_item: - if lidar is None: + if lidar is None and radar is None: raise ValueError( "Multiview FlexAttention requires per-camera VAE metadata; " "enable enable_per_camera_vae_encoding on the dataset." ) - # A pack carrying both streams cannot hold that metadata: it is written by the - # camera-major uint8 encode path, which a range clip never takes. Such a pack is - # single-camera by construction, so one view per item is the grid the mask needs, - # and the only thing it needs the count for. + # A pack carrying a second sensor cannot hold that metadata: it is written by the + # camera-major uint8 encode path, which neither a range clip nor a BEV clip takes. + # Such a pack is single-camera by construction, so one view per item is the grid + # the mask needs, and the only thing it needs the count for. views_per_vision_item = [1] * sum(vision_counts) lidar_counts = [0] * num_samples @@ -1887,12 +2052,23 @@ def _multiview_sensor_mask_items( # Step past the widest camera item so no LiDAR item can land on a camera's view. lidar_view_offset = max(views_per_vision_item, default=0) + radar_counts = [0] * num_samples + if radar is not None: + radar_counts = packed_seq.num_radar_items_per_sample or [1] * num_samples + # Every LiDAR item of a sample shares the one view above, so one more step clears the + # stream entirely and leaves the three sensors pairwise disjoint on the view axis. A pack + # without LiDAR takes no step for it, so a camera + radar pack numbers its radar exactly as + # a camera + LiDAR pack numbers its sweeps, and a radar-only pack stays on view 0. + radar_view_offset = lidar_view_offset + (1 if lidar is not None else 0) + sensor_mask_items: list[list[SensorMaskItem]] = [] vision_cursor = 0 lidar_cursor = 0 + radar_cursor = 0 for sample_idx in range(num_samples): sample_items: list[SensorMaskItem] = [] num_vision, num_lidar = vision_counts[sample_idx], lidar_counts[sample_idx] + num_radar = radar_counts[sample_idx] if vision is not None: for item_in_stream in range(num_vision): sample_items.append( @@ -1925,6 +2101,23 @@ def _multiview_sensor_mask_items( ) ) lidar_cursor += 1 + if radar is not None: + for item_in_stream in range(num_radar): + sample_items.append( + SensorMaskItem( + token_shape=radar.token_shapes[radar_cursor], + condition_mask=radar.condition_mask[radar_cursor], + num_views=1, + view_offset=radar_view_offset, + is_control=item_in_stream < num_radar - 1, + seconds_per_frame=radar.seconds_per_frame[radar_cursor], + # A BEV clip is not one of the rig's cameras either: it covers the + # scene around the whole rig, so it reads every camera's caption -- + # or, cut off from the text, none. + caption_access="all_captions" if radar_attends_captions else "no_captions", + ) + ) + radar_cursor += 1 sensor_mask_items.append(sample_items) return sensor_mask_items diff --git a/cosmos_framework/model/generator/mot/maskless_attention.py b/cosmos_framework/model/generator/mot/maskless_attention.py new file mode 100644 index 00000000..99dd331a --- /dev/null +++ b/cosmos_framework/model/generator/mot/maskless_attention.py @@ -0,0 +1,659 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +"""Causal replay using the upstream maskless sensor/caption passes and LSE merge. + +Replay roles determine each query group's keys. Grouping happens once before the +decoder, over metadata runs rather than a quadratic token mask. Queries sharing +the same keys share one varlen group, including across frames in a replay chunk. +The two sensor passes deliberately overlap, as in bidirectional maskless attention. +""" + +from collections.abc import Sequence +from dataclasses import dataclass, fields +from itertools import accumulate + +import attrs +import torch + +from cosmos_framework.model.attention import attention +from cosmos_framework.model.attention.backends import choose_backend +from cosmos_framework.model.attention.natten import ( + NATTEN_BLACKWELL_DETERMINISTIC_VERSION, + NATTEN_SUPPORTED, + natten_version_satisfies, +) +from cosmos_framework.model.attention.natten.checks import choose_natten_backend +from cosmos_framework.model.attention.utils.environment import is_torch_compiling +from cosmos_framework.model.generator.mot.flex_attention import SensorMaskItem +from cosmos_framework.model.generator.mot.merge_bridge import MergeAttentionsBridge +from cosmos_framework.model.generator.mot.multiview_maskless_attention import ( + _gather_from_packed, + _scatter_to_packed, +) +from cosmos_framework.model.generator.mot.causal_flex_attention import ( + _ROLE_PADDING, + _ROLE_UND, + TeacherForcingFlexMetadata, + _key_stream_fields, + _query_stream_fields, + _stream_metadata_groups, + _StreamFields, + _teacher_forcing_pair_predicate, +) +from cosmos_framework.model.generator.mot.merge_attention import merge_attentions_ac_safe + +if NATTEN_SUPPORTED and natten_version_satisfies(NATTEN_BLACKWELL_DETERMINISTIC_VERSION): + import natten.backends.blackwell_fmha as _blackwell_fmha +else: + _blackwell_fmha = None + +_REPLAY_KV_BATCH_TOKENS = 1 << 22 + + +@dataclass(frozen=True) +class ReplayMasklessPass: + """Nonempty varlen groups for one upstream maskless partition.""" + + name: str + q_gather: torch.Tensor # [N_Q_pass] + kv_gather: torch.Tensor # [N_KV_pass] + q_offsets: torch.Tensor # [G+1] + kv_offsets: torch.Tensor # [G+1] + q_max_len: int + kv_max_len: int + + +@dataclass(frozen=True) +class ReplayMasklessPlan: + """Layer-independent plan in global (post-Ulysses) Q/KV coordinates.""" + + passes: tuple[ReplayMasklessPass, ...] + q_len: int + kv_len: int + real_queries: torch.Tensor # [Q] + batches: tuple[ReplayMasklessPass, ...] = () + compiled_indices: tuple[torch.Tensor, ...] = () # Gather indices and packed batch offsets + batch_layout: torch.Tensor | None = None # [B,9] on CPU + + +def _batch_pass(partition: ReplayMasklessPass) -> tuple[ReplayMasklessPass, ...]: + """Bound repeated KV gathers without splitting an independent attention group.""" + q_offsets = partition.q_offsets.cpu().tolist() # list[G+1] + kv_offsets = partition.kv_offsets.cpu().tolist() # list[G+1] + batches: list[ReplayMasklessPass] = [] + first = 0 + while first < len(q_offsets) - 1: + last = first + 1 + while last < len(q_offsets) - 1 and kv_offsets[last + 1] - kv_offsets[first] <= _REPLAY_KV_BATCH_TOKENS: + last += 1 + batches.append( + ReplayMasklessPass( + name=partition.name, + q_gather=partition.q_gather[q_offsets[first] : q_offsets[last]], # [N_Q_batch] + kv_gather=partition.kv_gather[kv_offsets[first] : kv_offsets[last]], # [N_KV_batch] + q_offsets=partition.q_offsets[first : last + 1] - q_offsets[first], # [G_batch+1] + kv_offsets=partition.kv_offsets[first : last + 1] - kv_offsets[first], # [G_batch+1] + q_max_len=max(q_offsets[i + 1] - q_offsets[i] for i in range(first, last)), + kv_max_len=max(kv_offsets[i + 1] - kv_offsets[i] for i in range(first, last)), + ) + ) + first = last + return tuple(batches) + + +def _attend_batch( + q: torch.Tensor, # [1,N_Q_batch,H,D] + k: torch.Tensor, # [1,N_KV_batch,H_KV,D] + v: torch.Tensor, # [1,N_KV_batch,H_KV,D] + batch: ReplayMasklessPass, +) -> tuple[torch.Tensor, torch.Tensor]: # [1,N_Q_batch,H,D], [1,N_Q_batch,H] + return attention( + q, + k, + v, + cumulative_seqlen_Q=batch.q_offsets, + cumulative_seqlen_KV=batch.kv_offsets, + max_seqlen_Q=batch.q_max_len, + max_seqlen_KV=batch.kv_max_len, + return_lse=True, + ) # [1,N_Q_batch,H,D], [1,N_Q_batch,H] + + +def _blackwell_backward_config( + q: torch.Tensor, # [1,Q,H,D] + k: torch.Tensor, # [1,KV,H_KV,D] + v: torch.Tensor, # [1,KV,H_KV,D] +) -> tuple[int, int] | None: + """Reuse the selected native backward without changing backend preferences.""" + if _blackwell_fmha is None or q.device.type != "cuda": + return None + deterministic = torch.are_deterministic_algorithms_enabled() + backend = choose_backend( + query_shape=q.shape, + key_shape=k.shape, + value_shape=v.shape, + dtype=q.dtype, + device=q.device, + requires_grad=True, + is_causal=False, + causal_type=None, + is_varlen=True, + deterministic=deterministic, + return_lse=True, + is_compiling=is_torch_compiling(), + raise_error=False, + ) + if ( + backend != "natten" + or choose_natten_backend( + query_shape=q.shape, + key_shape=k.shape, + value_shape=v.shape, + dtype=q.dtype, + device=q.device, + requires_grad=True, + is_causal=False, + is_varlen=True, + deterministic=deterministic, + ) + != "blackwell-fmha" + ): + return None + return _blackwell_fmha.check_cutlass_blackwell_fmha_backward_config(input_tensor=q) + + +def _merge_batch( + output: torch.Tensor, # [1,Q,H,D], FP32 accumulation + merged_lse: torch.Tensor, # [1,Q,H], FP32 + indices: torch.Tensor, # [N_Q_batch], unique within this batch + out: torch.Tensor, # [1,N_Q_batch,H,D] + lse: torch.Tensor, # [1,N_Q_batch,H] +) -> None: + previous_lse = merged_lse[:, indices] # [1,N_Q_batch,H] + combined_lse = torch.logaddexp(previous_lse, lse) # [1,N_Q_batch,H] + combined_out = output[:, indices] * torch.exp(previous_lse - combined_lse).unsqueeze(-1) + out.float() * torch.exp( + lse - combined_lse + ).unsqueeze(-1) # [1,N_Q_batch,H,D] + output.index_copy_(1, indices, combined_out) # [1,Q,H,D] + merged_lse.index_copy_(1, indices, combined_lse) # [1,Q,H] + + +# Fuse merge intermediates independently of transformer compilation. +_compiled_merge_batch = torch.compile(_merge_batch, fullgraph=True, dynamic=True) + + +def _batched_forward( + q: torch.Tensor, # [1,Q,H,D] + k: torch.Tensor, # [1,KV,H_KV,D] + v: torch.Tensor, # [1,KV,H_KV,D] + batches: tuple[ReplayMasklessPass, ...], + *, + compiled: bool = False, +) -> tuple[torch.Tensor, torch.Tensor]: # [1,Q,H,D], [1,Q,H] + output = torch.zeros_like(q, dtype=torch.float32) # [1,Q,H,D] + merged_lse = torch.full(q.shape[:-1], -torch.inf, dtype=torch.float32, device=q.device) # [1,Q,H] + merge = _compiled_merge_batch if compiled else _merge_batch + for batch in batches: + out, lse = _attend_batch( + q[:, batch.q_gather], # [1,N_Q_batch,H,D] + k[:, batch.kv_gather], # [1,N_KV_batch,H_KV,D] + v[:, batch.kv_gather], # [1,N_KV_batch,H_KV,D] + batch, + ) # [1,N_Q_batch,H,D], [1,N_Q_batch,H] + lse = lse.squeeze(-1) if lse.ndim == 4 else lse # [1,N_Q_batch,H] + merge(output, merged_lse, batch.q_gather, out, lse) # [1,Q,H,D], [1,Q,H], updated in place + # The merged buffers own this group's result. Release its temporaries + # before gathering the next group or casting the final output. + del out, lse + output = output.to(q.dtype) # [1,Q,H,D] + return output, merged_lse # [1,Q,H,D], [1,Q,H] + + +def _batched_backward( + q: torch.Tensor, # [1,Q,H,D] + k: torch.Tensor, # [1,KV,H_KV,D] + v: torch.Tensor, # [1,KV,H_KV,D] + output: torch.Tensor, # [1,Q,H,D] + merged_lse: torch.Tensor, # [1,Q,H] + grad_output: torch.Tensor, # [1,Q,H,D] + batches: tuple[ReplayMasklessPass, ...], + backward_config: tuple[int, int] | None, + deterministic: bool, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: # Q: [1,Q,H,D], KV: [1,KV,H_KV,D] + gradients = [torch.zeros_like(tensor, dtype=torch.float32) for tensor in (q, k, v)] # each input shape + for batch in batches: + indices = (batch.q_gather, batch.kv_gather, batch.kv_gather) # [N_Q_batch], 2*[N_KV_batch] + if backward_config is not None: + assert _blackwell_fmha is not None + inputs = [tensor[:, index] for tensor, index in zip((q, k, v), indices)] # each gathered input shape + # The native backward already accepts Q/K/V and final O/LSE. + # Rebuilding its autograd graph would run a redundant forward. + grads = _blackwell_fmha.blackwell_fmha_backward( + *inputs, + output[:, batch.q_gather], # [1,N_Q_batch,H,D] + grad_output[:, batch.q_gather], # [1,N_Q_batch,H,D] + merged_lse[:, batch.q_gather], # [1,N_Q_batch,H] + False, + q.shape[-1] ** -0.5, + *backward_config, + batch.q_offsets, + batch.kv_offsets, + batch.q_max_len, + batch.kv_max_len, + deterministic, + ) # each gathered input shape + else: + with torch.enable_grad(): + inputs = [ + tensor[:, index].detach().requires_grad_() for tensor, index in zip((q, k, v), indices) + ] # [1,N_Q_batch,H,D], 2*[1,N_KV_batch,H_KV,D] + out, lse = _attend_batch(*inputs, batch) # [1,N_Q_batch,H,D], [1,N_Q_batch,H] + # Native backward must read the final merged O/LSE, just as it + # does through MergeAttentionsBridge. Recomputing unpatched + # component outputs would silently lose cross-pass gradients. + out.data.copy_(output[:, batch.q_gather]) # [1,N_Q_batch,H,D] + patched_lse = merged_lse[:, batch.q_gather] # [1,N_Q_batch,H] + lse.data.copy_(patched_lse.unsqueeze(-1) if lse.ndim == 4 else patched_lse) # native LSE shape + grads = torch.autograd.grad(out, inputs, grad_output[:, batch.q_gather]) # each input shape + del out, lse, patched_lse + for total, index, gradient in zip(gradients, indices, grads): + total.index_add_(1, index, gradient.float()) # full input shape + # Do not overlap a completed group's gathered inputs/gradients with + # the next native backward. Loop variables also retain their last tensors. + del inputs, grads, gradient, total + # Retire each FP32 accumulation buffer as its returned gradient is cast; + # retaining the list through all three casts overlaps both full copies. + dq = gradients.pop(0).to(q.dtype) # [1,Q,H,D] + dk = gradients.pop(0).to(k.dtype) # [1,KV,H_KV,D] + dv = gradients.pop(0).to(v.dtype) # [1,KV,H_KV,D] + return dq, dk, dv # input shapes + + +def _decode_batches( + indices: list[torch.Tensor], # 2*[N_pass] per partition, then 2*[sum(G_batch+1)] + layout: torch.Tensor, # [B,9] on CPU +) -> tuple[ReplayMasklessPass, ...]: + """Read runtime batch boundaries without specializing the compiled transformer.""" + batches = [] + for partition, qs, qe, ks, ke, offset_start, offset_end, q_max, kv_max in layout.tolist(): + batches.append( + ReplayMasklessPass( + name="compiled", + q_gather=indices[2 * partition][qs:qe], # [N_Q_batch] + kv_gather=indices[2 * partition + 1][ks:ke], # [N_KV_batch] + q_offsets=indices[-2][offset_start:offset_end], # [G_batch+1] + kv_offsets=indices[-1][offset_start:offset_end], # [G_batch+1] + q_max_len=q_max, + kv_max_len=kv_max, + ) + ) + return tuple(batches) + + +@torch.library.custom_op("cosmos3::bounded_replay_forward", mutates_args=()) +def _opaque_replay_forward( + q: torch.Tensor, # [1,Q,H,D] + k: torch.Tensor, # [1,KV,H_KV,D] + v: torch.Tensor, # [1,KV,H_KV,D] + indices: list[torch.Tensor], # Gather indices and packed batch offsets + layout: torch.Tensor, # [B,9] on CPU +) -> tuple[torch.Tensor, torch.Tensor]: # [1,Q,H,D], [1,Q,H] + return _batched_forward(q, k, v, _decode_batches(indices, layout), compiled=True) # [1,Q,H,D], [1,Q,H] + + +@_opaque_replay_forward.register_fake +def _opaque_replay_forward_fake( + q: torch.Tensor, # [1,Q,H,D] + k: torch.Tensor, # [1,KV,H_KV,D] + v: torch.Tensor, # [1,KV,H_KV,D] + indices: list[torch.Tensor], # Gather indices and packed batch offsets + layout: torch.Tensor, # [B,9] on CPU +) -> tuple[torch.Tensor, torch.Tensor]: # [1,Q,H,D], [1,Q,H] + return torch.empty_like(q), q.new_empty(q.shape[:-1], dtype=torch.float32) # [1,Q,H,D], [1,Q,H] + + +@torch.library.custom_op("cosmos3::bounded_replay_backward", mutates_args=()) +def _opaque_replay_backward( + q: torch.Tensor, # [1,Q,H,D] + k: torch.Tensor, # [1,KV,H_KV,D] + v: torch.Tensor, # [1,KV,H_KV,D] + output: torch.Tensor, # [1,Q,H,D] + merged_lse: torch.Tensor, # [1,Q,H] + grad_output: torch.Tensor, # [1,Q,H,D] + indices: list[torch.Tensor], # Gather indices and packed batch offsets + layout: torch.Tensor, # [B,9] on CPU + backward_config: list[int], + deterministic: bool, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: # Q: [1,Q,H,D], KV: [1,KV,H_KV,D] + config = (backward_config[0], backward_config[1]) if backward_config else None + return _batched_backward( + q, k, v, output, merged_lse, grad_output, _decode_batches(indices, layout), config, deterministic + ) # Q: [1,Q,H,D], KV: [1,KV,H_KV,D] + + +@_opaque_replay_backward.register_fake +def _opaque_replay_backward_fake( + q: torch.Tensor, # [1,Q,H,D] + k: torch.Tensor, # [1,KV,H_KV,D] + v: torch.Tensor, # [1,KV,H_KV,D] + output: torch.Tensor, # [1,Q,H,D] + merged_lse: torch.Tensor, # [1,Q,H] + grad_output: torch.Tensor, # [1,Q,H,D] + indices: list[torch.Tensor], # Gather indices and packed batch offsets + layout: torch.Tensor, # [B,9] on CPU + backward_config: list[int], + deterministic: bool, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: # Q: [1,Q,H,D], KV: [1,KV,H_KV,D] + return torch.empty_like(q), torch.empty_like(k), torch.empty_like(v) # input shapes + + +class _BatchedReplayAttention(torch.autograd.Function): + """Gather one KV batch at a time, retaining the native merged-O/LSE gradient contract.""" + + @staticmethod + def forward( + ctx: torch.autograd.function.FunctionCtx, + q: torch.Tensor, # [1,Q,H,D] + k: torch.Tensor, # [1,KV,H_KV,D] + v: torch.Tensor, # [1,KV,H_KV,D] + batches: tuple[ReplayMasklessPass, ...], + indices: tuple[torch.Tensor, ...] = (), # Gather indices and packed batch offsets + layout: torch.Tensor | None = None, # [B,9] on CPU + ) -> torch.Tensor: # [1,Q,H,D] + ctx.backward_config = _blackwell_backward_config(q, k, v) + ctx.deterministic = torch.are_deterministic_algorithms_enabled() + ctx.opaque = layout is not None + if layout is not None: + output, merged_lse = _opaque_replay_forward(q, k, v, list(indices), layout) # [1,Q,H,D], [1,Q,H] + ctx.save_for_backward(q, k, v, output, merged_lse, *indices, layout) + else: + # CUDA replay can fuse this merge while the surrounding model stays eager. + output, merged_lse = _batched_forward(q, k, v, batches, compiled=q.is_cuda) # [1,Q,H,D], [1,Q,H] + ctx.save_for_backward(q, k, v, output, merged_lse) + ctx.batches = batches + return output # [1,Q,H,D] + + @staticmethod + def backward( + ctx: torch.autograd.function.FunctionCtx, + grad_output: torch.Tensor, # [1,Q,H,D] + ) -> tuple[torch.Tensor | None, ...]: + q, k, v, output, merged_lse, *metadata = ctx.saved_tensors # Q: [1,Q,H,D], KV: [1,KV,H_KV,D] + if ctx.opaque: + gradients = _opaque_replay_backward( + q, + k, + v, + output, + merged_lse, + grad_output, + metadata[:-1], + metadata[-1], + list(ctx.backward_config or ()), + ctx.deterministic, + ) # Q: [1,Q,H,D], KV: [1,KV,H_KV,D] + else: + gradients = _batched_backward( + q, k, v, output, merged_lse, grad_output, ctx.batches, ctx.backward_config, ctx.deterministic + ) # Q: [1,Q,H,D], KV: [1,KV,H_KV,D] + return (*gradients, *((None,) * (len(ctx.needs_input_grad) - 3))) # input shapes + + +class _ConcatReplayKV(torch.autograd.Function): + """Keep a queued clean-cache gradient from retaining the full KV gradient.""" + + @staticmethod + def forward( + ctx: torch.autograd.function.FunctionCtx, + *parts: torch.Tensor, # each [S_part,H,D], last part is clean memory + ) -> torch.Tensor: # [S_total,H,D] + ctx.lengths = tuple(part.shape[0] for part in parts) + return torch.cat(parts) # [S_total,H,D] + + @staticmethod + def backward( + ctx: torch.autograd.function.FunctionCtx, + gradient: torch.Tensor, # [S_total,H,D] + ) -> tuple[torch.Tensor, ...]: # each [S_part,H,D] + parts = gradient.split(ctx.lengths) # tuple[[S_part,H,D]] + # Clean replay backpropagates after the noisy pass. Ordinary CatBackward + # returns views, keeping caption/current-token storage alive at every + # layer until then. Only the clean-memory slice needs independent storage. + return (*parts[:-1], parts[-1].clone()) # tuple[[S_part,H,D]] + + +def cat_replay_kv(parts: list[torch.Tensor]) -> torch.Tensor: # each [S_part,H,D]; returns [S_total,H,D] + """Concatenate replay KV with a compact gradient for the final clean-memory part.""" + return _ConcatReplayKV.apply(*parts) # [S_total,H,D] + + +def _run_fields(fields_: _StreamFields) -> tuple[_StreamFields, list[int]]: + """Copy one representative per metadata run to the CPU for grouping.""" + _, representatives = _stream_metadata_groups(fields_, fields_.sample_id.device) # [S], [R] + compact = _StreamFields( + **{field.name: getattr(fields_, field.name)[representatives].cpu() for field in fields(fields_)} # each [R] + ) + return compact, representatives.tolist() + [fields_.sample_id.numel()] + + +def _instant_ids( + fields_: _StreamFields, sensor_items: Sequence[Sequence[SensorMaskItem]] +) -> tuple[torch.Tensor, torch.Tensor]: # returns [R], [R] + """Use upstream camera-anchored midpoint buckets, including cached frame ids.""" + instants = torch.full_like(fields_.frame_id, -1) # [R] + has_cross_view = torch.zeros_like(fields_.is_control) # [R] + for sample, items in enumerate(sensor_items): + camera = next((item for item in items if item.caption_access == "camera"), None) + if camera is None: + raise ValueError("Maskless replay requires a camera anchor in every sample.") + view_count = len( + {view for item in items for view in range(item.view_offset, item.view_offset + item.num_views)} + ) + has_cross_view |= (fields_.sample_id == sample) & (view_count > 1) # [R] + for item in items: + belongs = ( + (fields_.sample_id == sample) + & (fields_.view_id >= item.view_offset) + & (fields_.view_id < item.view_offset + item.num_views) + ) # [R] + bucket = torch.floor( + (fields_.frame_id.double() + 0.5) * item.seconds_per_frame / camera.seconds_per_frame + ).long() # [R] + instants = torch.where(belongs, bucket, instants) # [R] + return instants, has_cross_view + + +def _gather_ranges(ranges: list[tuple[int, int]], device: torch.device) -> torch.Tensor: # returns [N] + """Expand CPU token intervals with one device allocation per index tensor.""" + lengths = [stop - start for start, stop in ranges] + offsets = list(accumulate(lengths, initial=0)) + starts = torch.tensor([start - offset for (start, _), offset in zip(ranges, offsets)], device=device) # [R] + repeats = torch.tensor(lengths, device=device) # [R] + return torch.repeat_interleave(starts, repeats, output_size=offsets[-1]) + torch.arange( + offsets[-1], device=device + ) # [N] + + +def _build_pass( + name: str, + visibility: torch.Tensor, # [R_Q,R_KV] on CPU + q_starts: list[int], + kv_starts: list[int], + device: torch.device, +) -> ReplayMasklessPass | None: + """Coalesce equal key sets; omit queries with no keys instead of empty kernels.""" + groups: dict[tuple[int, ...], list[int]] = {} + for query_run, row in enumerate(visibility): # row: [R_KV] + key_runs = tuple(row.nonzero().flatten().tolist()) + if key_runs: + groups.setdefault(key_runs, []).append(query_run) + if not groups: + return None + q_ranges: list[tuple[int, int]] = [] + kv_ranges: list[tuple[int, int]] = [] + q_lengths: list[int] = [] + kv_lengths: list[int] = [] + for keys, queries in groups.items(): + q_ranges.extend((q_starts[index], q_starts[index + 1]) for index in queries) + kv_ranges.extend((kv_starts[index], kv_starts[index + 1]) for index in keys) + q_lengths.append(sum(q_starts[index + 1] - q_starts[index] for index in queries)) + kv_lengths.append(sum(kv_starts[index + 1] - kv_starts[index] for index in keys)) + return ReplayMasklessPass( + name=name, + q_gather=_gather_ranges(q_ranges, device), # [N_Q_pass] + kv_gather=_gather_ranges(kv_ranges, device), # [N_KV_pass] + q_offsets=torch.tensor(list(accumulate(q_lengths, initial=0)), dtype=torch.int32, device=device), # [G+1] + kv_offsets=torch.tensor(list(accumulate(kv_lengths, initial=0)), dtype=torch.int32, device=device), # [G+1] + q_max_len=max(q_lengths), + kv_max_len=max(kv_lengths), + ) + + +def build_replay_maskless_plan( + metadata: TeacherForcingFlexMetadata, + sensor_items: Sequence[Sequence[SensorMaskItem]], +) -> ReplayMasklessPlan: + """Intersect replay roles/causality with the upstream maskless partitions.""" + policy = metadata.teacher_forcing_replay_policy + if policy.multiview_attention_scope not in ("same_view", "decomposed"): + raise ValueError("Maskless replay supports same_view or decomposed scope.") + if policy.decomposed_temporal_window_seconds is not None: + raise ValueError( + "Maskless replay requires decomposed_temporal_window_seconds=None for same-instant cross-view attention." + ) + q, q_starts = _run_fields(_query_stream_fields(metadata)) + kv, kv_starts = _run_fields(_key_stream_fields(metadata)) + # Spatial scope is applied by the partitions below. Keep every existing role, + # sample, condition, caption and causal rule in the shared replay predicate. + role_policy = attrs.evolve(policy, multiview_attention_scope="all_views") + q_ids = torch.arange(q.sample_id.numel())[:, None] # [R_Q,1] + kv_ids = torch.arange(kv.sample_id.numel())[None, :] # [1,R_KV] + unused = torch.tensor(0) # [] + allowed = _teacher_forcing_pair_predicate(q, kv, role_policy)(unused, unused, q_ids, kv_ids) # [R_Q,R_KV] + caption = kv.token_role_id[None, :] == _ROLE_UND # [1,R_KV] + padding = (q.token_role_id[:, None] == _ROLE_PADDING) & (kv.token_role_id[None, :] == _ROLE_PADDING) # [R_Q,R_KV] + same_view = q.view_id[:, None] == kv.view_id[None, :] # [R_Q,R_KV] + partitions = [("same_view", allowed & ~caption & (same_view | padding))] # list[[R_Q,R_KV]] + if policy.multiview_attention_scope == "decomposed": + q_instants, cross_enabled = _instant_ids(q, sensor_items) # [R_Q], [R_Q] + kv_instants, _ = _instant_ids(kv, sensor_items) # [R_KV], [R_KV] + # Match the bidirectional maskless model's camera-anchored instants. + # All spatial tokens at that instant remain visible, subject to replay roles. + cross_time = q_instants[:, None] == kv_instants[None, :] # [R_Q,R_KV] + cross = ( + allowed + & ~caption + & ~padding + & ~q.is_control[:, None] + & ~kv.is_control[None, :] + & cross_enabled[:, None] + & cross_time + ) # [R_Q,R_KV] + # Preserve production maskless weighting: a same-view key present in + # both partitions contributes twice to the merged softmax denominator. + partitions.append(("cross_instant", cross)) + partitions.append(("caption", allowed & caption)) # [R_Q,R_KV] + device = metadata.sample_id.device + passes = tuple( + plan + for name, visibility in partitions + if (plan := _build_pass(name, visibility, q_starts, kv_starts, device)) is not None + ) + if not passes: + raise ValueError("Maskless replay has no nonempty attention groups.") + batches: list[ReplayMasklessPass] = [] + layout: list[tuple[int, ...]] = [] + offset = 0 + if any(partition.kv_gather.numel() > _REPLAY_KV_BATCH_TOKENS for partition in passes): + for index, partition in enumerate(passes): + for batch in _batch_pass(partition): + qs = batch.q_gather.storage_offset() - partition.q_gather.storage_offset() + ks = batch.kv_gather.storage_offset() - partition.kv_gather.storage_offset() + end = offset + batch.q_offsets.numel() + layout.append( + ( + index, + qs, + qs + batch.q_gather.numel(), + ks, + ks + batch.kv_gather.numel(), + offset, + end, + batch.q_max_len, + batch.kv_max_len, + ) + ) + batches.append(batch) + offset = end + # Keep the large gather indices as views of the existing partitions. Only + # small offsets and CPU batch boundaries are packed for the opaque operator. + compiled_indices = ( + tuple(tensor for partition in passes for tensor in (partition.q_gather, partition.kv_gather)) + + (torch.cat([batch.q_offsets for batch in batches]), torch.cat([batch.kv_offsets for batch in batches])) + if batches + else () + ) # 2*[N_pass] per partition, then 2*[sum(G_batch+1)] + return ReplayMasklessPlan( + passes=passes, + q_len=metadata.q_len, + kv_len=metadata.seq_len, + real_queries=metadata.query.token_role_id != _ROLE_PADDING, # [Q] + batches=tuple(batches), + compiled_indices=compiled_indices, + batch_layout=torch.tensor(layout, dtype=torch.int64, device="cpu") if batches else None, # [B,9] on CPU + ) + + +def replay_maskless_attention( + q: torch.Tensor, # [1,Q,H,D] + k: torch.Tensor, # [1,KV,H_KV,D] + v: torch.Tensor, # [1,KV,H_KV,D] + plan: ReplayMasklessPlan, +) -> torch.Tensor: # [1,Q,H,D] + """Run native varlen passes with upstream bridges preserving cached-K/V gradients.""" + if is_torch_compiling(): + # State the plan invariants without branching on unbacked sequence lengths. + torch._check(q.shape[1] == plan.q_len) + torch._check(k.shape[1] == plan.kv_len) + torch._check(v.shape[1] == plan.kv_len) + elif q.shape[1] != plan.q_len or k.shape[1] != plan.kv_len or v.shape[1] != plan.kv_len: + raise ValueError("Maskless replay plan does not match the global Q/KV streams.") + if ( + is_torch_compiling() + and plan.batch_layout is not None + and torch.is_grad_enabled() + and _blackwell_backward_config(q, k, v) is not None + ): + # Batch count varies per pack. Do not unroll that Python loop into each + # compiled transformer or retain a separate kernel set for every count. + output = _BatchedReplayAttention.apply(q, k, v, (), plan.compiled_indices, plan.batch_layout) # [1,Q,H,D] + return torch.where(plan.real_queries[None, :, None, None], output, 0) # [1,Q,H,D] + if plan.batches: + output = _BatchedReplayAttention.apply(q, k, v, plan.batches) # [1,Q,H,D] + return torch.where(plan.real_queries[None, :, None, None], output, 0) # [1,Q,H,D] + outputs: list[torch.Tensor] = [] # list[[1,Q,H,D]] + lses: list[torch.Tensor] = [] # list[[1,Q,H]] + for partition in plan.passes: + out, lse = attention( + q[:, partition.q_gather], # [1,N_Q_pass,H,D] + k[:, partition.kv_gather], # [1,N_KV_pass,H_KV,D] + v[:, partition.kv_gather], # [1,N_KV_pass,H_KV,D] + cumulative_seqlen_Q=partition.q_offsets, + cumulative_seqlen_KV=partition.kv_offsets, + max_seqlen_Q=partition.q_max_len, + max_seqlen_KV=partition.kv_max_len, + return_lse=True, + ) # [1,N_Q_pass,H,D], [1,N_Q_pass,H] + out, lse = MergeAttentionsBridge.apply( + out, + lse, + _scatter_to_packed(partition.q_gather, plan.q_len), + _gather_from_packed(partition.q_gather), + ) # [1,Q,H,D], [1,Q,H] + outputs.append(out) + lses.append(lse) + output, _ = ( + (outputs[0], lses[0]) if len(outputs) == 1 else merge_attentions_ac_safe(outputs=outputs, lse_tensors=lses) + ) # [1,Q,H,D], [1,Q,H] + return torch.where(plan.real_queries[None, :, None, None], output, 0) # [1,Q,H,D] diff --git a/cosmos_framework/model/generator/mot/merge_attention.py b/cosmos_framework/model/generator/mot/merge_attention.py new file mode 100644 index 00000000..34693962 --- /dev/null +++ b/cosmos_framework/model/generator/mot/merge_attention.py @@ -0,0 +1,104 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +"""Activation-checkpoint-compatible LSE merge shared by interactive attention paths.""" + +from typing import Any + +import torch + + +class _ACSafeMergeAttentionsFn(torch.autograd.Function): + """AC-compatible drop-in for NATTEN's ``MergeAttentionsAutogradFn``. + + NATTEN's backward indexes ``ctx.saved_tensors`` in multiple slices + (``[:2]``, ``[2 : N+2]``, ``[N+2:]``), and each indexing access fires + the non-reentrant ``torch.utils.checkpoint`` unpack hook for *every* + saved tensor. The hook only permits one unpack per saved tensor, so + activation checkpointing + NATTEN merge_attentions raises + ``CheckpointError: Unpack is being triggered for a tensor that was + already unpacked once`` (see the replayed-LSE + AC=full long-video + TF path). + + This version preserves the same forward math and the same + storage-patching backward contract (see :class:`MergeAttentionsBridge` + docstring for the full description), but reads ``ctx.saved_tensors`` + exactly once. Numerics match ``naive_merge_attentions`` (iterative + pairwise LSE rescale) — which is what NATTEN's kernel implements up + to reduction order. + """ + + @staticmethod + def forward( + ctx: Any, + num_components: int, + *tensors: torch.Tensor, # outputs: [B,Q,H,D], LSEs: [B,Q,H] or [B,Q,H,1] + ) -> tuple[torch.Tensor, torch.Tensor]: + outputs = tensors[:num_components] + lses = tensors[num_components:] + output_dtype = outputs[0].dtype + normalized_lses = [lse.squeeze(-1) if lse.ndim == 4 else lse for lse in lses] # list[[B,Q,H]] + + merged_lse = normalized_lses[0] # [B,Q,H] + merged_out = outputs[0] # [B,Q,H,D] + for i in range(1, num_components): + new_lse = torch.logaddexp(merged_lse, normalized_lses[i]) # [B,Q,H] + w_old = torch.exp(merged_lse - new_lse).unsqueeze(-1) # [B,Q,H,1] + w_new = torch.exp(normalized_lses[i] - new_lse).unsqueeze(-1) # [B,Q,H,1] + merged_out = w_old * merged_out + w_new * outputs[i] # [B,Q,H,D] + merged_lse = new_lse # [B,Q,H] + merged_out = merged_out.to(output_dtype) # [B,Q,H,D] + + ctx.save_for_backward(merged_out, merged_lse, *outputs, *lses) + ctx.num_components = num_components + return merged_out, merged_lse # [B,Q,H,D], [B,Q,H] + + @staticmethod + def backward( + ctx: Any, + grad_merged_out: torch.Tensor, # [B,Q,H,D] + grad_merged_lse: torch.Tensor, # [B,Q,H] + ) -> tuple[torch.Tensor | None, ...]: + # Single access — avoid retriggering the AC unpack hook for any saved tensor. + saved = ctx.saved_tensors + merged_out = saved[0] # [B,Q,H,D] + merged_lse = saved[1] # [B,Q,H] + num = ctx.num_components + outputs = saved[2 : 2 + num] + lses = saved[2 + num : 2 + 2 * num] + + # Patch each component's storage with the merged O / LSE. The + # upstream attention kernel's backward will read these as its + # saved O / LSE and compute gradients as if it had produced the + # merged output. The original LSE shape is preserved (the + # forward squeezes a trailing singleton, so we re-broadcast). + for o in outputs: + o.data.copy_(merged_out.data) # [B,Q,H,D] + for l in lses: + if l.ndim == merged_lse.ndim + 1 and l.shape[-1] == 1: + l.data.copy_(merged_lse.data.unsqueeze(-1)) # [B,Q,H,1] + else: + l.data.copy_(merged_lse.data) # [B,Q,H] + + # Same upstream-grad contract as NATTEN: dL/dO_i = dL/dO_merged + # for every component; dL/dLSE_i is forwarded unchanged for + # parity (i4 attention treats LSE as non-differentiable, so this + # gradient is silently dropped at the kernel boundary). + grads = (None,) + (grad_merged_out,) * num + (grad_merged_lse,) * num + return grads + + +def merge_attentions_ac_safe( + outputs: list[torch.Tensor], # list[[B,Q,H,D]] + lse_tensors: list[torch.Tensor], # list[[B,Q,H] or [B,Q,H,1]] +) -> tuple[torch.Tensor, torch.Tensor]: + """AC-safe drop-in for ``cosmos_framework.model.attention.merge_attentions``. + + Use at call sites that live inside an activation-checkpointed module + boundary. Matches NATTEN's storage-patching backward contract so + upstream i4 attention kernels (whose LSE is not differentiable) + still receive correct gradients via their own saved O / LSE + backward formulas. + """ + assert len(outputs) == len(lse_tensors) >= 2 + return _ACSafeMergeAttentionsFn.apply(len(outputs), *outputs, *lse_tensors) # [B,Q,H,D], [B,Q,H] diff --git a/cosmos_framework/model/generator/mot/unified_mot.py b/cosmos_framework/model/generator/mot/unified_mot.py index 2c95692e..971aced8 100644 --- a/cosmos_framework/model/generator/mot/unified_mot.py +++ b/cosmos_framework/model/generator/mot/unified_mot.py @@ -1117,7 +1117,8 @@ def _impl_forward( _meta_tensor = get_gen_seq(pack) # [S_gen,D] device = _meta_tensor.device cos, sin = self.rotary_emb( - _meta_tensor, position_ids=position_ids.unsqueeze(0) if position_ids.ndim == 1 else position_ids.unsqueeze(1) + _meta_tensor, + position_ids=position_ids.unsqueeze(0) if position_ids.ndim == 1 else position_ids.unsqueeze(1), ) # if ndim == 2, the mrope position_ids is (3, seq_len); inject the batch dim in the # middle to get (3, 1, seq_len) so the rotary_emb's mrope branch broadcasts correctly. # In both branches Qwen3VLTextRotaryEmbedding.apply_interleaved_mrope collapses the @@ -1128,7 +1129,6 @@ def _impl_forward( from_all_seq(cos, pack), from_all_seq(sin, pack), ) - # Tracking the load balancing loss across all layers. For dense models, lbl_metadata_all # will be a dictionary with empty lists for each pathway. For MoE models, the lists # for each pathway will be populated with the load balancing loss metadata for each layer. diff --git a/cosmos_framework/model/generator/multiview_transfer_ar.py b/cosmos_framework/model/generator/multiview_transfer_ar.py index 16e59c0e..cb65451b 100644 --- a/cosmos_framework/model/generator/multiview_transfer_ar.py +++ b/cosmos_framework/model/generator/multiview_transfer_ar.py @@ -151,7 +151,8 @@ def create_session( if prefill_pack.vision is None or len(prefill_pack.vision.token_shapes) != 2: raise ValueError("Multiview transfer AR requires packed [control, target] vision metadata.") flex_backend = getattr(self.host.net, "flex_backend", None) - if flex_backend is None: + maskless_replay = getattr(self.host.net, "teacher_forcing_maskless", False) + if flex_backend is None and not maskless_replay: raise ValueError("Multiview transfer AR requires an initialized FlexAttention backend.") if text_view_ids is not None and text_view_ids != list(range(num_views)): raise ValueError( @@ -161,7 +162,7 @@ def create_session( control_shape, target_shape = prefill_pack.vision.token_shapes total_memory_tokens = control_shape[0] * control_shape[1] * control_shape[2] total_memory_tokens += target_shape[0] * target_shape[1] * target_shape[2] - kv_alignment = int(flex_backend.block_size[1]) + kv_alignment = 1 if maskless_replay else int(flex_backend.block_size[1]) memory_seq_len = ((total_memory_tokens + kv_alignment - 1) // kv_alignment) * kv_alignment target_condition_ranges = [(0, condition_count)] if condition_count else [] num_layers = int(self.host.net.num_hidden_layers) # type: ignore[attr-defined] diff --git a/cosmos_framework/model/generator/omni_mot_causal_model.py b/cosmos_framework/model/generator/omni_mot_causal_model.py index 9b1f8150..a80eacfd 100644 --- a/cosmos_framework/model/generator/omni_mot_causal_model.py +++ b/cosmos_framework/model/generator/omni_mot_causal_model.py @@ -13,6 +13,7 @@ from __future__ import annotations import contextlib +import copy import itertools from collections.abc import Callable, Generator, Iterable, Sequence from dataclasses import dataclass @@ -133,14 +134,14 @@ class reads them directly; they will migrate here once the corresponding # Select the replay implementation through one causal-model input instead # of reconstructing it from lower-level attention settings. The model - # resolves this selector to the internal two-way FlexAttention or three-way + # resolves this selector to the internal two-way Flex/maskless or three-way # attention layout while it builds the network. teacher_forcing_kv_implementation: TeacherForcingKVImplementation = attrs.field( default="singleview_threeway_kv", - validator=attrs.validators.in_(("multiview_flex_kv", "singleview_threeway_kv")), + validator=attrs.validators.in_(("multiview_flex_kv", "multiview_maskless_kv", "singleview_threeway_kv")), ) - # Backend-neutral connectivity for clean replay and transfer control. Both + # Backend-neutral connectivity for clean replay and transfer control. All # replay implementations consume the same policy object. teacher_forcing_replay_policy: TeacherForcingReplayPolicyConfig = attrs.Factory(TeacherForcingReplayPolicyConfig) @@ -515,7 +516,7 @@ def _resolve_teacher_forcing_replay_policy(value: Any) -> TeacherForcingReplayPo def _resolve_teacher_forcing_kv_implementation(value: Any) -> TeacherForcingKVImplementation: """Validate the public replay selector after LazyConfig resolution.""" - supported_implementations = ("multiview_flex_kv", "singleview_threeway_kv") + supported_implementations = ("multiview_flex_kv", "multiview_maskless_kv", "singleview_threeway_kv") if value not in supported_implementations: raise ValueError( f"teacher_forcing_kv_implementation must be one of {supported_implementations}, got {value!r}." @@ -528,9 +529,12 @@ def _validate_teacher_forcing_kv_strategy( causal_training_strategy: str, ) -> None: """Reject selectors that cannot be honored by the configured training strategy.""" - if implementation == "multiview_flex_kv" and causal_training_strategy not in _TEACHER_FORCING_REPLAY_STRATEGIES: + if ( + implementation != "singleview_threeway_kv" + and causal_training_strategy not in _TEACHER_FORCING_REPLAY_STRATEGIES + ): raise ValueError( - "teacher_forcing_kv_implementation='multiview_flex_kv' requires causal_training_strategy " + f"teacher_forcing_kv_implementation={implementation!r} requires causal_training_strategy " f"to be one of {_TEACHER_FORCING_REPLAY_STRATEGIES}, got {causal_training_strategy!r}." ) @@ -640,7 +644,7 @@ def build_net( lora_enabled: bool | None = None, ) -> torch.nn.Module: """Resolve the selected teacher-forcing KV implementation and build it.""" - uses_multiview_flex_kv = self._uses_multiview_flex_kv() + uses_multiview_replay_kv = self._uses_multiview_replay_kv() if self.config.causal_training_strategy not in _TEACHER_FORCING_REPLAY_STRATEGIES: return super().build_net(dtype, mp_policy=mp_policy, lora_enabled=lora_enabled) @@ -653,10 +657,14 @@ def build_net( joint_attn_implementation = self.config.joint_attn_implementation attention_scope = self.config.multiview_attention.mask.attention_scope decomposed_temporal_window_seconds = self.config.multiview_attention.mask.decomposed_temporal_window_seconds + multiview_backend = self.config.multiview_attention.backend + maskless_replay = ( + video_temporal_causal and self._get_teacher_forcing_kv_implementation() == "multiview_maskless_kv" + ) # One knob rather than two that had to agree: the pathway is what selects multiview # attention, so there is no second flag to save and restore alongside it. - self.config.joint_attn_implementation = "multiview" if uses_multiview_flex_kv else "three_way" - if uses_multiview_flex_kv: + self.config.joint_attn_implementation = "multiview" if uses_multiview_replay_kv else "three_way" + if uses_multiview_replay_kv: # Core validates temporal causality as a three-way-only layout. The # replay mask supplies causality for this two-way path. self.config.video_temporal_causal = False @@ -664,8 +672,14 @@ def build_net( self.config.multiview_attention.mask.decomposed_temporal_window_seconds = ( replay_policy.decomposed_temporal_window_seconds ) + if maskless_replay: + self.config.multiview_attention.backend = "maskless" + # Core's bidirectional geometry has no sliding-window fold. + # Replay replaces that geometry with its own same-instant + # visibility plan before decoder execution. + self.config.multiview_attention.mask.decomposed_temporal_window_seconds = None try: - if uses_multiview_flex_kv: + if uses_multiview_replay_kv: with patch.object( omni_mot_model_module, "Cosmos3VFMNetwork", @@ -674,17 +688,23 @@ def build_net( net = super().build_net(dtype, mp_policy=mp_policy, lora_enabled=lora_enabled) else: net = super().build_net(dtype, mp_policy=mp_policy, lora_enabled=lora_enabled) + if maskless_replay: + # Each network owns the resolved backend geometry, independently + # of the student/teacher selectors restored on the model below. + net.config.multiview_attention_config = copy.deepcopy(self.config.multiview_attention) finally: self.config.video_temporal_causal = video_temporal_causal self.config.joint_attn_implementation = joint_attn_implementation self.config.multiview_attention.mask.attention_scope = attention_scope self.config.multiview_attention.mask.decomposed_temporal_window_seconds = decomposed_temporal_window_seconds + self.config.multiview_attention.backend = multiview_backend - if uses_multiview_flex_kv: + if uses_multiview_replay_kv: net.config.video_temporal_causal = video_temporal_causal net.video_temporal_causal = video_temporal_causal setattr(net, "teacher_forcing_replay_policy", replay_policy) setattr(net, "teacher_forcing_frames_per_chunk", self.config.teacher_forcing_frames_per_chunk) + setattr(net, "teacher_forcing_maskless", maskless_replay) return net def maybe_convert_linears_to_nvfp4(self) -> None: @@ -733,7 +753,7 @@ def memory_init_training( # The legacy chunkwise path requires divisibility and therefore drops # trailing latent frames in lockstep across modalities. The multiview Flex # path represents camera-major partial tails explicitly in its mask metadata. - if not self._uses_multiview_flex_kv(): + if not self._uses_multiview_replay_kv(): gen_data_clean = self._truncate_for_chunkwise_tf(gen_data_clean) self._assert_chunkwise_tf_shape(gen_data_clean) @@ -826,19 +846,19 @@ def _is_chunkwise_tf(self) -> bool: "teacher_forcing_dcm", ) - def _uses_multiview_flex_kv(self) -> bool: - """Whether this run selected replayed multiview Flex K/V.""" + def _uses_multiview_replay_kv(self) -> bool: + """Whether this run selected the shared two-way multiview replay layout.""" implementation = self._get_teacher_forcing_kv_implementation() _validate_teacher_forcing_kv_strategy(implementation, self.config.causal_training_strategy) - return ( - self.config.causal_training_strategy in _TEACHER_FORCING_REPLAY_STRATEGIES - and implementation == "multiview_flex_kv" + return self.config.causal_training_strategy in _TEACHER_FORCING_REPLAY_STRATEGIES and implementation in ( + "multiview_flex_kv", + "multiview_maskless_kv", ) @override def _pack_input_sequence(self, *args: Any, **kwargs: Any) -> PackedSequence: - """Keep the standard multiview layout when Flex supplies causality.""" - if not self._uses_multiview_flex_kv(): + """Keep the standard multiview layout when replay supplies causality.""" + if not self._uses_multiview_replay_kv(): return super()._pack_input_sequence(*args, **kwargs) video_temporal_causal = self.config.video_temporal_causal self.config.video_temporal_causal = False @@ -1096,7 +1116,7 @@ def _build_clean_tf_cache( ) -> TeacherForcingMemoryState: """Build replayed clean K/V for one denoiser network.""" clean_target_indexes: torch.Tensor | None = None - if self._uses_multiview_flex_kv(): + if self._uses_multiview_replay_kv(): if packed_sequence.vision is None or packed_sequence.num_views_per_vision_item is None: raise ValueError("Two-way Flex teacher forcing requires multiview vision metadata.") # Match the packer's per-sample order: RGB items, then LiDAR items. @@ -1114,9 +1134,10 @@ def _build_clean_tf_cache( device=packed_sequence.text_ids.device, ) # [S_clean_real] flex_backend = getattr(net, "flex_backend", None) - if flex_backend is None: + maskless_replay = getattr(net, "teacher_forcing_maskless", False) + if flex_backend is None and not maskless_replay: raise ValueError("Two-way Flex teacher forcing requires the network FlexAttention backend.") - kv_alignment = flex_backend.block_size[1] + kv_alignment = 1 if maskless_replay else flex_backend.block_size[1] selected_clean_target_padded_capacity = ( (clean_target_indexes.numel() + kv_alignment - 1) // kv_alignment ) * kv_alignment @@ -1134,7 +1155,7 @@ def _build_clean_tf_cache( selected_clean_gen_token_indexes=clean_target_indexes, ) clean_pack = make_teacher_forcing_clean_pack(packed_sequence) - if self._uses_multiview_flex_kv(): + if self._uses_multiview_replay_kv(): clean_pack.teacher_forcing_pass = "clean" ctx = torch.no_grad() if detach_clean_kv else contextlib.nullcontext() with ctx: @@ -1170,7 +1191,7 @@ def _validate_teacher_forcing_pack(self, packed_seq: PackedSequence) -> None: f"or aligned [control, target] items; got {packed_seq.num_vision_items_per_sample} " f"with {len(packed_seq.sample_lens)} logical samples." ) - if self._uses_multiview_flex_kv(): + if self._uses_multiview_replay_kv(): if packed_seq.vision is None or packed_seq.action is not None or packed_seq.sound is not None: raise ValueError("Two-way Flex teacher forcing supports RGB and optional LiDAR generation batches.") if packed_seq.num_views_per_vision_item is None: @@ -1920,7 +1941,7 @@ def iter_samples_from_batch_autoregressive( self._reset_ar_forward_cuda_graph_runtime_for_generation() reset_ar_post_saturation_runtime_for_generation(self) - if mode == "video_transfer" and self._uses_multiview_flex_kv(): + if mode == "video_transfer" and self._uses_multiview_replay_kv(): if self.config.compile.enabled: raise ValueError("Multiview transfer AR requires eager attention; run with --no-use-torch-compile.") yield from self._iter_samples_multiview_transfer_autoregressive( @@ -2763,7 +2784,7 @@ def _iter_samples_multiview_transfer_autoregressive( has_negative_prompt: bool, ) -> Generator[dict[str, Any], torch.Tensor | None, None]: """Generate the target for a two-item camera-major multiview transfer sample.""" - if not self._uses_multiview_flex_kv(): + if not self._uses_multiview_replay_kv(): raise ValueError("Multiview transfer AR requires replayed two-way Flex teacher-forcing configuration.") if self.config.action_gen or self.config.sound_gen: raise ValueError("Multiview transfer AR supports vision-only models.") diff --git a/cosmos_framework/model/generator/omni_mot_causal_model_test.py b/cosmos_framework/model/generator/omni_mot_causal_model_test.py index 8f8828f2..684413cf 100644 --- a/cosmos_framework/model/generator/omni_mot_causal_model_test.py +++ b/cosmos_framework/model/generator/omni_mot_causal_model_test.py @@ -25,12 +25,20 @@ import pytest import torch +from cosmos_framework.configs.base.defaults.replay_attention import ( + TEACHER_FORCING_KV_IMPLEMENTATIONS, + TeacherForcingKVImplementation, +) + @pytest.mark.L0 @pytest.mark.CPU @pytest.mark.parametrize("view_scoped", [False, True]) @pytest.mark.parametrize("detach_clean_kv", [False, True]) -def test_clean_tf_cache_preserves_target_indexes_and_caption_layout(view_scoped: bool, detach_clean_kv: bool) -> None: +@pytest.mark.parametrize("maskless_replay", [False, True]) +def test_clean_tf_cache_preserves_target_indexes_and_caption_layout( + view_scoped: bool, detach_clean_kv: bool, maskless_replay: bool +) -> None: """Exercise the real mask builder and index helper through the model entry point.""" from cosmos_framework.data.generator.sequence_packing.sequence import ModalityData, PackedSequence from cosmos_framework.model.generator.omni_mot_causal_model import OmniMoTCausalModel @@ -55,8 +63,11 @@ def test_clean_tf_cache_preserves_target_indexes_and_caption_layout(view_scoped: packed_sequence.text_caption_lens = [[1, 1]] packed_sequence.text_caption_view_ids = [[0, 1]] model = MagicMock() - model._uses_multiview_flex_kv.return_value = True - net = SimpleNamespace(flex_backend=SimpleNamespace(block_size=(128, 128))) + model._uses_multiview_replay_kv.return_value = True + net = SimpleNamespace( + teacher_forcing_maskless=maskless_replay, + flex_backend=None if maskless_replay else SimpleNamespace(block_size=(128, 128)), + ) clean_pack = MagicMock() with patch( "cosmos_framework.model.generator.omni_mot_causal_model.make_teacher_forcing_clean_pack", @@ -70,7 +81,7 @@ def test_clean_tf_cache_preserves_target_indexes_and_caption_layout(view_scoped: selected_indexes = call_kwargs["selected_clean_gen_token_indexes"] # [N_clean] torch.testing.assert_close(selected_indexes, torch.tensor([5, 7])) # [N_clean] assert call_kwargs["detach_clean_kv"] is detach_clean_kv - assert packed_sequence.teacher_forcing_selected_clean_target_padded_capacity == 128 + assert packed_sequence.teacher_forcing_selected_clean_target_padded_capacity == (2 if maskless_replay else 128) assert packed_sequence.teacher_forcing_pass == "noisy" assert packed_sequence.text_caption_lens == ([[1, 1]] if view_scoped else []) assert packed_sequence.text_caption_view_ids == ([[0, 1]] if view_scoped else []) @@ -283,7 +294,7 @@ def test_target_only_teacher_forcing_requires_one_logical_sample( model = MagicMock() model.config.video_temporal_causal = True model.config.teacher_forcing_target_only_no_text_pass2 = True - model._uses_multiview_flex_kv.return_value = False + model._uses_multiview_replay_kv.return_value = False packed_sequence = PackedSequence( sample_lens=sample_lens, vision=ModalityData( @@ -423,13 +434,37 @@ def test_teacher_forcing_replay_policy_resolves_lazy_config_and_runs_validation( _resolve_teacher_forcing_replay_policy(mutated_policy) +@pytest.mark.L0 +@pytest.mark.CPU +@pytest.mark.parametrize("implementation", TEACHER_FORCING_KV_IMPLEMENTATIONS) +def test_public_teacher_forcing_kv_implementation_survives_config_resolution( + implementation: TeacherForcingKVImplementation, +) -> None: + """Every public backend must pass both attrs validation and LazyConfig resolution.""" + from cosmos_framework.utils.lazy_config import LazyCall as L + from cosmos_framework.model.generator.omni_mot_causal_model import ( + OmniMoTCausalModel, + OmniMoTCausalModelConfig, + _resolve_teacher_forcing_kv_implementation, + ) + + config = OmniMoTCausalModelConfig(teacher_forcing_kv_implementation=implementation) + lazy_model = L(OmniMoTCausalModel)(config=config, _recursive_=False) + + assert ( + _resolve_teacher_forcing_kv_implementation(lazy_model.config.teacher_forcing_kv_implementation) + == implementation + ) + + @pytest.mark.L0 @pytest.mark.CPU def test_teacher_forcing_kv_implementation_validates_lazy_config_value() -> None: """LazyConfig must not bypass validation for the public replay selector.""" from cosmos_framework.model.generator.omni_mot_causal_model import _resolve_teacher_forcing_kv_implementation - assert _resolve_teacher_forcing_kv_implementation("multiview_flex_kv") == "multiview_flex_kv" + for implementation in ("multiview_flex_kv", "multiview_maskless_kv", "singleview_threeway_kv"): + assert _resolve_teacher_forcing_kv_implementation(implementation) == implementation with pytest.raises(ValueError, match="teacher_forcing_kv_implementation"): _resolve_teacher_forcing_kv_implementation("unknown") @@ -437,30 +472,34 @@ def test_teacher_forcing_kv_implementation_validates_lazy_config_value() -> None @pytest.mark.L0 @pytest.mark.CPU @pytest.mark.parametrize("causal_training_strategy", ["teacher_forcing", "teacher_forcing_dcm"]) -def test_multiview_flex_selection_does_not_infer_from_backend_knobs(causal_training_strategy: str) -> None: - """The public selector, rather than internal attention flags, chooses Flex replay.""" +@pytest.mark.parametrize("implementation", ["multiview_flex_kv", "multiview_maskless_kv"]) +def test_multiview_replay_selection_does_not_infer_from_backend_knobs( + causal_training_strategy: str, implementation: str +) -> None: + """The public selector, rather than internal attention flags, chooses multiview replay.""" from cosmos_framework.model.generator.omni_mot_causal_model import OmniMoTCausalModel model = object.__new__(OmniMoTCausalModel) torch.nn.Module.__init__(model) model.config = SimpleNamespace( causal_training_strategy=causal_training_strategy, - teacher_forcing_kv_implementation="multiview_flex_kv", + teacher_forcing_kv_implementation=implementation, joint_attn_implementation="three_way", multiview_attention=SimpleNamespace(enabled=False), ) - assert model._uses_multiview_flex_kv() is True + assert model._uses_multiview_replay_kv() is True del model._teacher_forcing_kv_implementation_runtime model.config.teacher_forcing_kv_implementation = "singleview_threeway_kv" model.config.joint_attn_implementation = "multiview" - assert model._uses_multiview_flex_kv() is False + assert model._uses_multiview_replay_kv() is False @pytest.mark.L0 @pytest.mark.CPU -def test_multiview_flex_selection_rejects_non_replay_strategy() -> None: - """A Flex selector cannot silently fall through to a non-replay backend.""" +@pytest.mark.parametrize("implementation", ["multiview_flex_kv", "multiview_maskless_kv"]) +def test_multiview_replay_selection_rejects_non_replay_strategy(implementation: str) -> None: + """A multiview replay selector cannot silently fall through to a non-replay backend.""" from cosmos_framework.model.generator.omni_mot_model import OmniMoTModel from cosmos_framework.model.generator.omni_mot_causal_model import OmniMoTCausalModel @@ -468,7 +507,7 @@ def test_multiview_flex_selection_rejects_non_replay_strategy() -> None: torch.nn.Module.__init__(model) model.config = SimpleNamespace( causal_training_strategy="none", - teacher_forcing_kv_implementation="multiview_flex_kv", + teacher_forcing_kv_implementation=implementation, ) with patch.object(OmniMoTModel, "build_net") as base_build: @@ -480,8 +519,13 @@ def test_multiview_flex_selection_rejects_non_replay_strategy() -> None: @pytest.mark.L0 @pytest.mark.CPU -def test_teacher_forcing_dcm_build_routes_multiview_flex_selector() -> None: - """TF-dCM constructs the selected interactive Flex network instead of falling through.""" +@pytest.mark.parametrize("implementation", ["multiview_flex_kv", "multiview_maskless_kv"]) +@pytest.mark.parametrize("video_temporal_causal", [False, True]) +@pytest.mark.parametrize("causal_training_strategy", ["teacher_forcing", "teacher_forcing_dcm"]) +def test_teacher_forcing_build_keeps_student_and_teacher_backends_separate( + implementation: str, video_temporal_causal: bool, causal_training_strategy: str +) -> None: + """Construct the selected student network while preserving bidirectional teacher settings.""" from cosmos_framework.model.generator.omni_mot_model import OmniMoTModel from cosmos_framework.configs.base.defaults.replay_attention import TeacherForcingReplayPolicyConfig from cosmos_framework.model.generator import omni_mot_causal_model as causal_model_module @@ -491,26 +535,29 @@ def test_teacher_forcing_dcm_build_routes_multiview_flex_selector() -> None: replay_policy = TeacherForcingReplayPolicyConfig( control_visibility="causal", multiview_attention_scope="same_view", - decomposed_temporal_window_seconds=0.5, + decomposed_temporal_window_seconds=None if implementation == "multiview_maskless_kv" else 0.5, ) model = object.__new__(OmniMoTCausalModel) torch.nn.Module.__init__(model) model.config = SimpleNamespace( - causal_training_strategy="teacher_forcing_dcm", - teacher_forcing_kv_implementation="multiview_flex_kv", + causal_training_strategy=causal_training_strategy, + teacher_forcing_kv_implementation=implementation, teacher_forcing_replay_policy=replay_policy, teacher_forcing_frames_per_chunk=4, - video_temporal_causal=True, + video_temporal_causal=video_temporal_causal, joint_attn_implementation="three_way", multiview_attention=SimpleNamespace( enabled=False, + backend="flex_flash", mask=SimpleNamespace( attention_scope="all_views", + control_attends_sensor=False, decomposed_temporal_window_seconds=None, ), ), ) network = SimpleNamespace(config=SimpleNamespace()) + use_maskless = implementation == "multiview_maskless_kv" and video_temporal_causal def _fake_base_build( dtype: torch.dtype, @@ -523,20 +570,35 @@ def _fake_base_build( assert lora_enabled is None assert causal_model_module.omni_mot_model_module.Cosmos3VFMNetwork is InteractiveCosmos3VFMNetwork assert model.config.joint_attn_implementation == "multiview" + assert model.config.multiview_attention.backend == ("maskless" if use_maskless else "flex_flash") + assert model.config.multiview_attention.mask.control_attends_sensor is False + assert model.config.multiview_attention.mask.decomposed_temporal_window_seconds == ( + None if use_maskless else replay_policy.decomposed_temporal_window_seconds + ) + network.config.multiview_attention_config = model.config.multiview_attention return network with patch.object(OmniMoTModel, "build_net", side_effect=_fake_base_build): result = model.build_net(torch.bfloat16) assert result is network - assert network.config.video_temporal_causal is True - assert network.video_temporal_causal is True + assert network.config.video_temporal_causal is video_temporal_causal + assert network.video_temporal_causal is video_temporal_causal + assert network.teacher_forcing_maskless is use_maskless assert network.teacher_forcing_replay_policy is replay_policy assert network.teacher_forcing_frames_per_chunk == 4 assert model.config.joint_attn_implementation == "three_way" assert model.config.joint_attn_implementation != "multiview" assert model.config.multiview_attention.mask.attention_scope == "all_views" assert model.config.multiview_attention.mask.decomposed_temporal_window_seconds is None + assert model.config.multiview_attention.backend == "flex_flash" + assert model.config.multiview_attention.mask.control_attends_sensor is False + if use_maskless: + assert network.config.multiview_attention_config is not model.config.multiview_attention + assert network.config.multiview_attention_config.backend == "maskless" + assert network.config.multiview_attention_config.mask.control_attends_sensor is False + assert network.config.multiview_attention_config.mask.decomposed_temporal_window_seconds is None + assert network.teacher_forcing_replay_policy.decomposed_temporal_window_seconds is None class TestTeacherForcingTransferControlDropout: @@ -579,7 +641,7 @@ def test_target_only_replay_accepts_pack_after_control_dropout(self, dropout_rat model = self._make_model(dropout_rate) model.config.video_temporal_causal = True model.config.teacher_forcing_target_only_no_text_pass2 = True - model._uses_multiview_flex_kv.return_value = False + model._uses_multiview_replay_kv.return_value = False gen_data_clean = self._make_data() result = OmniMoTCausalModel._maybe_drop_teacher_forcing_transfer_control( model, gen_data_clean, {"dataset_name": ["video_transfer_4modality_480"]} @@ -1060,7 +1122,7 @@ def test_multiview_transfer_ar_mode_dispatches_to_specialized_iterator() -> None model = MagicMock() model.config.compile.enabled = False - model._uses_multiview_flex_kv.return_value = True + model._uses_multiview_replay_kv.return_value = True expected = {"vision": torch.zeros(1, 4, 2, 1, 1)} # [B,C,V*T,H,W] model._iter_samples_multiview_transfer_autoregressive.return_value = iter([expected]) data_batch = { @@ -1142,7 +1204,7 @@ def build_prefill(*_args: object, **_kwargs: object) -> SimpleNamespace: ) model = MagicMock() - model._uses_multiview_flex_kv.return_value = True + model._uses_multiview_replay_kv.return_value = True model.config = SimpleNamespace( action_gen=False, sound_gen=False, @@ -1712,7 +1774,7 @@ class TextTokensObserved(Exception): fps_vision=fps_vision, ) model = MagicMock() - model._uses_multiview_flex_kv.return_value = True + model._uses_multiview_replay_kv.return_value = True model.config.action_gen = False model.config.sound_gen = False model.input_caption_key = "caption" @@ -2121,7 +2183,7 @@ def _make_model_mock(self) -> MagicMock: m.config.fixed_step_sampler_config.t_list = [1.0, 0.5] m.config.fixed_step_sampler_config.sample_type = "ode" m.config.rectified_flow_inference_config.num_train_timesteps = 1000 - m._uses_multiview_flex_kv.return_value = False + m._uses_multiview_replay_kv.return_value = False return m def _make_gen_data(self, mode: str) -> SimpleNamespace: @@ -2541,7 +2603,7 @@ def test_multiview_transfer_with_compile_still_requires_eager_attention(self) -> """The multiview flex path keeps its eager-only guard.""" model = self._make_model_mock() model.config.compile.enabled = True - model._uses_multiview_flex_kv.return_value = True + model._uses_multiview_replay_kv.return_value = True with pytest.raises(ValueError, match="Multiview transfer AR requires eager attention"): self._run(model, "video_transfer") @@ -3797,13 +3859,13 @@ def test_forward_cuda_graph_scope_requires_the_static_path() -> None: model = _forward_graph_model_mock() model.config.compile.ar_post_saturation_mode = "cuda-graph" - model._uses_multiview_flex_kv.return_value = False + model._uses_multiview_replay_kv.return_value = False with pytest.raises(ValueError, match="cuda_graph_scope='forward'"): list(OmniMoTCausalModel.iter_samples_from_batch_autoregressive(model, {}, mode="text2video")) # CFG parallelism bypasses the whole-forward graphs and would recompile the static blocks per frame. model = _forward_graph_model_mock() - model._uses_multiview_flex_kv.return_value = False + model._uses_multiview_replay_kv.return_value = False model.parallel_dims = MagicMock(cfgp_enabled=True) with pytest.raises(ValueError, match="requires cfgp_size=1"): list(OmniMoTCausalModel.iter_samples_from_batch_autoregressive(model, {}, mode="text2video")) diff --git a/cosmos_framework/model/generator/omni_mot_model.py b/cosmos_framework/model/generator/omni_mot_model.py index 3d0c8e5f..1006e68a 100644 --- a/cosmos_framework/model/generator/omni_mot_model.py +++ b/cosmos_framework/model/generator/omni_mot_model.py @@ -82,6 +82,7 @@ uses_aux_loss_free_load_balancing, uses_ema_router_bias, ) +from cosmos_framework.model.generator.utils.rig_view_embedding import vision_view_ids from cosmos_framework.model.generator.utils.safetensors_loader import ( load_language_model as load_language_model_safetensors, ) @@ -408,6 +409,31 @@ def set_up_tokenizers(self) -> None: self.tokenizer_lidar_gen.reset_dtype() log.info(f"LiDAR tokenizer initialized: {type(self.tokenizer_lidar_gen).__name__}") + # 2c. Radar VAE for the BEV stream, a third sensor alongside camera and LiDAR and + # standing to the sequence exactly as LiDAR does. Its latent_ch is as wide as the + # LiDAR VAE's, but the network sizes the radar projections from radar_state_ch, so + # that is the field this VAE has to agree with. + self.tokenizer_radar_gen: VideoTokenizerInterface | None = None + if self.config.radar_tokenizer is not None and self.config.load_vision_tokenizer: + self.tokenizer_radar_gen = lazy_instantiate(self.config.radar_tokenizer) + if self.config.radar_state_ch is None: + raise ValueError("radar_state_ch must be set when radar_tokenizer is configured.") + if self.tokenizer_radar_gen.latent_ch != self.config.radar_state_ch: + raise ValueError( + f"Radar tokenizer latent_ch {self.tokenizer_radar_gen.latent_ch} != " + f"radar_state_ch {self.config.radar_state_ch}; the radar projection heads are sized " + "from radar_state_ch, so a mismatch would silently project the wrong width." + ) + if self.config.radar_fps is None or self.config.radar_fps <= 0: + raise ValueError( + f"radar_fps must be a positive cycle rate when radar_tokenizer is configured, got " + f"{self.config.radar_fps}. Without it the packer has no way to convert a radar " + "latent index into seconds and would silently place scans on the camera's rate." + ) + if hasattr(self.tokenizer_radar_gen, "reset_dtype"): + self.tokenizer_radar_gen.reset_dtype() + log.info(f"Radar tokenizer initialized: {type(self.tokenizer_radar_gen).__name__}") + # 3. Sound/audio tokenizer (optional) if self.config.sound_gen: assert self.config.sound_tokenizer is not None, "sound_tokenizer must be provided when sound_gen is True" @@ -478,9 +504,13 @@ def build_net( network_config = Cosmos3VFMNetworkConfig( vlm_config=language_model.config, latent_patch_size=self.config.diffusion_expert_config.patch_spatial, + lidar_patch_spatial_hw=self.config.diffusion_expert_config.lidar_patch_spatial_hw, + num_view_embeddings=self.config.diffusion_expert_config.num_view_embeddings, latent_downsample_factor=self.config.latent_downsample_factor, latent_channel_size=self.config.state_ch, lidar_latent_channel_size=self.config.lidar_state_ch, + radar_latent_channel_size=self.config.radar_state_ch, + radar_patch_spatial_hw=self.config.diffusion_expert_config.radar_patch_spatial_hw, max_latent_h=self.config.diffusion_expert_config.max_vae_latent_side_after_patchify, max_latent_w=self.config.diffusion_expert_config.max_vae_latent_side_after_patchify, max_latent_t=self.config.state_t, @@ -498,6 +528,7 @@ def build_net( joint_attn_implementation=self.config.joint_attn_implementation, multiview_attention_config=self.config.multiview_attention, timestep_scale=1.0 / float(num_train_timesteps) * self.config.diffusion_expert_config.timestep_range, + timestep_range=self.config.diffusion_expert_config.timestep_range, action_dim=self.config.max_action_dim, num_embodiment_domains=self.config.num_embodiment_domains, action_io_projector_type=self.config.action_io_projector_type, @@ -923,13 +954,15 @@ def _pack_input_sequence( plus three optional flags. """ assert self.tokenizer_vision_gen is not None - return pack_input_sequence( + packed = pack_input_sequence( sequence_plans=sequence_plans, input_text_indexes=input_text_indexes, gen_data_clean=gen_data_clean, input_timesteps=input_timesteps, special_tokens=self.llm_special_tokens, latent_patch_size=self.config.diffusion_expert_config.patch_spatial, + lidar_patch_spatial_hw=self.config.diffusion_expert_config.lidar_patch_spatial_hw, + radar_patch_spatial_hw=self.config.diffusion_expert_config.radar_patch_spatial_hw, skip_text_tokens=skip_text_tokens, include_end_of_generation_token=include_end_of_generation_token, unified_3d_mrope_reset_spatial_ids=self.config.diffusion_expert_config.unified_3d_mrope_reset_spatial_ids, @@ -941,11 +974,15 @@ def _pack_input_sequence( lidar_temporal_compression_factor=( self.tokenizer_lidar_gen.temporal_compression_factor if self.tokenizer_lidar_gen is not None else None ), + radar_temporal_compression_factor=( + self.tokenizer_radar_gen.temporal_compression_factor if self.tokenizer_radar_gen is not None else None + ), vision_temporal_position_mode=self.config.diffusion_expert_config.vision_temporal_position_mode, video_temporal_causal=self.config.video_temporal_causal, action_dim=self.config.max_action_dim, initial_mrope_temporal_offset=initial_mrope_temporal_offset, ) + return packed def _get_temporal_positions_vision( self, @@ -1206,8 +1243,8 @@ def _pack_training_payload( gen_data_clean_payload = { field.name: getattr(gen_data_clean, field.name) for field in dataclasses.fields(GenerationDataClean) } - # Raw pixels/audio/action/LiDAR are unused after tokenization; omit them from the CP cache. - for key in ("raw_state_vision", "raw_state_sound", "raw_state_action", "raw_state_lidar"): + # Raw pixels/audio/action/LiDAR/radar are unused after tokenization; omit them from the CP cache. + for key in ("raw_state_vision", "raw_state_sound", "raw_state_action", "raw_state_lidar", "raw_state_radar"): gen_data_clean_payload.pop(key) return { "input_text_indexes": input_text_indexes, @@ -1352,6 +1389,7 @@ def training_step( num_vision_latent_frames=num_vision_latent_frames, num_tokens=num_tokens_per_sample, iteration=iteration, + sequence_plans=sequence_plans, ) # [B, T_vis] each # Optional independent action schedule (sampled from rectified_flow_action with @@ -1497,6 +1535,22 @@ def training_step( else: timesteps_lidar, sigmas_lidar = (None, None) + # Radar borrows the sample's vision clock on the same terms LiDAR does, and is + # likewise expanded before the vision rebinding below. + if gen_data_clean.num_radar_items_per_sample is not None: + assert timesteps_vision.shape[1] == 1, ( + "Radar reuses the sample's vision timestep, which requires a single timestep per sample " + "(diffusion forcing is not supported for joint camera + radar batches)" + ) + timesteps_radar = _expand_per_sample_to_per_vision_item( + timesteps_vision, gen_data_clean.num_radar_items_per_sample + ) # [B_radar_items, 1] + sigmas_radar = _expand_per_sample_to_per_vision_item( + sigmas_vision, gen_data_clean.num_radar_items_per_sample + ) # [B_radar_items, 1] + else: + timesteps_radar, sigmas_radar = (None, None) + timesteps_vision = _expand_per_sample_to_per_vision_item( timesteps_vision, gen_data_clean.num_vision_items_per_sample ) # [B_items, T_vis] @@ -1514,6 +1568,7 @@ def training_step( sigmas_action=sigmas_action, sigmas_sound=sigmas_sound, sigmas_lidar=sigmas_lidar, + sigmas_radar=sigmas_radar, iteration=iteration, ) self._replace_clean_with_noised(packed_sequence, gen_data_noised) @@ -1537,6 +1592,7 @@ def training_step( timesteps_action=timesteps_action, timesteps_sound=timesteps_sound, timesteps_lidar=timesteps_lidar, + timesteps_radar=timesteps_radar, ) _vision_tokens = len(packed_sequence.vision.sequence_indexes) if packed_sequence.vision else 0 @@ -1709,6 +1765,7 @@ def _compute_losses( timesteps_action: torch.Tensor | None = None, timesteps_sound: torch.Tensor | None = None, timesteps_lidar: torch.Tensor | None = None, + timesteps_radar: torch.Tensor | None = None, ) -> tuple[torch.Tensor, dict[str, torch.Tensor]]: """Compute flow matching loss and auxiliary load balancing losses. @@ -1729,7 +1786,9 @@ def _compute_losses( rf_cfg = self.config.rectified_flow_training_config normalize_by_active = rf_cfg.normalize_loss_by_active - exclude_fully_conditioned_items = self.config.causal_training_strategy == "teacher_forcing" + exclude_fully_conditioned_items = rf_cfg.exclude_fully_conditioned_items + if exclude_fully_conditioned_items is None: + exclude_fully_conditioned_items = self.config.causal_training_strategy == "teacher_forcing" if self.config.vision_gen: # Only a batch that generates no camera stream, as the LiDAR-only recipe does, may # arrive with vision unpacked; for anything else that would silently sink the vision @@ -1743,8 +1802,9 @@ def _compute_losses( # loss over the network's zero-weighted vision predictions. That is what keeps # vae2llm and llm2vae in the backward graph on every rank, which FSDP requires. # Transfer teacher forcing flattens clean controls and the generated target into - # one vision-item list. Keep the per-item vector aligned for logging, but do not - # let zero-loss controls dilute the scalar training objective. + # one vision-item list. Keep the per-item vector aligned for logging. By default, + # zero-loss controls do not dilute TF's scalar objective; an explicit item-mean + # override can instead retain them to match the bidirectional initializer. fm_loss_vision, fm_loss_vision_per_instance = compute_flow_matching_loss( pred=out_net["preds_vision"], target=gen_data_noised.vt_target_vision, @@ -1773,8 +1833,8 @@ def _compute_losses( "LiDAR condition mask must be a list of tensors for loss computation" ) assert gen_data_noised.vt_target_lidar is not None, "LiDAR targets required when the batch has LiDAR" - # Match vision teacher forcing's target-item normalization. Clean HD-map - # controls must not dilute the LiDAR mean and change the sensor loss mix. + # Match vision's configured item normalization. Apply the same inclusion + # rule to clean HD-map controls so the sensor loss mix remains consistent. fm_loss_lidar, _ = compute_flow_matching_loss( pred=out_net["preds_lidar"], target=gen_data_noised.vt_target_lidar, @@ -1809,6 +1869,40 @@ def _compute_losses( total_loss += dummy_loss # [] losses_dict["flow_matching_loss_lidar"] = dummy_loss # [] + # Same condition the network builds its radar projections under, so the two cannot + # disagree about whether this run has a radar stream to supervise. + if self.config.vision_gen and self.config.radar_state_ch is not None: + if data_batch_packed.radar is not None: + assert isinstance(data_batch_packed.radar.condition_mask, list), ( + "Radar condition mask must be a list of tensors for loss computation" + ) + assert gen_data_noised.vt_target_radar is not None, "Radar targets required when the batch has radar" + fm_loss_radar, _ = compute_flow_matching_loss( + pred=out_net["preds_radar"], + target=gen_data_noised.vt_target_radar, + condition_mask=data_batch_packed.radar.condition_mask, + timesteps=timesteps_radar if timesteps_radar is not None else timesteps, + has_valid_tokens=has_noisy_tokens(data_batch_packed.radar), + rectified_flow=self.rectified_flow_video, + tensor_kwargs_fp32=self.tensor_kwargs_fp32, + normalize_by_active=normalize_by_active, + ) + radar_loss_scale = rf_cfg.radar_loss_scale if rf_cfg.radar_loss_scale is not None else rf_cfg.loss_scale + total_loss += fm_loss_radar * radar_loss_scale # [] + losses_dict["flow_matching_loss_radar"] = fm_loss_radar # [] + else: + # No radar data in this batch. Connect the network's dummy preds_radar to the + # loss so radar2llm / llm2radar stay in the backward graph, for the reason + # spelled out for LiDAR above. + preds_radar = out_net["preds_radar"] + assert preds_radar, ( + "preds_radar must carry the network's zero-weighted probe so radar2llm / " + "llm2radar stay in the backward graph on a batch with no radar" + ) + dummy_loss = 0.0 * sum(p.sum() for p in preds_radar) # [] + total_loss += dummy_loss # [] + losses_dict["flow_matching_loss_radar"] = dummy_loss # [] + if self.config.action_gen: if data_batch_packed.action is not None: assert isinstance(data_batch_packed.action.condition_mask, list), ( @@ -2022,6 +2116,7 @@ def _get_train_noise_level_vision( resolutions: list[str] | str | None = None, num_tokens: list[int] | None = None, iteration: int | None = None, + sequence_plans: list[SequencePlan] | None = None, ) -> tuple[torch.Tensor, torch.Tensor]: """ Sample the rectified flow interpolation coefficient (timesteps) and obtain the corresponding @@ -2038,6 +2133,9 @@ def _get_train_noise_level_vision( Can be a single string (applied to all samples) or a list of strings (one per sample). If None, defaults to self.config.resolution (can be used for other modalities). num_tokens: Number of tokens for each sample (before 2x2 merge). Needed for dynamic shift. + sequence_plans: Effective per-sample modality metadata after CP input selection. + Batches with LiDAR and no RGB use shift_lidar when configured. RGB and joint + RGB/LiDAR batches share the vision schedule. None retains the vision schedule. Returns: (timesteps, sigmas): Both [B,1] for TF/base, or [B,T_max] for diffusion_forcing. @@ -2055,6 +2153,14 @@ def _get_train_noise_level_vision( rf_config = self.config.rectified_flow_training_config shift_image = getattr(rf_config, "shift_image", None) shift_config = shift_image if is_image_batch and shift_image is not None else rf_config.shift + if ( + sequence_plans + and getattr(rf_config, "shift_lidar", None) is not None + and all(plan.has_lidar and not plan.has_vision for plan in sequence_plans) + ): + shift_config = rf_config.shift_lidar + if not isinstance(shift_config, int) or shift_config <= 0: + raise ValueError("shift_lidar must be a positive integer") if isinstance(shift_config, int): # Int-based shift: use directly for all samples shifts = torch.full((batch_size,), shift_config, dtype=torch.float32) @@ -2211,6 +2317,7 @@ def _add_noise_to_input( sigmas_action: torch.Tensor | None = None, sigmas_sound: torch.Tensor | None = None, sigmas_lidar: torch.Tensor | None = None, + sigmas_radar: torch.Tensor | None = None, iteration: int | None = None, ) -> GenerationDataNoised: """ @@ -2234,6 +2341,8 @@ def _add_noise_to_input( sigmas_lidar: ``[n_lidar_items, 1]`` sigma per LiDAR item, required when the batch carries LiDAR. Rows are the owning sample's vision sigma so both sensors are noised at the same point on the flow. + sigmas_radar: ``[n_radar_items, 1]`` sigma per radar item, required when the batch + carries radar. Read exactly as ``sigmas_lidar``. Returns: GenerationDataNoised: A dataclass containing the noise, noisy data (xt), and velocity field (vt). @@ -2317,6 +2426,31 @@ def _add_noise_to_input( xt_lidar = None vt_lidar = None + # Radar: same rectified flow as video, on its own BEV grid latents. + x0_radar = gen_data_clean.x0_tokens_radar # list of [C,T,H,W] + if x0_radar is not None and len(x0_radar) > 0: + assert packed_sequence.radar is not None, "Packed radar data required when radar tokens exist" + assert isinstance(packed_sequence.radar.condition_mask, list), ( + "Radar condition mask must be a list of tensors for noise scheduling" + ) + assert sigmas_radar is not None, "sigmas_radar required when radar tokens exist" + epsilon_radar = [ + torch.randn(x0_i.size(), generator=noise_gen, **self.tensor_kwargs_fp32) for x0_i in x0_radar + ] # list of [C,T,H,W] + context_parallel_broadcast_tensor_list(epsilon_radar, self.parallel_dims) + # sigmas_radar[i] is (1,) → view (1,1,1), broadcast against condition_mask [T,1,1]. + sigmas_radar_list = [ + sigmas_radar[i].view(-1, 1, 1) * (1.0 - packed_sequence.radar.condition_mask[i]) + for i in range(len(x0_radar)) + ] # list of [T,1,1] + xt_radar, vt_radar = self.rectified_flow_video.get_interpolation(epsilon_radar, x0_radar, sigmas_radar_list) + xt_radar = [xt_i.to(**self.tensor_kwargs) for xt_i in xt_radar] # list of [C,T,H,W] + else: + epsilon_radar = None + sigmas_radar_list = None + xt_radar = None + vt_radar = None + # Action (x0_tokens_action is already a dense list with no None entries). # Gate on action_gen: the dataset may emit action tensors for models that # don't consume them (e.g. camera dataset on a vision-only config), in @@ -2424,6 +2558,11 @@ def _add_noise_to_input( xt_tokens_lidar=xt_lidar, vt_target_lidar=vt_lidar, sigmas_lidar=sigmas_lidar_list, + # radar + epsilon_radar=epsilon_radar, + xt_tokens_radar=xt_radar, + vt_target_radar=vt_radar, + sigmas_radar=sigmas_radar_list, # action epsilon_action=epsilon_action, xt_tokens_action=xt_action, @@ -2450,6 +2589,8 @@ def _replace_clean_with_noised( packed_sequence.vision.tokens = gen_data_noised.xt_tokens_vision if packed_sequence.lidar is not None and gen_data_noised.xt_tokens_lidar is not None: packed_sequence.lidar.tokens = gen_data_noised.xt_tokens_lidar + if packed_sequence.radar is not None and gen_data_noised.xt_tokens_radar is not None: + packed_sequence.radar.tokens = gen_data_noised.xt_tokens_radar if packed_sequence.action is not None and gen_data_noised.xt_tokens_action is not None: action_all_conditioning = all( torch.all(condition_mask == 1).item() for condition_mask in packed_sequence.action.condition_mask @@ -2705,6 +2846,25 @@ def _prepare_inference_data( ) # [C,T,H,W] noise_lidar_list.append(cond_mask * x0_token + (1.0 - cond_mask) * pure_noise_i) # [C,T,H,W] + # 5c. Initialize radar noise the same way, from its own stream's condition masks. + noise_radar_list: list[torch.Tensor] | None = None + if gen_data_clean.x0_tokens_radar is not None: + assert packed_sequence.radar is not None, "Packed radar data required when the batch carries radar" + assert isinstance(packed_sequence.radar.condition_mask, list), "Radar condition mask required" + radar_counts = gen_data_clean.num_radar_items_per_sample or [1] * n_sample + seed_radar = [seed[sample_idx] for sample_idx, count in enumerate(radar_counts) for _ in range(count)] + noise_radar_list = [] + for i, (x0_token, cond_mask) in enumerate( + zip(gen_data_clean.x0_tokens_radar, packed_sequence.radar.condition_mask, strict=True) + ): + pure_noise_i = misc.arch_invariant_rand( + tuple(x0_token.shape), + self.tensor_kwargs_fp32["dtype"], + self.tensor_kwargs_fp32["device"], + seed_radar[i], + ) # [C,T,H,W] + noise_radar_list.append(cond_mask * x0_token + (1.0 - cond_mask) * pure_noise_i) # [C,T,H,W] + # 6. Initialize action noise if action_gen is True has_action = self.config.action_gen and any(plan.has_action for plan in sequence_plans) # Actions are denoising targets only when the packer marked action tokens @@ -2779,9 +2939,9 @@ def _prepare_inference_data( ) # [sound_channels,T_sound] noise_sound_list.append(noise_sound_i) - # 8. Concatenate vision, LiDAR, action, and sound noise per sample (flattened) - # Order: [vision | lidar (if present) | action (if present) | sound (if present)], - # matching the order the packer lays a sample out in. + # 8. Concatenate vision, LiDAR, radar, action, and sound noise per sample (flattened) + # Order: [vision | lidar (if present) | radar (if present) | action (if present) | + # sound (if present)], matching the order the packer lays a sample out in. # noise_action_list and noise_sound_list are dense (only modality-having samples), # so we use separate indexes. initial_noise: list[torch.Tensor] = [] @@ -2789,6 +2949,7 @@ def _prepare_inference_data( condition_mask: list[torch.Tensor] = [] idx_vision = 0 idx_lidar = 0 + idx_radar = 0 idx_action = 0 idx_sound = 0 @@ -2823,6 +2984,20 @@ def _prepare_inference_data( condition_mask_parts.append((mask_lidar * torch.ones_like(x0_lidar)).reshape(-1)) # [N_lidar] idx_lidar += 1 + if noise_radar_list is not None and sequence_plans[i].has_radar: + assert packed_sequence.radar is not None + assert gen_data_clean.x0_tokens_radar is not None + radar_counts = gen_data_clean.num_radar_items_per_sample + for _ in range(radar_counts[i] if radar_counts is not None else 1): + parts.append(noise_radar_list[idx_radar].reshape(-1)) + x0_radar = gen_data_clean.x0_tokens_radar[idx_radar] # [C,T,H,W] + mask_radar = packed_sequence.radar.condition_mask[idx_radar].to( # [T,1,1] + dtype=x0_radar.dtype, device=x0_radar.device + ) + condition_reference_parts.append(x0_radar.reshape(-1)) # [N_radar] + condition_mask_parts.append((mask_radar * torch.ones_like(x0_radar)).reshape(-1)) # [N_radar] + idx_radar += 1 + if noise_action_list is not None and sequence_plans[i].has_action: assert packed_sequence.action is not None assert packed_sequence.action.condition_mask is not None @@ -2976,6 +3151,7 @@ def _update_inference_pack_template( packed_sequence: PackedSequence, noise_x_vision: list[torch.Tensor], noise_x_lidar: list[torch.Tensor] | None, + noise_x_radar: list[torch.Tensor] | None, noise_x_action: list[torch.Tensor] | None, noise_x_sound: list[torch.Tensor] | None, timestep: torch.Tensor, @@ -2990,6 +3166,11 @@ def _update_inference_pack_template( packed_sequence.lidar.tokens = [x.to(**self.tensor_kwargs) for x in noise_x_lidar] # list[[C,T,H,W]] self._copy_timestep_to_template(packed_sequence.lidar.timesteps, timestep) + if noise_x_radar is not None: + assert packed_sequence.radar is not None, "packed_sequence.radar must exist when radar noise is present" + packed_sequence.radar.tokens = [x.to(**self.tensor_kwargs) for x in noise_x_radar] # list[[C,T,H,W]] + self._copy_timestep_to_template(packed_sequence.radar.timesteps, timestep) + if noise_x_action is not None: assert packed_sequence.action is not None, "packed_sequence.action must exist when action noise is present" packed_sequence.action.tokens = [x.to(**self.tensor_kwargs) for x in noise_x_action] # list[[T,D]] @@ -3091,16 +3272,21 @@ def _get_velocity( has_lidar = gen_data_clean.x0_tokens_lidar is not None num_lidar_items = gen_data_clean.num_lidar_items_per_sample - # Split flattened noise_x into vision, LiDAR, action, and sound parts per sample + has_radar = gen_data_clean.x0_tokens_radar is not None + num_radar_items = gen_data_clean.num_radar_items_per_sample + + # Split flattened noise_x into vision, LiDAR, radar, action, and sound parts per sample # Order must match _prepare_inference_data: - # [vision | lidar (if present) | action (if present) | sound (if present)] + # [vision | lidar (if present) | radar (if present) | action (if present) | sound (if present)] noise_x_vision: list[torch.Tensor] = [] noise_x_lidar: list[torch.Tensor] | None = [] if has_lidar else None + noise_x_radar: list[torch.Tensor] | None = [] if has_radar else None noise_x_action: list[torch.Tensor] | None = [] if has_noisy_actions else None noise_x_sound: list[torch.Tensor] | None = [] if has_sound else None vision_offset = 0 # tracks position in the flat x0_tokens_vision list lidar_offset = 0 # tracks position in the flat x0_tokens_lidar list + radar_offset = 0 # tracks position in the flat x0_tokens_radar list idx_action = 0 idx_sound = 0 for i in range(n_samples): @@ -3123,6 +3309,15 @@ def _get_velocity( offset += lidar_dim lidar_offset += n_lidar + if noise_x_radar is not None and sequence_plans[i].has_radar: + n_radar = num_radar_items[i] if num_radar_items is not None else 1 + for j in range(n_radar): + radar_shape = gen_data_clean.x0_tokens_radar[radar_offset + j].shape + radar_dim = int(torch.prod(torch.tensor(radar_shape))) + noise_x_radar.append(noise_x[i][offset : offset + radar_dim].reshape(radar_shape)) + offset += radar_dim + radar_offset += n_radar + if has_noisy_actions and noise_x_action is not None and sequence_plans[i].has_action: assert gen_data_clean.x0_tokens_action is not None action_shape = gen_data_clean.x0_tokens_action[idx_action].shape @@ -3170,11 +3365,17 @@ def _get_velocity( fps_sound=gen_data_clean.fps_sound if has_sound else None, num_vision_items_per_sample=num_items, num_views_per_vision_item=gen_data_clean.num_views_per_vision_item, + vision_view_ids=gen_data_clean.vision_view_ids, # LiDAR fields raw_state_lidar=gen_data_clean.raw_state_lidar, x0_tokens_lidar=noise_x_lidar if has_lidar else None, fps_lidar=gen_data_clean.fps_lidar, num_lidar_items_per_sample=num_lidar_items, + # Radar fields + raw_state_radar=gen_data_clean.raw_state_radar, + x0_tokens_radar=noise_x_radar if has_radar else None, + fps_radar=gen_data_clean.fps_radar, + num_radar_items_per_sample=num_radar_items, # Multi-control transfer: carry per-control weights so the packer can # populate vision_item_split_lens / control_weights on the packed # sequence. Without this, multi_control_two_way_attention never runs @@ -3200,6 +3401,10 @@ def _get_velocity( assert packed_sequence.lidar is not None, "packed_sequence.lidar must exist when the batch has LiDAR" packed_sequence.lidar.tokens = [x.to(**self.tensor_kwargs) for x in noise_x_lidar] # list[[C,T,H,W]] + if noise_x_radar is not None: + assert packed_sequence.radar is not None, "packed_sequence.radar must exist when the batch has radar" + packed_sequence.radar.tokens = [x.to(**self.tensor_kwargs) for x in noise_x_radar] # list[[C,T,H,W]] + if has_noisy_actions and noise_x_action is not None: assert packed_sequence.action is not None, "packed_sequence.action must exist when has_action is True" packed_sequence.action.tokens = [x.to(**self.tensor_kwargs) for x in noise_x_action] # list[[T,D]] @@ -3215,6 +3420,7 @@ def _get_velocity( packed_sequence_template, noise_x_vision, noise_x_lidar, + noise_x_radar, noise_x_action, noise_x_sound, timestep, @@ -3258,6 +3464,18 @@ def _get_velocity( else: velocity_lidar.append(torch.zeros_like(pred)) # [C,T,H,W] + # Handle radar velocity + velocity_radar: list[torch.Tensor] | None = None + if has_radar and packed_sequence.radar is not None and isinstance(packed_sequence.radar.condition_mask, list): + velocity_radar = [] + for pred, cond_mask in zip(out["preds_radar"], packed_sequence.radar.condition_mask, strict=True): + # pred: [C,T,H,W], cond_mask: [T,1,1] + noisy_mask = 1.0 - cond_mask + if noisy_mask.sum() > 0: + velocity_radar.append(pred * noisy_mask.to(dtype=pred.dtype, device=pred.device)) # [C,T,H,W] + else: + velocity_radar.append(torch.zeros_like(pred)) # [C,T,H,W] + # Handle action velocity velocity_action: list[torch.Tensor] | None = None if ( @@ -3305,11 +3523,12 @@ def _get_velocity( else: velocity_sound.append(torch.zeros_like(pred)) # [sound_channels,T_sound] - # Concatenate vision, LiDAR, action, and sound velocities per sample (flattened) - # Order must match _prepare_inference_data: [vision | lidar | action | sound] + # Concatenate vision, LiDAR, radar, action, and sound velocities per sample (flattened) + # Order must match _prepare_inference_data: [vision | lidar | radar | action | sound] velocity_output: list[torch.Tensor] = [] vis_offset = 0 lidar_out_offset = 0 + radar_out_offset = 0 idx_action = 0 idx_sound = 0 for i in range(n_samples): @@ -3325,6 +3544,11 @@ def _get_velocity( parts.append(velocity_lidar[lidar_out_offset].reshape(-1)) lidar_out_offset += 1 + if velocity_radar is not None and sequence_plans[i].has_radar: + for _ in range(num_radar_items[i] if num_radar_items is not None else 1): + parts.append(velocity_radar[radar_out_offset].reshape(-1)) + radar_out_offset += 1 + if velocity_action is not None and sequence_plans[i].has_action: parts.append(velocity_action[idx_action].reshape(-1)) idx_action += 1 @@ -3460,7 +3684,7 @@ def generate_samples_from_batch( n_sample: int | None = None, has_negative_prompt: bool = False, num_steps: int = 35, - shift: float = 5.0, + shift: float | None = None, sigma_max: float = 80.0, skip_text_tokens_for_cfg: bool = False, normalize_cfg: bool = False, @@ -3503,7 +3727,8 @@ def generate_samples_from_batch( n_sample (int | None): Number of samples to generate; defaults to batch size. has_negative_prompt (bool): If True, use negative prompt for unconditional branch. num_steps (int): Number of sampling steps for the diffusion process. - shift (float): Time shift parameter for the sampler. + shift (float | None): Explicit sampler shift. With shift_lidar configured, None uses + the LiDAR or RGB-resolution training shift; otherwise it retains the legacy value 5. sigma_max (float): Maximum sigma for the EDM sampler. skip_text_tokens_for_cfg (bool): If True, skip text tokens in unconditional branch. normalize_cfg (bool): If True, normalize the CFG output. @@ -3570,6 +3795,19 @@ def generate_samples_from_batch( ValueError: If the seed is a single integer. This is not supported anymore: `seed` must be a list of integers, one for each sample. """ + if shift is None: + shift = 5.0 + rf_config = getattr(self.config, "rectified_flow_training_config", None) + if getattr(rf_config, "shift_lidar", None) is not None: + if not self._has_vision_stream(data_batch): + shift = float(rf_config.shift_lidar) + elif isinstance(rf_config.shift, int): + shift = float(rf_config.shift) + else: + size = data_batch["image_size"][-1].reshape(-1, 4)[0] # [4] + resolution = get_vision_data_resolution((int(size[0]), int(size[1]))) + shift = float(rf_config.shift[resolution]) + if isinstance(seed, int): raise ValueError( "Single integer seed is not supported anymore: `seed` must be a list of integers, one for each sample." @@ -4131,21 +4369,25 @@ def x0_fn(noise_x: torch.Tensor, sigma: torch.Tensor) -> torch.Tensor: if _mixed_precision_runtime is not None: _mixed_precision_runtime.reset() - # Split flattened latents back into vision latents, LiDAR latents, external actions, - # and sound latents. Mirror the per-sample logic from _prepare_inference_data: - # Order: [vision | lidar (if present) | action (if present) | sound (if present)] - # lidar/action/sound lists are dense (only modality-having samples), so use separate indexes. + # Split flattened latents back into vision latents, LiDAR latents, radar latents, + # external actions, and sound latents. Mirror the per-sample logic from + # _prepare_inference_data: + # Order: [vision | lidar (if present) | radar (if present) | action (if present) | sound (if present)] + # lidar/radar/action/sound lists are dense (only modality-having samples), so use separate indexes. result_vision: list[torch.Tensor] = [] result_lidar: list[torch.Tensor] = [] + result_radar: list[torch.Tensor] = [] result_action: list[torch.Tensor] = [] result_sound: list[torch.Tensor] = [] action_processing_records = get_action_processing_records(data_batch) idx_vision = 0 idx_lidar = 0 + idx_radar = 0 idx_action = 0 idx_sound = 0 num_vision_items = gen_data_clean.num_vision_items_per_sample num_lidar_items = gen_data_clean.num_lidar_items_per_sample + num_radar_items = gen_data_clean.num_radar_items_per_sample for i in range(n_sample): offset = 0 @@ -4184,6 +4426,23 @@ def x0_fn(noise_x: torch.Tensor, sigma: torch.Tensor) -> torch.Tensor: offset += lidar_dim idx_lidar += n_lidar + # Extract radar if present, on the same rule: every item the plan generates. + if sequence_plans[i].has_radar: + assert gen_data_clean.x0_tokens_radar is not None + n_radar = num_radar_items[i] if num_radar_items is not None else 1 + for j in range(n_radar): + radar_shape = gen_data_clean.x0_tokens_radar[idx_radar + j].shape + radar_dim = int(torch.prod(torch.tensor(radar_shape))) + if is_item_generated( + sequence_plans[i].condition_frame_indexes_radar, + item_idx=j, + num_items=n_radar, + latent_t=int(radar_shape[2]), + ): + result_radar.append(latents[i][offset : offset + radar_dim].reshape(radar_shape)) + offset += radar_dim + idx_radar += n_radar + # Extract action if present if has_noisy_actions and sequence_plans[i].has_action: assert gen_data_clean.x0_tokens_action is not None @@ -4213,6 +4472,8 @@ def x0_fn(noise_x: torch.Tensor, sigma: torch.Tensor) -> torch.Tensor: result: dict[str, list[torch.Tensor]] = {"vision": result_vision} if result_lidar: result["lidar"] = result_lidar + if result_radar: + result["radar"] = result_radar if has_noisy_actions and result_action: result["action"] = result_action if self.config.sound_gen and len(result_sound) > 0: @@ -4352,6 +4613,20 @@ def _slice_gen_data_clean(self, gen_data_clean: GenerationDataClean, start: int, subset_num_lidar_items = None if num_lidar_items is None else num_lidar_items[start:limit] fps_lidar = gen_data_clean.fps_lidar[start:limit] if gen_data_clean.fps_lidar is not None else None + # The radar stream is grouped by sample the same way, so it takes its own slice. + num_radar_items = gen_data_clean.num_radar_items_per_sample + if gen_data_clean.x0_tokens_radar is None: + subset_x0_radar = None + subset_raw_radar = None + subset_num_radar_items = None + else: + radar_counts = num_radar_items if num_radar_items is not None else [1] * gen_data_clean.batch_size + radar_slice = slice(sum(radar_counts[:start]), sum(radar_counts[:limit])) + subset_x0_radar = gen_data_clean.x0_tokens_radar[radar_slice] + subset_raw_radar = gen_data_clean.raw_state_radar[radar_slice] if gen_data_clean.raw_state_radar else None + subset_num_radar_items = None if num_radar_items is None else num_radar_items[start:limit] + fps_radar = gen_data_clean.fps_radar[start:limit] if gen_data_clean.fps_radar is not None else None + if has_action: subset_raw_action = ( gen_data_clean.raw_state_action[start:limit] if gen_data_clean.raw_state_action else None @@ -4398,10 +4673,19 @@ def _slice_gen_data_clean(self, gen_data_clean: GenerationDataClean, start: int, action_valid_mask=action_valid_mask, num_vision_items_per_sample=subset_num_items, num_views_per_vision_item=subset_num_views_per_vision_item, + vision_view_ids=( + gen_data_clean.vision_view_ids[vision_item_slice] + if gen_data_clean.vision_view_ids is not None + else None + ), # list[[V]] raw_state_lidar=subset_raw_lidar, x0_tokens_lidar=subset_x0_lidar, fps_lidar=fps_lidar, num_lidar_items_per_sample=subset_num_lidar_items, + raw_state_radar=subset_raw_radar, + x0_tokens_radar=subset_x0_radar, + fps_radar=fps_radar, + num_radar_items_per_sample=subset_num_radar_items, ) @torch.no_grad() @@ -4527,6 +4811,15 @@ def _require_lidar_tokenizer(self) -> VideoTokenizerInterface: ) return self.tokenizer_lidar_gen + def _require_radar_tokenizer(self) -> VideoTokenizerInterface: + """Return the radar VAE, or say which config knob is missing.""" + if self.tokenizer_radar_gen is None: + raise ValueError( + "This batch carries a radar stream, but no radar tokenizer is loaded. " + "Set radar_tokenizer and radar_state_ch on the model config." + ) + return self.tokenizer_radar_gen + def _normalize_uint8_vision_item(self, state: torch.Tensor) -> torch.Tensor: """Move one uint8 vision item to the model device as fp32 and normalize it to ``[-1,1]``.""" return normalize_uint8_item(state, self.tensor_kwargs_fp32) @@ -4945,6 +5238,9 @@ def get_data_and_condition( # LiDAR range clips: their own VAE, so they never travel among the vision items. raw_state_lidar, x0_tokens_lidar, num_lidar_items_per_sample = self._encode_lidar_stream(data_batch, batch_size) + # Radar BEV clips: likewise their own VAE, so likewise outside the vision items. + raw_state_radar, x0_tokens_radar, num_radar_items_per_sample = self._encode_radar_stream(data_batch, batch_size) + output_raw_state_vision = raw_state_vision if retain_raw_state_vision and num_views_per_vision_item is not None: # Camera pixels arrive as uint8 levels and need the [-1,1] map. An item that is @@ -4995,6 +5291,13 @@ def get_data_and_condition( **self.tensor_kwargs ) + # Radar cycle rate for mRoPE, read from the config for the reason the LiDAR rate is. + fps_radar = None + if x0_tokens_radar is not None: + fps_radar = torch.full((batch_size,), float(self.config.radar_fps), dtype=torch.float32).to( + **self.tensor_kwargs + ) + # Sound FPS for RoPE alignment (constant, from config) if x0_tokens_sound is not None: sound_batch_size = len(x0_tokens_sound) @@ -5006,6 +5309,15 @@ def get_data_and_condition( else: fps_sound = None + physical_view_ids = None + if self.config.diffusion_expert_config.num_view_embeddings: + physical_view_ids = vision_view_ids( + data_batch, + batch_size=batch_size, + item_counts=num_vision_items_per_sample, + views_per_item=num_views_per_vision_item, + num_embeddings=self.config.diffusion_expert_config.num_view_embeddings, + ) # list[[V]] control_weights: list[list[float]] | None = data_batch.get("control_weights", None) return GenerationDataClean( batch_size=batch_size, @@ -5024,10 +5336,15 @@ def get_data_and_condition( action_family=action_family, num_vision_items_per_sample=num_vision_items_per_sample, num_views_per_vision_item=num_views_per_vision_item, + vision_view_ids=physical_view_ids, raw_state_lidar=raw_state_lidar, x0_tokens_lidar=x0_tokens_lidar, fps_lidar=fps_lidar, num_lidar_items_per_sample=num_lidar_items_per_sample, + raw_state_radar=raw_state_radar, + x0_tokens_radar=x0_tokens_radar, + fps_radar=fps_radar, + num_radar_items_per_sample=num_radar_items_per_sample, raw_action_dim=raw_action_dim, action_valid_mask=action_valid_mask, control_weights=control_weights, @@ -5064,6 +5381,14 @@ def _get_lidar_only_data_and_condition(self, data_batch: dict[str, Any]) -> Gene fps_lidar = torch.full((batch_size,), float(self.config.lidar_fps), dtype=torch.float32).to( **self.tensor_kwargs ) + # A camera-less AV batch may still carry radar beside its LiDAR, so the third stream is + # encoded here too rather than silently dropped. + raw_state_radar, x0_tokens_radar, num_radar_items = self._encode_radar_stream(data_batch, batch_size) + fps_radar = None + if x0_tokens_radar is not None: + fps_radar = torch.full((batch_size,), float(self.config.radar_fps), dtype=torch.float32).to( + **self.tensor_kwargs + ) return GenerationDataClean( batch_size=batch_size, is_image_batch=False, @@ -5074,6 +5399,10 @@ def _get_lidar_only_data_and_condition(self, data_batch: dict[str, Any]) -> Gene x0_tokens_lidar=x0_tokens_lidar, fps_lidar=fps_lidar, num_lidar_items_per_sample=num_lidar_items, + raw_state_radar=raw_state_radar, + x0_tokens_radar=x0_tokens_radar, + fps_radar=fps_radar, + num_radar_items_per_sample=num_radar_items, control_weights=data_batch.get("control_weights", None), ) @@ -5142,6 +5471,82 @@ def _encode_lidar_stream( x0_tokens_lidar = [self.encode_lidar(state).contiguous().float() for state in raw_state_lidar] return raw_state_lidar, x0_tokens_lidar, num_lidar_items_per_sample + def _encode_radar_stream( + self, + data_batch: dict[str, Any], + batch_size: int, + ) -> tuple[list[torch.Tensor] | None, list[torch.Tensor] | None, list[int] | None]: + """Encode the batch's radar BEV clips into x0 latent tokens. + + ``data_batch["radar"]`` arrives from the loader in the same per-sample shape the LiDAR + key does — one entry per sample, each a list of tokenizer-native clips — and is + flattened here the same way. The clips are on the native polar BEV grid and the + tokenizer owns their normalization. The per-sample counts are written back to the + batch for the reasons they are for LiDAR. + + Returns: + The raw BEV clips, their latents, and the per-sample item counts; all + ``None`` when the batch carries no radar. + """ + raw = data_batch.get("radar", None) + if raw is None: + return None, None, None + + num_radar_items_per_sample = data_batch.get("num_radar_items_per_sample", None) + if num_radar_items_per_sample is None: + # Per-sample form: group the items and record the counts. + if len(raw) != batch_size: + raise ValueError(f"The radar stream needs one entry per sample: got {len(raw)} for {batch_size}.") + grouped = [list(entry) if isinstance(entry, (list, tuple)) else [entry] for entry in raw] + num_radar_items_per_sample = [len(items) for items in grouped] + flat_items = [item for items in grouped for item in items] + else: + # Flat form: already grouped on an earlier pass over this batch. + num_radar_items_per_sample = [int(count) for count in num_radar_items_per_sample] + flat_items = list(raw) + if len(num_radar_items_per_sample) != batch_size: + raise ValueError( + "num_radar_items_per_sample must have one count per sample: " + f"got {len(num_radar_items_per_sample)} counts for batch size {batch_size}." + ) + if len(flat_items) != sum(num_radar_items_per_sample): + raise ValueError( + "Radar items must match num_radar_items_per_sample: " + f"got {len(flat_items)} items for {sum(num_radar_items_per_sample)} expected." + ) + + raw_state_radar: list[torch.Tensor] = [] + for item in flat_items: + if not isinstance(item, torch.Tensor): + raise TypeError(f"Radar items must be tensors, got {type(item).__name__}.") + if item.dim() == 4: # [C,T,H,W] + item = item.unsqueeze(0) # [1,C,T,H,W] + elif item.dim() != 5: + raise ValueError(f"Radar items must have shape [C,T,H,W] or [B,C,T,H,W], got {tuple(item.shape)}.") + if not torch.is_floating_point(item): + raise TypeError(f"Radar BEV maps must arrive as floating-point clips, got {item.dtype}.") + raw_state_radar.append(item.to(**self.tensor_kwargs_fp32)) # [1,C,T,H,W] + + data_batch["radar"] = raw_state_radar + data_batch["num_radar_items_per_sample"] = num_radar_items_per_sample + + x0_tokens_radar = self._encode_radar_items(raw_state_radar) + return raw_state_radar, x0_tokens_radar, num_radar_items_per_sample + + def _encode_radar_items(self, raw_state_radar: list[torch.Tensor]) -> list[torch.Tensor]: + """Encode each radar clip on its own batch dim. + + Transfer samples carry a map-control clip and a radar-target clip of + equal ``[1,C,T,H,W]``. Stacking them (B=2) with a T=32 stream window + made the V1 VAE's circular pad at the 16-head stage allocate ~17 GiB, + which OOMs an 8-node run once the camera tokens already sit on the GPU. + One clip at a time plus a 16-frame window fits; the extra VAE pass is + ~1.4s on a ~24s step. + """ + if not raw_state_radar: + return [] + return [self.encode_radar(state).contiguous().float() for state in raw_state_radar] # list[[1,Cz,Tz,Hz,Wz]] + def _normalize_video_databatch_inplace(self, data_batch: dict[str, torch.Tensor]) -> None: """ Normalizes video data in-place on a CUDA device to reduce data loading overhead. @@ -5861,6 +6266,7 @@ def denoise( dict containing: - "preds_vision": list[Tensor[C,T,H,W]], one per sample. - "preds_lidar": Velocity prediction for the LiDAR stream (if the network carries it). + - "preds_radar": Velocity prediction for the radar stream (if the network carries it). - "preds_action": Velocity prediction for action modality (if action_gen enabled). - "preds_sound": Velocity prediction for sound modality (if sound_gen enabled). - "lbl_metadata_und": Load balancing metadata for understanding pathway (if present). @@ -5877,6 +6283,8 @@ def denoise( output_dict["preds_vision"] = out_net["preds_vision"] if "preds_lidar" in out_net: output_dict["preds_lidar"] = out_net["preds_lidar"] + if "preds_radar" in out_net: + output_dict["preds_radar"] = out_net["preds_radar"] if self.config.action_gen and "preds_action" in out_net: output_dict["preds_action"] = out_net["preds_action"] if self.config.sound_gen and "preds_sound" in out_net: @@ -6678,6 +7086,21 @@ def decode_lidar(self, latent: torch.Tensor) -> torch.Tensor: """ return self._require_lidar_tokenizer().decode(latent) + @torch.no_grad() + def encode_radar(self, state: torch.Tensor) -> torch.Tensor: + """Encode a radar BEV clip with the radar VAE.""" + return self._require_radar_tokenizer().encode(state) + + @torch.no_grad() + def decode_radar(self, latent: torch.Tensor) -> torch.Tensor: + """Decode radar latents to BEV pixels with the radar VAE. + + The stream keeps its own VAE's channel count end to end, so a latent arrives here + at the width its decoder expects, on the native polar grid the tokenizer was + trained on. + """ + return self._require_radar_tokenizer().decode(latent) + @torch.no_grad() def encode_sound(self, waveform: torch.Tensor) -> torch.Tensor: """Encode audio waveform into latent tokens. diff --git a/cosmos_framework/model/generator/utils/data_and_condition.py b/cosmos_framework/model/generator/utils/data_and_condition.py index ce9b47ae..764d8507 100644 --- a/cosmos_framework/model/generator/utils/data_and_condition.py +++ b/cosmos_framework/model/generator/utils/data_and_condition.py @@ -16,7 +16,7 @@ class GenerationDataClean: """ Container for tokenized states and conditioning info (clean states) - for the multi-modal (vision, lidar, sound, action) MoT training. + for the multi-modal (vision, lidar, radar, sound, action) MoT training. Used for the VFM generation model. """ @@ -39,6 +39,13 @@ class GenerationDataClean: # camera clips along the latent temporal axis (camera-major), so latent_t is # num_views * frames_per_view. None when per-camera VAE encoding is disabled. num_views_per_vision_item: list[int] | None = None + # Physical camera IDs: one [V] integer tensor per flattened x0_tokens_vision item, + # in the same camera-major order as its latents. IDs come from view_indices_selection + # and are repeated for each control and target without renumbering selected cameras: + # camera 8 with one control and one target gives [tensor([8]), tensor([8])]. + # None when rig embeddings are disabled or no RGB is present. LiDAR uses the final + # embedding row separately and never appears in these tensors. + vision_view_ids: list[torch.Tensor] | None = None # LiDAR (list of per-item range-view latents, flattened over samples the way # x0_tokens_vision is). A range clip is its own modality with its own VAE and its own @@ -48,6 +55,14 @@ class GenerationDataClean: fps_lidar: torch.Tensor | None = None num_lidar_items_per_sample: list[int] | None = None + # Radar (per-item BEV latents, flattened over samples exactly as the LiDAR items are). + # Radar is a third sensor stream with its own VAE and its own ~20 Hz cycle rate, so like + # LiDAR it never appears among the vision items. + raw_state_radar: list[torch.Tensor] | None = None + x0_tokens_radar: list[torch.Tensor] | None = None + fps_radar: torch.Tensor | None = None + num_radar_items_per_sample: list[int] | None = None + # Audio (Sound) raw_state_sound: torch.Tensor | None = None x0_tokens_sound: torch.Tensor | None = None @@ -72,7 +87,7 @@ class GenerationDataClean: class GenerationDataNoised: """Container for states after noise addition, along with other helper attributes for the flow-matching (gt velocity and noise) - for the multi-modal (vision, lidar, sound, action) MoT training. + for the multi-modal (vision, lidar, radar, sound, action) MoT training. Used for the VFM generation model. """ @@ -89,6 +104,12 @@ class GenerationDataNoised: vt_target_lidar: torch.Tensor | None = None sigmas_lidar: torch.Tensor | None = None + # Radar + epsilon_radar: torch.Tensor | None = None + xt_tokens_radar: torch.Tensor | None = None + vt_target_radar: torch.Tensor | None = None + sigmas_radar: torch.Tensor | None = None + # Audio (Sound) epsilon_sound: torch.Tensor | None = None xt_tokens_sound: torch.Tensor | None = None diff --git a/cosmos_framework/model/generator/utils/rig_view_embedding.py b/cosmos_framework/model/generator/utils/rig_view_embedding.py new file mode 100644 index 00000000..896bc0b9 --- /dev/null +++ b/cosmos_framework/model/generator/utils/rig_view_embedding.py @@ -0,0 +1,51 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +"""Physical rig identity for camera-major control and target token grids.""" + +import math +from typing import Any + +import torch + + +def vision_view_ids( + data_batch: dict[str, Any], + *, + batch_size: int, + item_counts: list[int] | None, + views_per_item: list[int] | None, + num_embeddings: int, +) -> list[torch.Tensor]: # returns one [V] ID tensor per vision item + """Repeat each sample's canonical camera IDs for its controls and target.""" + raw_ids = data_batch.get("view_indices_selection") + if raw_ids is None or len(raw_ids) != batch_size: + raise ValueError("Rig view embeddings require view_indices_selection for every RGB sample") + counts = item_counts if item_counts is not None else [1] * batch_size + result: list[torch.Tensor] = [] + for sample_ids, count in zip(raw_ids, counts, strict=True): + ids = torch.as_tensor(sample_ids, dtype=torch.long).reshape(-1) # [V] + if ids.numel() == 0 or bool(((ids < 0) | (ids >= num_embeddings - 1)).any()): + raise ValueError("RGB view IDs must be canonical camera IDs; the final embedding is reserved for LiDAR") + for _ in range(count): + if views_per_item is not None and ids.numel() != views_per_item[len(result)]: + raise ValueError("Physical camera IDs do not match the encoded camera count") + result.append(ids) + return result + + +def add_view_embeddings( + tokens: torch.Tensor, # [N,D] + token_shapes: list[tuple[int, int, int]], + view_ids: list[torch.Tensor], # one [V] tensor per item + embedding: torch.nn.Embedding, +) -> torch.Tensor: # [N,D] + """Apply IDs in camera-major order to every spatial and temporal token.""" + offsets: list[torch.Tensor] = [] + for shape, ids in zip(token_shapes, view_ids, strict=True): + count = math.prod(shape) + if shape[0] % ids.numel() != 0: + raise ValueError("Camera-major temporal extent must be divisible by the number of view IDs") + per_view = embedding(ids) # [V,D] + offsets.append(per_view.repeat_interleave(count // ids.numel(), dim=0)) # [N_item,D] + return tokens + torch.cat(offsets, dim=0).to(tokens.dtype) # [N,D] diff --git a/cosmos_framework/model/generator/utils/rig_view_embedding_test.py b/cosmos_framework/model/generator/utils/rig_view_embedding_test.py new file mode 100644 index 00000000..4f65f178 --- /dev/null +++ b/cosmos_framework/model/generator/utils/rig_view_embedding_test.py @@ -0,0 +1,51 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +"""Physical view IDs survive controls, packing and unconditional text paths.""" + +import pytest +import torch + +from cosmos_framework.model.generator.utils.rig_view_embedding import add_view_embeddings, vision_view_ids + +pytestmark = [pytest.mark.L0, pytest.mark.CPU] + + +def test_ids_repeat_per_item_without_renumbering() -> None: + ids = vision_view_ids( + {"view_indices_selection": [[8], [2, 5]], "ai_caption": ["", ""]}, + batch_size=2, + item_counts=[2, 3], + views_per_item=[1, 1, 2, 2, 2], + num_embeddings=12, + ) + assert [item.tolist() for item in ids] == [[8], [8], [2, 5], [2, 5], [2, 5]] + + +@pytest.mark.parametrize("raw_ids", [None, [[11]], [[-1]], [[]]]) +def test_invalid_camera_ids_fail_before_forward(raw_ids: list[list[int]] | None) -> None: + with pytest.raises(ValueError): + vision_view_ids( + {"view_indices_selection": raw_ids}, batch_size=1, item_counts=[2], views_per_item=[1, 1], num_embeddings=12 + ) + + +def test_camera_major_offsets_and_gradients() -> None: + embedding = torch.nn.Embedding(12, 3) + with torch.no_grad(): + embedding.weight.copy_(torch.arange(12).reshape(12, 1).expand(12, 3)) # [12,3] + tokens = torch.zeros(12, 3) # [N,D] + ids = [torch.tensor([8]), torch.tensor([2, 5])] # per-item [V] + actual = add_view_embeddings(tokens, [(2, 1, 2), (4, 1, 2)], ids, embedding) # [N,D] + assert actual[:, 0].tolist() == [8] * 4 + [2] * 4 + [5] * 4 + actual.sum().backward() + assert embedding.weight.grad is not None + assert embedding.weight.grad[:, 0].tolist() == [0, 0, 4, 0, 0, 4, 0, 0, 4, 0, 0, 0] + + +def test_zero_initialization_preserves_ga_tokens() -> None: + embedding = torch.nn.Embedding(12, 3) + torch.nn.init.zeros_(embedding.weight) # [12,D] + tokens = torch.randn(4, 3) # [N,D] + actual = add_view_embeddings(tokens, [(2, 1, 2)], [torch.tensor([10])], embedding) # [N,D] + torch.testing.assert_close(actual, tokens, rtol=0, atol=0) diff --git a/cosmos_framework/utils/generator/data_utils.py b/cosmos_framework/utils/generator/data_utils.py index 14e79cec..ddaeb29d 100644 --- a/cosmos_framework/utils/generator/data_utils.py +++ b/cosmos_framework/utils/generator/data_utils.py @@ -96,7 +96,10 @@ def slice_data_batch( The ``lidar`` field is flattened the same way and follows ``num_lidar_items_per_sample``, since a sample's range clips need not be as - many as its camera clips. + many as its camera clips. ``radar`` follows ``num_radar_items_per_sample`` + for the same reason: a joint camera + radar sample flattens the HD-map + control and radar target into two items that must stay together when the + sample callback draws the first sample. Args: data_batch: The data batch to slice. @@ -121,6 +124,8 @@ def flat_range(counts: Any) -> tuple[int, int]: # The LiDAR stream is flattened the same way, and counts its own items. num_lidar_items = data_batch.get("num_lidar_items_per_sample") lidar_start, lidar_limit = flat_range(num_lidar_items) if num_lidar_items is not None else (start, limit) + num_radar_items = data_batch.get("num_radar_items_per_sample") + radar_start, radar_limit = flat_range(num_radar_items) if num_radar_items is not None else (start, limit) multi_item_fields = set(multi_item_fields) @@ -130,6 +135,8 @@ def flat_range(counts: Any) -> tuple[int, int]: s, e = flat_start, flat_limit elif key == "lidar" and num_lidar_items is not None: s, e = lidar_start, lidar_limit + elif key == "radar" and num_radar_items is not None: + s, e = radar_start, radar_limit else: s, e = start, limit if isinstance(value, torch.Tensor): diff --git a/cosmos_framework/utils/generator/spatial_patch.py b/cosmos_framework/utils/generator/spatial_patch.py new file mode 100644 index 00000000..6df29695 --- /dev/null +++ b/cosmos_framework/utils/generator/spatial_patch.py @@ -0,0 +1,12 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +"""Normalize square or rectangular spatial patches at config boundaries.""" + + +def normalize_spatial_patch_hw(patch_hw: int | tuple[int, int]) -> tuple[int, int]: + """Accept a square side or (height, width), including Hydra's sequence values.""" + height, width = (patch_hw, patch_hw) if isinstance(patch_hw, int) else patch_hw + if any(not isinstance(side, int) or isinstance(side, bool) or side <= 0 for side in (height, width)): + raise ValueError("Spatial patch height and width must be positive integers") + return height, width diff --git a/cosmos_framework/utils/generator/video_preprocess.py b/cosmos_framework/utils/generator/video_preprocess.py index 1b633e27..9c0d56c7 100644 --- a/cosmos_framework/utils/generator/video_preprocess.py +++ b/cosmos_framework/utils/generator/video_preprocess.py @@ -6,18 +6,22 @@ from PIL import Image -def tensor_to_pil_images(video_tensor: torch.Tensor) -> list[Image.Image]: +def tensor_to_pil_images(video_tensor: torch.Tensor, *, channels_first: bool | None = None) -> list[Image.Image]: """Convert a video tensor of shape (C, T, H, W) or (T, C, H, W) into a list of PIL images. Args: video_tensor: Video tensor with shape (C, T, H, W) or (T, C, H, W). + channels_first: Specify the layout when it cannot be inferred, such as for three-frame RGB videos. + If omitted, the helper retains its original layout inference. Returns: One PIL image per frame. """ + if channels_first is None: + channels_first = video_tensor.shape[0] == 3 and video_tensor.shape[1] > 3 # (C, T, H, W) -> (T, C, H, W) - if video_tensor.shape[0] == 3 and video_tensor.shape[1] > 3: - video_tensor = video_tensor.permute(1, 0, 2, 3) + if channels_first: + video_tensor = video_tensor.permute(1, 0, 2, 3) # [T,C,H,W] # (T, C, H, W) -> (T, H, W, C) and detach to CPU numpy. video_np = video_tensor.permute(0, 2, 3, 1).cpu().numpy()