From 370da8ea025a7981bd374c31d41b6072faec2c2d Mon Sep 17 00:00:00 2001 From: Liang Feng Date: Wed, 23 Sep 2026 12:57:41 +0000 Subject: [PATCH] Release 2026-09-23 from i4 8962d9fffb7304ced8f87e3c43eb8699331f1999 _source_commit: 8962d9fffb7304ced8f87e3c43eb8699331f1999-dirty _dest_commit: cf5d68c00d97ccd2480a2320ed652b92dec63102 --- .file_mapping.json | 9 ++- .../configs/base/defaults/model_config.py | 4 ++ .../augmentors/first_caption_paragraph.py | 34 +++++++++ .../first_caption_paragraph_test.py | 35 ++++++++++ .../data/generator/joint_dataloader.py | 12 +++- .../generator/multiview/caption_chunks.py | 7 ++ .../generator/multiview/caption_format.py | 14 ++++ .../generator/sequence_packing/packers.py | 8 ++- .../generator/sequence_packing/sequence.py | 14 ++-- .../generator/mot/cosmos3_vfm_network.py | 69 ++++++++++++------- .../model/generator/omni_mot_model.py | 2 + .../utils/generator/spatial_patch.py | 12 ++++ 12 files changed, 181 insertions(+), 39 deletions(-) create mode 100644 cosmos_framework/data/generator/augmentors/first_caption_paragraph.py create mode 100644 cosmos_framework/data/generator/augmentors/first_caption_paragraph_test.py create mode 100644 cosmos_framework/utils/generator/spatial_patch.py diff --git a/.file_mapping.json b/.file_mapping.json index b96ddff3..9a571805 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": "8962d9fffb7304ced8f87e3c43eb8699331f1999-dirty", + "_dest_commit": "cf5d68c00d97ccd2480a2320ed652b92dec63102", + "_generated_at": "2026-09-23T12:57:25Z", "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", @@ -569,6 +571,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", diff --git a/cosmos_framework/configs/base/defaults/model_config.py b/cosmos_framework/configs/base/defaults/model_config.py index 777f8a59..9efd9cd7 100644 --- a/cosmos_framework/configs/base/defaults/model_config.py +++ b/cosmos_framework/configs/base/defaults/model_config.py @@ -47,6 +47,10 @@ class DiffusionExpertConfig: enable_sound_modality_embedding: bool = True patch_spatial: int = 2 + # None preserves the shared camera/LiDAR 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 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). ) 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/joint_dataloader.py b/cosmos_framework/data/generator/joint_dataloader.py index 0f9ca647..389698e6 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", @@ -461,6 +462,7 @@ 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, ) -> None: """ Initialize the JointDataLoader with multiple datasets. @@ -479,6 +481,7 @@ 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. 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 @@ -543,6 +546,9 @@ def __init__( f"lidar_temporal_compression_factor must be positive, got {self.lidar_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.max_sequence_length = max_sequence_length self.max_samples_per_batch = max_samples_per_batch self.sound_latent_fps = sound_latent_fps @@ -668,8 +674,8 @@ 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 @@ -1138,6 +1144,7 @@ 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, ) -> 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}.") @@ -1189,6 +1196,7 @@ 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, ) 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/sequence_packing/packers.py b/cosmos_framework/data/generator/sequence_packing/packers.py index ab985abb..854cd9a0 100644 --- a/cosmos_framework/data/generator/sequence_packing/packers.py +++ b/cosmos_framework/data/generator/sequence_packing/packers.py @@ -169,6 +169,7 @@ 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, + lidar_patch_spatial_hw: int | tuple[int, int] | None = None, ) -> PackedSequence: """ Pack a sequence of input strings and VAE latents into a packed tensor format. @@ -189,7 +190,8 @@ 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. 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 @@ -702,7 +704,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, diff --git a/cosmos_framework/data/generator/sequence_packing/sequence.py b/cosmos_framework/data/generator/sequence_packing/sequence.py index be297c14..01e71729 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 @@ -412,7 +413,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 +432,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. @@ -465,7 +466,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 +486,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 diff --git a/cosmos_framework/model/generator/mot/cosmos3_vfm_network.py b/cosmos_framework/model/generator/mot/cosmos3_vfm_network.py index fbc4be56..90de5446 100644 --- a/cosmos_framework/model/generator/mot/cosmos3_vfm_network.py +++ b/cosmos_framework/model/generator/mot/cosmos3_vfm_network.py @@ -51,6 +51,7 @@ get_causal_seq, get_full_only_seq, ) +from cosmos_framework.utils.generator.spatial_patch import normalize_spatial_patch_hw class Cosmos3VFMNetworkConfig(PretrainedConfig): @@ -93,12 +94,16 @@ 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, **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.latent_downsample_factor = latent_downsample_factor self.latent_channel_size = latent_channel_size self.lidar_latent_channel_size = lidar_latent_channel_size @@ -232,6 +237,7 @@ 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.timestep_shift = config.timestep_shift self.timestep_scale = config.timestep_scale self.latent_downsample = config.latent_downsample_factor * config.latent_patch_size @@ -249,11 +255,13 @@ def __init__(self, language_model, config: Cosmos3VFMNetworkConfig): # 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) if config.enable_vision_modality_embeddings: @@ -436,8 +444,10 @@ def patchify_and_pack_latents( 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 +463,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 +478,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 +501,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 +516,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 +541,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] @@ -799,6 +813,7 @@ def _encode_lidar( packed_sequence, vae2llm=self.lidar2llm, latent_channel=self.lidar_latent_channel, + patch_size=self.lidar_patch_spatial_hw, modality_embed=None, target_dtype=target_dtype, ) @@ -813,6 +828,7 @@ def _encode_grid_stream( latent_channel: int, modality_embed: torch.Tensor | None, target_dtype: torch.dtype, + patch_size: int | tuple[int, int] | None = None, ) -> List[Tuple[int, int, int]] | None: """Patchify, project and scatter one stream of VAE latent grids. @@ -836,7 +852,7 @@ 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 modality_embed is not None: @@ -903,6 +919,7 @@ 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, ) ) @@ -917,6 +934,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 +976,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( diff --git a/cosmos_framework/model/generator/omni_mot_model.py b/cosmos_framework/model/generator/omni_mot_model.py index 3d0c8e5f..64c59a55 100644 --- a/cosmos_framework/model/generator/omni_mot_model.py +++ b/cosmos_framework/model/generator/omni_mot_model.py @@ -478,6 +478,7 @@ 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, latent_downsample_factor=self.config.latent_downsample_factor, latent_channel_size=self.config.state_ch, lidar_latent_channel_size=self.config.lidar_state_ch, @@ -930,6 +931,7 @@ def _pack_input_sequence( 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, 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, 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