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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 6 additions & 3 deletions .file_mapping.json
Original file line number Diff line number Diff line change
@@ -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",
Expand Down Expand Up @@ -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",
Expand Down Expand Up @@ -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",
Expand Down
4 changes: 4 additions & 0 deletions cosmos_framework/configs/base/defaults/model_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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).
)
Expand Down
Original file line number Diff line number Diff line change
@@ -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
Original file line number Diff line number Diff line change
@@ -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
12 changes: 10 additions & 2 deletions cosmos_framework/data/generator/joint_dataloader.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down Expand Up @@ -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.
Expand All @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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}.")
Expand Down Expand Up @@ -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
Expand Down
7 changes: 7 additions & 0 deletions cosmos_framework/data/generator/multiview/caption_chunks.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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.

Expand All @@ -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.
Expand Down Expand Up @@ -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,
Expand Down
14 changes: 14 additions & 0 deletions cosmos_framework/data/generator/multiview/caption_format.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
from __future__ import annotations

import json
import re
from collections.abc import Mapping, Sequence
from typing import Any, Final

Expand All @@ -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")
Expand Down
8 changes: 6 additions & 2 deletions cosmos_framework/data/generator/sequence_packing/packers.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand All @@ -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
Expand Down Expand Up @@ -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,
Expand Down
14 changes: 7 additions & 7 deletions cosmos_framework/data/generator/sequence_packing/sequence.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand All @@ -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.
Expand Down Expand Up @@ -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,
Expand All @@ -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
Expand Down
Loading
Loading