diff --git a/.file_mapping.json b/.file_mapping.json
index 2ddbf36bf..9cbaf068d 100644
--- a/.file_mapping.json
+++ b/.file_mapping.json
@@ -1,7 +1,7 @@
{
- "_source_commit": "2962eb13e76ad641803d5f498c2822555c88e79c-dirty",
- "_dest_commit": "6c9d3916e0c5b62f8bed9694ae805d3d4c04c454",
- "_generated_at": "2026-09-18T13:40:26Z",
+ "_source_commit": "4d9b6cfd3731fdbbca883184937d40a64d2dca52-dirty",
+ "_dest_commit": "c23e51f2f157ae3e51cfcd86ebfb5464850894f2",
+ "_generated_at": "2026-09-20T05:50:07Z",
"files": {
"imaginaire/__init__.py": "cosmos_framework/__init__.py",
"imaginaire/attention/__init__.py": "cosmos_framework/model/attention/__init__.py",
@@ -292,6 +292,20 @@
"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/hr_lr_degradation/__init__.py": "cosmos_framework/data/generator/augmentors/hr_lr_degradation/__init__.py",
+ "projects/cosmos3/cosmos3/datasets/augmentors/hr_lr_degradation/augmentor.py": "cosmos_framework/data/generator/augmentors/hr_lr_degradation/augmentor.py",
+ "projects/cosmos3/cosmos3/datasets/augmentors/hr_lr_degradation/augmentor_test.py": "cosmos_framework/data/generator/augmentors/hr_lr_degradation/augmentor_test.py",
+ "projects/cosmos3/cosmos3/datasets/augmentors/hr_lr_degradation/bench.py": "cosmos_framework/data/generator/augmentors/hr_lr_degradation/bench.py",
+ "projects/cosmos3/cosmos3/datasets/augmentors/hr_lr_degradation/codec.py": "cosmos_framework/data/generator/augmentors/hr_lr_degradation/codec.py",
+ "projects/cosmos3/cosmos3/datasets/augmentors/hr_lr_degradation/contact_sheet.py": "cosmos_framework/data/generator/augmentors/hr_lr_degradation/contact_sheet.py",
+ "projects/cosmos3/cosmos3/datasets/augmentors/hr_lr_degradation/contact_sheet_test.py": "cosmos_framework/data/generator/augmentors/hr_lr_degradation/contact_sheet_test.py",
+ "projects/cosmos3/cosmos3/datasets/augmentors/hr_lr_degradation/degrade.py": "cosmos_framework/data/generator/augmentors/hr_lr_degradation/degrade.py",
+ "projects/cosmos3/cosmos3/datasets/augmentors/hr_lr_degradation/degrade_test.py": "cosmos_framework/data/generator/augmentors/hr_lr_degradation/degrade_test.py",
+ "projects/cosmos3/cosmos3/datasets/augmentors/hr_lr_degradation/diffjpeg.py": "cosmos_framework/data/generator/augmentors/hr_lr_degradation/diffjpeg.py",
+ "projects/cosmos3/cosmos3/datasets/augmentors/hr_lr_degradation/kernels.py": "cosmos_framework/data/generator/augmentors/hr_lr_degradation/kernels.py",
+ "projects/cosmos3/cosmos3/datasets/augmentors/hr_lr_degradation/ops.py": "cosmos_framework/data/generator/augmentors/hr_lr_degradation/ops.py",
+ "projects/cosmos3/cosmos3/datasets/augmentors/hr_lr_degradation/packing_test.py": "cosmos_framework/data/generator/augmentors/hr_lr_degradation/packing_test.py",
+ "projects/cosmos3/cosmos3/datasets/augmentors/hr_lr_degradation/profiles.py": "cosmos_framework/data/generator/augmentors/hr_lr_degradation/profiles.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",
@@ -339,6 +353,7 @@
"projects/cosmos3/cosmos3/datasets/augmentors/transfer_control_input/seg.py": "cosmos_framework/data/generator/augmentors/transfer_control_input/seg.py",
"projects/cosmos3/cosmos3/datasets/augmentors/transfer_control_transform.py": "cosmos_framework/data/generator/augmentors/transfer_control_transform.py",
"projects/cosmos3/cosmos3/datasets/augmentors/transfer_control_transform_test.py": "cosmos_framework/data/generator/augmentors/transfer_control_transform_test.py",
+ "projects/cosmos3/cosmos3/datasets/augmentors/transfer_target_length_test.py": "cosmos_framework/data/generator/augmentors/transfer_target_length_test.py",
"projects/cosmos3/cosmos3/datasets/augmentors/video_editing_transform.py": "cosmos_framework/data/generator/augmentors/video_editing_transform.py",
"projects/cosmos3/cosmos3/datasets/augmentors/video_editing_transform_test.py": "cosmos_framework/data/generator/augmentors/video_editing_transform_test.py",
"projects/cosmos3/cosmos3/datasets/augmentors/video_parsing.py": "cosmos_framework/data/generator/augmentors/video_parsing.py",
@@ -371,6 +386,8 @@
"projects/cosmos3/cosmos3/models/mot/__init__.py": "cosmos_framework/model/generator/mot/__init__.py",
"projects/cosmos3/cosmos3/models/mot/action_io_projector.py": "cosmos_framework/model/generator/mot/action_io_projector.py",
"projects/cosmos3/cosmos3/models/mot/action_io_projector_test.py": "cosmos_framework/model/generator/mot/action_io_projector_test.py",
+ "projects/cosmos3/cosmos3/models/mot/activation_marks.py": "cosmos_framework/model/generator/mot/activation_marks.py",
+ "projects/cosmos3/cosmos3/models/mot/activation_marks_test.py": "cosmos_framework/model/generator/mot/activation_marks_test.py",
"projects/cosmos3/cosmos3/models/mot/attention.py": "cosmos_framework/model/generator/mot/attention.py",
"projects/cosmos3/cosmos3/models/mot/attention_test.py": "cosmos_framework/model/generator/mot/attention_test.py",
"projects/cosmos3/cosmos3/models/mot/cfgp_ar_test.py": "cosmos_framework/model/generator/mot/cfgp_ar_test.py",
@@ -460,12 +477,15 @@
"projects/cosmos3/cosmos3/models/reasoner/qwen3_vl_moe/shared_expert_test.py": "cosmos_framework/model/generator/reasoner/qwen3_vl_moe/shared_expert_test.py",
"projects/cosmos3/cosmos3/models/utils/__init__.py": "cosmos_framework/model/generator/utils/__init__.py",
"projects/cosmos3/cosmos3/models/utils/data_and_condition.py": "cosmos_framework/model/generator/utils/data_and_condition.py",
+ "projects/cosmos3/cosmos3/models/utils/data_and_condition_test.py": "cosmos_framework/model/generator/utils/data_and_condition_test.py",
"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/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",
+ "projects/cosmos3/cosmos3/models/utils/sr_latent_noise.py": "cosmos_framework/model/generator/utils/sr_latent_noise.py",
+ "projects/cosmos3/cosmos3/models/utils/sr_latent_noise_test.py": "cosmos_framework/model/generator/utils/sr_latent_noise_test.py",
"projects/cosmos3/cosmos3/models/vision_encoder.py": "cosmos_framework/model/generator/vision_encoder.py",
"projects/cosmos3/cosmos3/models/vlm_model.py": "cosmos_framework/model/generator/vlm_model.py",
"projects/cosmos3/cosmos3/processors/__init__.py": "cosmos_framework/data/generator/processors/__init__.py",
@@ -492,6 +512,7 @@
"projects/cosmos3/cosmos3/sequence_packing/natten.py": "cosmos_framework/data/generator/sequence_packing/natten.py",
"projects/cosmos3/cosmos3/sequence_packing/packers.py": "cosmos_framework/data/generator/sequence_packing/packers.py",
"projects/cosmos3/cosmos3/sequence_packing/runtime.py": "cosmos_framework/data/generator/sequence_packing/runtime.py",
+ "projects/cosmos3/cosmos3/sequence_packing/runtime_test.py": "cosmos_framework/data/generator/sequence_packing/runtime_test.py",
"projects/cosmos3/cosmos3/sequence_packing/sequence.py": "cosmos_framework/data/generator/sequence_packing/sequence.py",
"projects/cosmos3/cosmos3/sequence_packing/temporal_causal.py": "cosmos_framework/data/generator/sequence_packing/temporal_causal.py",
"projects/cosmos3/cosmos3/tokenizers/audio/__init__.py": "cosmos_framework/model/generator/tokenizers/audio/__init__.py",
@@ -584,6 +605,26 @@
"projects/cosmos3/interactive/models/utils/nvfp4_test.py": "cosmos_framework/model/generator/utils/nvfp4_test.py",
"projects/cosmos3/interactive/sequence_packing.py": "cosmos_framework/data/generator/sequence_packing/autoregressive.py",
"projects/cosmos3/interactive/utils/data_batch.py": "cosmos_framework/utils/generator/data_batch.py",
+ "projects/cosmos3/tokenizer/lidar_tokenizer/__init__.py": "cosmos_framework/model/generator/tokenizers/lidar/__init__.py",
+ "projects/cosmos3/tokenizer/lidar_tokenizer/checkpoint.py": "cosmos_framework/model/generator/tokenizers/lidar/checkpoint.py",
+ "projects/cosmos3/tokenizer/lidar_tokenizer/checkpoint_test.py": "cosmos_framework/model/generator/tokenizers/lidar/checkpoint_test.py",
+ "projects/cosmos3/tokenizer/lidar_tokenizer/dtypes.py": "cosmos_framework/model/generator/tokenizers/lidar/dtypes.py",
+ "projects/cosmos3/tokenizer/lidar_tokenizer/dtypes_test.py": "cosmos_framework/model/generator/tokenizers/lidar/dtypes_test.py",
+ "projects/cosmos3/tokenizer/lidar_tokenizer/lidar_tokenizer_v1.py": "cosmos_framework/model/generator/tokenizers/lidar/lidar_tokenizer_v1.py",
+ "projects/cosmos3/tokenizer/lidar_tokenizer/network/__init__.py": "cosmos_framework/model/generator/tokenizers/lidar/network/__init__.py",
+ "projects/cosmos3/tokenizer/lidar_tokenizer/network/distributions.py": "cosmos_framework/model/generator/tokenizers/lidar/network/distributions.py",
+ "projects/cosmos3/tokenizer/lidar_tokenizer/network/encoding.py": "cosmos_framework/model/generator/tokenizers/lidar/network/encoding.py",
+ "projects/cosmos3/tokenizer/lidar_tokenizer/network/ops.py": "cosmos_framework/model/generator/tokenizers/lidar/network/ops.py",
+ "projects/cosmos3/tokenizer/lidar_tokenizer/network/rope3d.py": "cosmos_framework/model/generator/tokenizers/lidar/network/rope3d.py",
+ "projects/cosmos3/tokenizer/lidar_tokenizer/network/transformer_vae.py": "cosmos_framework/model/generator/tokenizers/lidar/network/transformer_vae.py",
+ "projects/cosmos3/tokenizer/lidar_tokenizer/normalization.py": "cosmos_framework/model/generator/tokenizers/lidar/normalization.py",
+ "projects/cosmos3/tokenizer/lidar_tokenizer/normalization_test.py": "cosmos_framework/model/generator/tokenizers/lidar/normalization_test.py",
+ "projects/cosmos3/tokenizer/lidar_tokenizer/paths.py": "cosmos_framework/model/generator/tokenizers/lidar/paths.py",
+ "projects/cosmos3/tokenizer/lidar_tokenizer/postprocessing.py": "cosmos_framework/model/generator/tokenizers/lidar/postprocessing.py",
+ "projects/cosmos3/tokenizer/lidar_tokenizer/postprocessing_test.py": "cosmos_framework/model/generator/tokenizers/lidar/postprocessing_test.py",
+ "projects/cosmos3/tokenizer/lidar_tokenizer/preprocessing.py": "cosmos_framework/model/generator/tokenizers/lidar/preprocessing.py",
+ "projects/cosmos3/tokenizer/lidar_tokenizer/range_projection.py": "cosmos_framework/model/generator/tokenizers/lidar/range_projection.py",
+ "projects/cosmos3/tokenizer/lidar_tokenizer/range_projection_test.py": "cosmos_framework/model/generator/tokenizers/lidar/range_projection_test.py",
"projects/cosmos3/tokenizer/uniae/checkpoint_identity.py": "cosmos_framework/model/tokenizer/checkpoint_identity.py",
"projects/cosmos3/tokenizer/uniae/checkpoint_io.py": "cosmos_framework/model/tokenizer/checkpoint_io.py",
"projects/cosmos3/tokenizer/uniae/evaluation/lpips_cache.py": "cosmos_framework/model/tokenizer/evaluation/lpips_cache.py",
diff --git a/cosmos_framework/configs/base/defaults/activation_checkpointing.py b/cosmos_framework/configs/base/defaults/activation_checkpointing.py
index 3eb0b0bbc..1e80e957a 100644
--- a/cosmos_framework/configs/base/defaults/activation_checkpointing.py
+++ b/cosmos_framework/configs/base/defaults/activation_checkpointing.py
@@ -17,6 +17,20 @@
import attrs
+# The forward attention op of every backend the frontend can dispatch, by the name
+# the checkpoint policy matches against.
+#
+# * natten: fmha_forward
+# * flash2: _flash_attn_varlen_forward (varlen), _flash_attn_forward (dense)
+# * flash3: _flash_attn_forward (dense, varlen)
+#
+# cuDNN is deliberately absent since it does not support varlen attention.
+ATTENTION_FORWARD_OPS_REGEX = [
+ "fmha_forward",
+ "_flash_attn_varlen_forward",
+ "_flash_attn_forward",
+]
+
@attrs.define(slots=False)
class ActivationCheckpointingConfig:
@@ -45,10 +59,34 @@ class ActivationCheckpointingConfig:
# Regex patterns for ops to save when using selective AC. Ignored if
# mode is "full" or "none".
+ #
+ # Defaults to attention on whichever backend the frontend picks, which is what
+ # every config asking for selective AC wants and is not what a name like "fmha"
+ # delivers: that covers NATTEN alone, and NATTEN is only the *selected* backend
+ # where the others refuse the call. On sm100 cuDNN and flash2 both reject varlen
+ # so NATTEN wins and "fmha" matches; on sm90 flash3 is ranked first and takes it,
+ # leaving nothing in the region named "fmha", so selective AC there kept nothing
+ # and silently recomputed every attention it was configured to save.
+ #
+ # Copied because attrs hands this list to the instance, and a config mutating it
+ # would edit the module constant for every other config in the process.
save_ops_regex: list[str] = attrs.field(
- factory=lambda: ["fmha"],
+ factory=lambda: list(ATTENTION_FORWARD_OPS_REGEX),
)
+ # Narrow ``save_ops_regex`` to the call sites that asked to be kept.
+ #
+ # The regex matches dispatched op names, which cannot separate calls running
+ # the same kernel. The decomposed multiview attention runs four FMHA calls per
+ # layer worth very different amounts to keep -- the same-view fold is ~96% of
+ # forward attention time and ~94% of backward, the rest are cheap to recompute.
+ #
+ # With this on, a matching op is kept only where the model marked it with
+ # ``activation_marks.mark_next_activation``. The regex still says which ops are
+ # eligible; the mark says which of them are worth it. Off by default, which is
+ # the behaviour every existing config already has.
+ save_only_marked_ops: bool = False
+
# Stash and restore RNG state across recompute boundaries. Required for
# deterministic output vs. non-checkpointed passes; slower otherwise.
preserve_rng_state: bool = True
diff --git a/cosmos_framework/configs/base/defaults/compile.py b/cosmos_framework/configs/base/defaults/compile.py
index 3a0be06f9..93b24f22f 100644
--- a/cosmos_framework/configs/base/defaults/compile.py
+++ b/cosmos_framework/configs/base/defaults/compile.py
@@ -39,6 +39,18 @@ class CompileConfig:
# Whether to use CUDA graphs for faster inference. This option does not work during training.
use_cuda_graphs: bool = False
+ # Granularity of CUDA-graph capture for AR inference (only with ``enabled`` and
+ # ``use_cuda_graphs``). "block": every decoder block is compiled with
+ # ``torch.compile(mode="reduce-overhead")`` and replays its own CUDA-graph tree, so a
+ # forward still pays one graph launch per block plus the Python between blocks.
+ # "forward": blocks are compiled without CUDA-graph trees and the AR loop captures one
+ # explicit graph per whole forward (denoise / KV-refresh × CFG branch) on the
+ # static-shape KV path from cache index 1 onward.
+ cuda_graph_scope: Literal["block", "forward"] = attrs.field(
+ default="block",
+ validator=attrs.validators.in_({"block", "forward"}),
+ )
+
# AR-inference-specific behavior once the rolling KV window saturates.
# "default" uses the global compile settings for the entire generation.
# "static-compile" keeps the normal pre-saturation path, then uses dedicated
diff --git a/cosmos_framework/configs/base/defaults/model_config.py b/cosmos_framework/configs/base/defaults/model_config.py
index cde908218..777f8a59f 100644
--- a/cosmos_framework/configs/base/defaults/model_config.py
+++ b/cosmos_framework/configs/base/defaults/model_config.py
@@ -16,6 +16,7 @@
from cosmos_framework.configs.base.defaults.reasoner import VLMConfig
from cosmos_framework.model.generator.mot.action_io_projector import ACTION_IO_PROJECTOR_TYPES
from cosmos_framework.model.generator.utils.load_balancing_stats import LBLConfig
+from cosmos_framework.model.generator.utils.sr_latent_noise import SRLatentConditionNoiseConfig
# Mirrors ``cosmos3.common.args.AttentionIOLayout``. Defined locally on purpose: importing
# the ``cosmos3`` workspace package at module scope makes the whole cosmos3 config tree
@@ -351,6 +352,10 @@ class OmniMoTModelConfig:
sound_dim: int | None = None # Sound latent channel size (e.g., 64 for AVAE 48kHz)
sound_latent_fps: int = 25 # Sound tokenizer's latent rate (e.g., 48kHz / 1920 hop = 25 Hz)
+ # Super-resolution: Gaussian noise on the LR conditioning latent of SR samples during training (L1).
+ # None disables it. See cosmos_framework/model/generator/utils/sr_latent_noise.py.
+ sr_latent_condition_noise: SRLatentConditionNoiseConfig | None = None
+
# When False, removes bias from vae2llm, sound2llm, and the two Linear layers inside
# time_embedder. These biases seem to inject token-constant DC offsets that dominate
# the MoE router input and create prompt invariant routing. This is observed empirically
diff --git a/cosmos_framework/data/generator/augmentor_provider.py b/cosmos_framework/data/generator/augmentor_provider.py
index 83e24ebdc..7fc757b65 100644
--- a/cosmos_framework/data/generator/augmentor_provider.py
+++ b/cosmos_framework/data/generator/augmentor_provider.py
@@ -695,7 +695,7 @@ def get_video_augmentor_v3(
return augmentors
-# Use video_basic_augmentor_v3_json_caption instead.
+# Deprecated: sound is trained only with json caption; use video_basic_augmentor_v3_json_caption instead.
@augmentor_register("video_basic_augmentor_v3_with_audio")
def get_video_augmentor_v3_with_audio(
resolution: str,
@@ -1519,3 +1519,136 @@ def image_basic_augmentor_json_caption(
}
return augmentation
+
+
+def _insert_relative(augmentors: dict, anchor: str, new_key: str, new_value, *, after: bool) -> dict:
+ """Return a copy of ``augmentors`` with ``new_key`` inserted right before or right after ``anchor``."""
+ if anchor not in augmentors:
+ raise KeyError(f"{anchor!r} not found in pipeline; cannot insert {new_key}")
+ out: dict = {}
+ for key, value in augmentors.items():
+ if key == anchor and not after:
+ out[new_key] = new_value
+ out[key] = value
+ if key == anchor and after:
+ out[new_key] = new_value
+ return out
+
+
+def _insert_before(augmentors: dict, anchor_keys: tuple[str, ...], new_key: str, new_value) -> dict:
+ """Return a copy of ``augmentors`` with ``new_key`` inserted before the first present anchor key."""
+ anchor = next((k for k in anchor_keys if k in augmentors), None)
+ if anchor is None:
+ raise KeyError(f"None of {anchor_keys} found in pipeline; cannot insert {new_key}")
+ return _insert_relative(augmentors, anchor, new_key, new_value, after=False)
+
+
+def _insert_low_res_stage(augmentors: dict, add_low_res) -> dict:
+ """Place ``AddLowRes`` so the LR is derived from exactly the HR frame the model will see.
+
+ - Reflection-padding path (causal VAE): LR is made from the *unpadded* frame, before ``reflection_padding``;
+ ``SRToTrainingFormat`` pads LR separately to half the HR bucket, so LR and HR stay aligned at the top-left.
+ - Crop path (non-causal / UniAE, ``crop_to_multiple``): LR is made *after* the centre crop. Making it before
+ would derive LR from pixels the HR no longer contains (spatial misalignment) and, when the crop changes the
+ size, a larger LR than the target that ``SRToTrainingFormat`` cannot pad down.
+ """
+ if "reflection_padding" in augmentors:
+ return _insert_relative(augmentors, "reflection_padding", "add_low_res", add_low_res, after=False)
+ if "crop_to_multiple" in augmentors:
+ return _insert_relative(augmentors, "crop_to_multiple", "add_low_res", add_low_res, after=True)
+ raise KeyError("Pipeline has neither reflection_padding nor crop_to_multiple; cannot place add_low_res")
+
+
+@augmentor_register("video_basic_augmentor_v3_json_caption_sr")
+def get_video_augmentor_v3_json_caption_sr(
+ resolution: str,
+ sr_scale: float = 2.0,
+ sr_profiles: dict[str, float] | str = "p1_first_order",
+ sr_seed_salt: str = "",
+ sr_chunk_frames: int = 8,
+ sr_device: str = "cpu",
+ sr_jpeg_backend: str = "auto",
+ sr_poisson_mode: str = "auto",
+ sr_share_vision_temporal_positions: bool = False,
+ **kwargs: object,
+) -> dict[str, object]:
+ """``video_basic_augmentor_v3_json_caption`` plus an on-the-fly HR-to-LR conditioning stream.
+
+ Adds ``AddLowRes`` (writes ``video_lr`` at ``1/sr_scale`` of the HR frame, uint8) right before
+ reflection padding, and ``SRToTrainingFormat`` as the last stage, which pads LR to half the HR
+ bucket and packs ``video = [lr, hr]`` with per-item ``image_size`` and a two-item SequencePlan.
+ All other stages (caption, chunked decode, sequence plan, sound) are inherited unchanged.
+ """
+ from cosmos_framework.data.generator.augmentors.hr_lr_degradation import augmentor as sr_augmentor
+
+ augmentors = get_video_augmentor_v3_json_caption(resolution=resolution, **kwargs)
+ add_low_res = L(sr_augmentor.AddLowRes)(
+ input_keys=["video"],
+ output_keys=["video_lr"],
+ args={
+ "scale": sr_scale,
+ "profiles": sr_profiles,
+ "seed_salt": sr_seed_salt,
+ "modality": "video",
+ "chunk_frames": sr_chunk_frames,
+ "device": sr_device,
+ "jpeg_backend": sr_jpeg_backend,
+ "poisson_mode": sr_poisson_mode,
+ },
+ )
+ augmentors = _insert_low_res_stage(augmentors, add_low_res)
+ augmentors["sr_to_training_format"] = L(sr_augmentor.SRToTrainingFormat)(
+ input_keys=["video", "video_lr"],
+ args={
+ "media_key": "video",
+ "lr_key": "video_lr",
+ "scale": sr_scale,
+ "share_vision_temporal_positions": sr_share_vision_temporal_positions,
+ "dataset_name": "video_sr",
+ },
+ )
+ return augmentors
+
+
+@augmentor_register("image_basic_augmentor_with_tokenization_sr")
+def image_basic_augmentor_with_tokenization_sr(
+ resolution: str,
+ sr_scale: float = 2.0,
+ sr_profiles: dict[str, float] | str = "p1_first_order",
+ sr_seed_salt: str = "",
+ sr_device: str = "cpu",
+ sr_jpeg_backend: str = "auto",
+ sr_poisson_mode: str = "auto",
+ **kwargs: object,
+) -> dict[str, object]:
+ """``image_basic_augmentor_with_tokenization`` plus an on-the-fly HR-to-LR conditioning image.
+
+ ``AddLowRes`` runs before ``reflection_padding`` (on the resized uint8 image), the LR copy gets
+ its own ``Normalize`` so both items reach the model as float in [-1, 1], and
+ ``SRToTrainingFormat`` packs ``images = [lr, hr]`` with per-item ``image_size``.
+ """
+ from cosmos_framework.data.generator.augmentors.hr_lr_degradation import augmentor as sr_augmentor
+
+ augmentors = image_basic_augmentor_with_tokenization(resolution=resolution, **kwargs)
+ add_low_res = L(sr_augmentor.AddLowRes)(
+ input_keys=["images"],
+ output_keys=["images_lr"],
+ args={
+ "scale": sr_scale,
+ "profiles": sr_profiles,
+ "seed_salt": sr_seed_salt,
+ "modality": "image",
+ "chunk_frames": 1,
+ "device": sr_device,
+ "jpeg_backend": sr_jpeg_backend,
+ "poisson_mode": sr_poisson_mode,
+ },
+ )
+ augmentors = _insert_before(augmentors, ("reflection_padding",), "add_low_res", add_low_res)
+ normalize_lr = L(normalize.Normalize)(input_keys=["images_lr"], args={"mean": 0.5, "std": 0.5})
+ augmentors = _insert_before(augmentors, ("text_transform",), "normalize_lr", normalize_lr)
+ augmentors["sr_to_training_format"] = L(sr_augmentor.SRToTrainingFormat)(
+ input_keys=["images", "images_lr"],
+ args={"media_key": "images", "lr_key": "images_lr", "scale": sr_scale, "dataset_name": "image_sr"},
+ )
+ return augmentors
diff --git a/cosmos_framework/data/generator/augmentors/hr_lr_degradation/__init__.py b/cosmos_framework/data/generator/augmentors/hr_lr_degradation/__init__.py
new file mode 100644
index 000000000..b4cb9e0f3
--- /dev/null
+++ b/cosmos_framework/data/generator/augmentors/hr_lr_degradation/__init__.py
@@ -0,0 +1,35 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+"""On-the-fly HR-to-LR degradation operators for super-resolution training.
+
+The package is device-agnostic (CPU dataloader workers or GPU) and fully seeded:
+``degrade_hr_to_lr(hr, profile, scale, seed)`` returns the same LR for the same
+inputs on every call, and the sampled parameters are returned as a record.
+
+Modules:
+- ``kernels``: blur kernel generators (Real-ESRGAN / BasicSR lineage), numpy.
+- ``diffjpeg``: torch JPEG round trip (DiffJPEG lineage), any device.
+- ``ops``: per-clip primitives on ``[T,C,H,W]`` float tensors in [0, 1].
+- ``profiles``: dataclass configs and the named profile registry (P0, P1, ...).
+- ``degrade``: the entry point that runs a profile on an HR clip or image.
+"""
+
+from cosmos_framework.data.generator.augmentors.hr_lr_degradation.degrade import (
+ DegradationResult,
+ degrade_hr_to_lr,
+)
+from cosmos_framework.data.generator.augmentors.hr_lr_degradation.profiles import (
+ PROFILES,
+ CleanResizeProfile,
+ RealESRGANProfile,
+ get_profile,
+)
+
+__all__ = [
+ "PROFILES",
+ "CleanResizeProfile",
+ "DegradationResult",
+ "RealESRGANProfile",
+ "degrade_hr_to_lr",
+ "get_profile",
+]
diff --git a/cosmos_framework/data/generator/augmentors/hr_lr_degradation/augmentor.py b/cosmos_framework/data/generator/augmentors/hr_lr_degradation/augmentor.py
new file mode 100644
index 000000000..0abcd8c08
--- /dev/null
+++ b/cosmos_framework/data/generator/augmentors/hr_lr_degradation/augmentor.py
@@ -0,0 +1,236 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+"""Dataloader augmentors that turn a single-stream sample into an (LR, HR) super-resolution sample.
+
+Two stages, meant for the cosmos3 Lance video/image pipelines:
+
+``AddLowRes`` runs right after the media is at its final HR resolution and before reflection
+padding. It writes ``data_dict[output_key]`` (uint8, same layout as the input) at ``1/scale`` of
+the HR size using a seeded degradation profile, plus ``data_dict[record_key]``, a JSON string with every sampled
+parameter (a string collates cleanly; a dict of variable-length lists does not). The seed derives from the sample key so a sample always gets the same LR.
+
+``SRToTrainingFormat`` runs last. It pads LR to half the HR padding bucket, packs
+``data_dict[media_key] = [lr, hr]`` (the joint dataloader treats every item before the last as
+pure conditioning, as in ``TransferToTrainingFormat``), writes one ``image_size`` entry per item so
+``OmniMoTModel._remove_padding_from_latent`` crops each latent correctly, and marks the
+``SequencePlan`` with ``share_vision_temporal_positions=False`` because the two items have
+different latent grids (``sequence_packing/packers.py`` asserts equal grids when sharing).
+"""
+
+from __future__ import annotations
+
+import hashlib
+import json
+import random
+from typing import Any, Mapping, Optional
+
+import numpy as np
+import torch
+import torchvision.transforms.functional as transforms_F
+
+from cosmos_framework.data.imaginaire.webdataset.augmentors.augmentor import Augmentor
+from cosmos_framework.utils import log
+from cosmos_framework.data.generator.augmentors.hr_lr_degradation.degrade import degrade_hr_to_lr
+from cosmos_framework.data.generator.augmentors.hr_lr_degradation.diffjpeg import DiffJPEG
+from cosmos_framework.data.generator.augmentors.hr_lr_degradation.profiles import get_profile
+from cosmos_framework.data.generator.sequence_packing import SequencePlan
+
+DEFAULT_RECORD_KEY = "degradation_record"
+
+
+def seed_from_sample(data_dict: Mapping[str, Any], salt: str = "") -> int:
+ """Stable 63-bit seed from the sample identity (``__key__``), or a random one when absent."""
+ key = data_dict.get("__key__")
+ if key is None:
+ return random.getrandbits(63)
+ digest = hashlib.blake2b(f"{key}|{salt}".encode(), digest_size=8).digest()
+ return int.from_bytes(digest, "little") & 0x7FFF_FFFF_FFFF_FFFF
+
+
+DEFAULT_FPS = 24.0
+
+
+def clip_fps(data_dict: Mapping[str, Any]) -> float:
+ """Effective frame rate of the frames in the sample, for the codec stage.
+
+ ``conditioning_fps`` is the native rate divided by the sampled stride, i.e. the rate at which the kept
+ frames actually play and the rate the model, captions and mRoPE use. ``fps`` is the source file's
+ native rate and is only right when the stride is 1. Images carry neither and fall back to a default
+ (the codec is skipped for them anyway).
+ """
+ for key in ("conditioning_fps", "fps"):
+ value = data_dict.get(key)
+ if value is None:
+ continue
+ if isinstance(value, torch.Tensor):
+ value = value.reshape(-1)[0].item()
+ if float(value) > 0:
+ return float(value)
+ return DEFAULT_FPS
+
+
+def _as_uint8_media(frames: Any) -> torch.Tensor: # returns [C,T,H,W] or [C,H,W] uint8
+ if isinstance(frames, np.ndarray):
+ frames = torch.from_numpy(frames)
+ if not isinstance(frames, torch.Tensor):
+ raise TypeError(f"AddLowRes expects a tensor or ndarray, got {type(frames).__name__}")
+ if frames.dtype == torch.uint8:
+ return frames
+ if frames.is_floating_point():
+ if frames.numel() > 0 and frames.min() < 0.0:
+ raise ValueError("AddLowRes must run before normalisation (got values below 0)")
+ return (frames.clamp(0.0, 1.0) * 255.0).round().to(torch.uint8)
+ raise TypeError(f"Unsupported media dtype {frames.dtype}")
+
+
+class AddLowRes(Augmentor):
+ """Add a degraded low-resolution copy of ``input_keys[0]`` under ``output_keys[0]``.
+
+ args:
+ scale: HR-to-LR factor (default 2).
+ profiles: mapping profile name -> sampling weight, or a single profile name string.
+ seed_salt: extra string mixed into the per-sample seed (use to decorrelate ablation arms).
+ chunk_frames: frames per degradation step (peak-memory bound).
+ device: ``"cpu"`` (dataloader workers) or a CUDA device string for GPU-side use.
+ jpeg_backend / poisson_mode: forwarded to ``degrade_hr_to_lr``.
+ modality: ``"image"`` or ``"video"``; when set, profiles declared for the other modality are rejected
+ at construction (a video regime on single images would have no compression term at all).
+ record_key: where the parameter record is written.
+ """
+
+ def __init__(self, input_keys: list, output_keys: Optional[list] = None, args: Optional[dict] = None) -> None:
+ super().__init__(input_keys, output_keys, args)
+ args = dict(args or {})
+ if len(self.input_keys) != 1:
+ raise ValueError("AddLowRes takes exactly one input key")
+ self.output_key = (output_keys or [f"{self.input_keys[0]}_lr"])[0]
+ self.scale = float(args.get("scale", 2.0))
+ profiles = args.get("profiles", "p1_first_order")
+ if isinstance(profiles, str):
+ profiles = {profiles: 1.0}
+ self.profile_names = list(profiles.keys())
+ weights = np.asarray([float(profiles[n]) for n in self.profile_names], dtype=np.float64) # [P]
+ if weights.sum() <= 0:
+ raise ValueError("profile weights must sum to a positive number")
+ self.profile_weights = weights / weights.sum() # [P]
+ self.modality = args.get("modality")
+ for name in self.profile_names:
+ prof = get_profile(name) # fail early on typos
+ if self.modality is not None and prof.modality not in ("any", self.modality):
+ raise ValueError(f"profile {name!r} is for {prof.modality} data, this AddLowRes serves {self.modality}")
+ self.seed_salt = str(args.get("seed_salt", ""))
+ self.chunk_frames = int(args.get("chunk_frames", 8))
+ self.device = torch.device(args.get("device", "cpu"))
+ self.jpeg_backend = str(args.get("jpeg_backend", "auto"))
+ self.poisson_mode = str(args.get("poisson_mode", "auto"))
+ self.record_key = str(args.get("record_key", DEFAULT_RECORD_KEY))
+ self._jpeger: DiffJPEG | None = None
+
+ def _jpeger_for(self, device: torch.device) -> DiffJPEG:
+ if self._jpeger is None:
+ self._jpeger = DiffJPEG(differentiable=False)
+ if self._jpeger.y_table.device != device:
+ self._jpeger.to(device)
+ return self._jpeger
+
+ def __call__(self, data_dict: dict) -> dict | None:
+ media = data_dict.get(self.input_keys[0])
+ if media is None:
+ log.warning(f"AddLowRes: missing {self.input_keys[0]} in {data_dict.get('__key__', 'unknown')}")
+ return None
+ hr = _as_uint8_media(media) # [C,T,H,W] or [C,H,W]
+ seed = seed_from_sample(data_dict, self.seed_salt)
+ # The profile draw must not share a stream with the plan: degrade_hr_to_lr re-creates default_rng(seed),
+ # so drawing from default_rng(seed) here would make the mixture choice and the plan's first probability
+ # gate the same uniform (e.g. a 30% clean regime whose only gate then fires 100% of the time).
+ profile_rng = np.random.default_rng(seed_from_sample(data_dict, self.seed_salt + "|profile"))
+ profile_name = self.profile_names[int(profile_rng.choice(len(self.profile_names), p=self.profile_weights))]
+ result = degrade_hr_to_lr(
+ hr.to(self.device, non_blocking=False),
+ profile_name,
+ scale=self.scale,
+ seed=seed,
+ chunk_frames=self.chunk_frames,
+ jpeger=self._jpeger_for(self.device),
+ jpeg_backend=self.jpeg_backend,
+ poisson_mode=self.poisson_mode,
+ fps=clip_fps(data_dict), # only the codec stage (P3) uses it
+ )
+ data_dict[self.output_key] = result.lr.cpu() # [C,T,h,w] or [C,h,w] uint8
+ # JSON string: records differ in length between samples, so a dict would break default_collate.
+ data_dict[self.record_key] = json.dumps(result.record)
+ return data_dict
+
+
+def _pad_to(frames: torch.Tensor, target_h: int, target_w: int) -> torch.Tensor: # frames: [...,H,W]
+ """One-sided reflect padding (bottom/right), edge padding when the pad exceeds the content."""
+ h, w = frames.shape[-2:]
+ pad_right, pad_bottom = target_w - w, target_h - h
+ if pad_right < 0 or pad_bottom < 0:
+ raise ValueError(f"Cannot pad {(h, w)} to smaller target {(target_h, target_w)}")
+ if pad_right == 0 and pad_bottom == 0:
+ return frames
+ mode = "edge" if (pad_right >= w or pad_bottom >= h) else "reflect"
+ return transforms_F.pad(frames, [0, 0, pad_right, pad_bottom], padding_mode=mode) # [...,tH,tW]
+
+
+class SRToTrainingFormat(Augmentor):
+ """Pack (LR, HR) into the two-item conditioning format with per-item ``image_size``.
+
+ args:
+ media_key: ``"video"`` or ``"images"`` (the HR key; the LR key defaults to ``f"{media_key}_lr"``).
+ lr_key: override for the LR key.
+ scale: HR-to-LR factor; the LR padding bucket is the HR bucket divided by this.
+ share_vision_temporal_positions: keep False for native-resolution LR (default).
+ dataset_name: value written to ``data_dict["dataset_name"]``.
+ drop_lr_key: remove the standalone LR key after packing (default True).
+ """
+
+ def __init__(self, input_keys: list, output_keys: Optional[list] = None, args: Optional[dict] = None) -> None:
+ super().__init__(input_keys, output_keys, args)
+ args = dict(args or {})
+ self.media_key = str(args.get("media_key", "video"))
+ self.lr_key = str(args.get("lr_key", f"{self.media_key}_lr"))
+ self.scale = float(args.get("scale", 2.0))
+ self.share_vision_temporal_positions = bool(args.get("share_vision_temporal_positions", False))
+ default_name = "image_sr" if self.media_key == "images" else f"{self.media_key}_sr"
+ self.dataset_name = str(args.get("dataset_name", default_name))
+ self.drop_lr_key = bool(args.get("drop_lr_key", True))
+
+ def __call__(self, data_dict: dict) -> dict | None:
+ hr = data_dict.get(self.media_key)
+ lr = data_dict.get(self.lr_key)
+ if hr is None or lr is None or not isinstance(hr, torch.Tensor) or not isinstance(lr, torch.Tensor):
+ log.warning(
+ f"SRToTrainingFormat: missing {self.media_key} or {self.lr_key} in {data_dict.get('__key__', 'unknown')}",
+ rank0_only=False,
+ )
+ return None
+ hr_size = data_dict.get("image_size")
+ if hr_size is None:
+ hr_size = torch.tensor([hr.shape[-2], hr.shape[-1], hr.shape[-2], hr.shape[-1]], dtype=torch.float) # [4]
+ hr_size = torch.as_tensor(hr_size, dtype=torch.float).reshape(-1) # [4] = [tH,tW,oH,oW]
+ target_h, target_w = int(hr_size[0].item()), int(hr_size[1].item())
+ lr_target_h = int(np.ceil(target_h / self.scale))
+ lr_target_w = int(np.ceil(target_w / self.scale))
+ lr_orig_h, lr_orig_w = int(lr.shape[-2]), int(lr.shape[-1])
+ lr_padded = _pad_to(lr, lr_target_h, lr_target_w) # [C,T,th,tw] or [C,th,tw]
+ if lr_padded.dtype != hr.dtype:
+ # HR may already be normalised float (image pipeline); match dtype so the model treats both alike.
+ if hr.is_floating_point() and lr_padded.dtype == torch.uint8:
+ raise ValueError(
+ "HR is float but LR is uint8: add a Normalize stage for the LR key before SRToTrainingFormat"
+ )
+ lr_size = torch.tensor([lr_target_h, lr_target_w, lr_orig_h, lr_orig_w], dtype=torch.float) # [4]
+
+ data_dict[self.media_key] = [lr_padded, hr]
+ data_dict["image_size"] = [lr_size, hr_size]
+ data_dict["dataset_name"] = self.dataset_name
+ plan = data_dict.get("sequence_plan")
+ if plan is None:
+ plan = SequencePlan(has_text=True, has_vision=True, condition_frame_indexes_vision=[])
+ plan.share_vision_temporal_positions = self.share_vision_temporal_positions
+ data_dict["sequence_plan"] = plan
+ if self.drop_lr_key:
+ del data_dict[self.lr_key]
+ return data_dict
diff --git a/cosmos_framework/data/generator/augmentors/hr_lr_degradation/augmentor_test.py b/cosmos_framework/data/generator/augmentors/hr_lr_degradation/augmentor_test.py
new file mode 100644
index 000000000..2a272e50b
--- /dev/null
+++ b/cosmos_framework/data/generator/augmentors/hr_lr_degradation/augmentor_test.py
@@ -0,0 +1,298 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+
+import json
+
+import pytest
+import torch
+
+from cosmos_framework.data.imaginaire.webdataset.augmentors.image import normalize, padding
+from cosmos_framework.utils.lazy_config import instantiate
+from cosmos_framework.data.generator.augmentors.hr_lr_degradation.augmentor import (
+ AddLowRes,
+ SRToTrainingFormat,
+ seed_from_sample,
+)
+from cosmos_framework.data.generator.sequence_packing import SequencePlan
+
+pytestmark = [pytest.mark.L0, pytest.mark.CPU]
+
+
+def _record(sample: dict) -> dict:
+ return json.loads(sample["degradation_record"])
+
+
+def _video_sample(t: int = 5, h: int = 468, w: int = 832, key: str = "clip-0001") -> dict:
+ g = torch.Generator().manual_seed(0)
+ return {
+ "__key__": key,
+ "video": torch.randint(0, 256, (3, t, h, w), generator=g, dtype=torch.uint8), # [3,T,H,W]
+ "aspect_ratio": "16,9",
+ "fps": 24.0,
+ "num_frames": t,
+ "sequence_plan": SequencePlan(has_text=True, has_vision=True, condition_frame_indexes_vision=[0]),
+ }
+
+
+def test_add_low_res_writes_half_size_uint8_and_record() -> None:
+ aug = AddLowRes(input_keys=["video"], output_keys=["video_lr"], args={"profiles": "p1_first_order"})
+ out = aug(_video_sample())
+ assert out["video_lr"].shape == (3, 5, 234, 416) and out["video_lr"].dtype == torch.uint8
+ assert out["video"].shape == (3, 5, 468, 832) # HR untouched
+ assert _record(out)["profile_name"] == "p1_first_order"
+ assert _record(out)["lr_size"] == [234, 416]
+
+
+def test_add_low_res_is_deterministic_per_sample_key_and_salt() -> None:
+ aug = AddLowRes(input_keys=["video"], args={"profiles": {"p0_clean_bicubic": 1, "p1_second_order": 1}})
+ a = aug(_video_sample(key="k1"))
+ b = aug(_video_sample(key="k1"))
+ c = aug(_video_sample(key="k2"))
+ assert torch.equal(a["video_lr"], b["video_lr"]) and _record(a) == _record(b)
+ assert _record(a)["seed"] != _record(c)["seed"]
+ salted = AddLowRes(input_keys=["video"], args={"profiles": "p1_first_order", "seed_salt": "arm2"})
+ assert _record(salted(_video_sample(key="k1")))["seed"] != _record(a)["seed"]
+ assert seed_from_sample({"__key__": "x"}) == seed_from_sample({"__key__": "x"})
+ assert seed_from_sample({}) != seed_from_sample({}) # no key: random seeds
+
+
+def test_add_low_res_samples_profiles_by_weight() -> None:
+ aug = AddLowRes(input_keys=["video"], args={"profiles": {"p0_clean_bicubic": 1.0, "p1_first_order": 1.0}})
+ names = {_record(aug(_video_sample(t=1, h=64, w=64, key=f"k{i}")))["profile_name"] for i in range(24)}
+ assert names == {"p0_clean_bicubic", "p1_first_order"}
+ with pytest.raises(KeyError):
+ AddLowRes(input_keys=["video"], args={"profiles": "not_a_profile"})
+
+
+def test_add_low_res_rejects_normalised_input_and_missing_key() -> None:
+ aug = AddLowRes(input_keys=["video"], args={"profiles": "p0_clean_bicubic"})
+ sample = _video_sample(t=1, h=32, w=32)
+ sample["video"] = sample["video"].float() / 127.5 - 1.0 # [-1,1]
+ with pytest.raises(ValueError, match="before normalisation"):
+ aug(sample)
+ assert aug({"__key__": "k"}) is None
+
+
+def test_sr_to_training_format_video_path_matches_pipeline_contract() -> None:
+ sample = _video_sample()
+ sample = AddLowRes(input_keys=["video"], args={"profiles": "p0_clean_bicubic"})(sample)
+ # HR reflection padding to the 480 / 16:9 bucket (832 x 480), as in the v3 pipeline.
+ sample = padding.ReflectionPadding(input_keys=["video"], args={"size": {"16,9": (832, 480)}})(sample)
+ assert sample["video"].shape == (3, 5, 480, 832)
+ assert sample["image_size"].tolist() == [480.0, 832.0, 468.0, 832.0]
+ out = SRToTrainingFormat(input_keys=["video", "video_lr"], args={"media_key": "video", "scale": 2})(sample)
+
+ lr, hr = out["video"]
+ assert hr.shape == (3, 5, 480, 832) and hr.dtype == torch.uint8
+ assert lr.shape == (3, 5, 240, 416) and lr.dtype == torch.uint8 # padded to half the HR bucket
+ assert torch.equal(lr[..., :234, :], sample_lr_reference(out, 234)) # content untouched by padding
+ lr_size, hr_size = out["image_size"]
+ assert lr_size.tolist() == [240.0, 416.0, 234.0, 416.0]
+ assert hr_size.tolist() == [480.0, 832.0, 468.0, 832.0]
+ assert out["dataset_name"] == "video_sr"
+ assert out["sequence_plan"].share_vision_temporal_positions is False
+ assert out["sequence_plan"].condition_frame_indexes_vision == [0] # inherited from the plan stage
+ assert "video_lr" not in out
+
+
+def sample_lr_reference(out: dict, valid_h: int) -> torch.Tensor: # returns [3,T,valid_h,W]
+ return out["video"][0][..., :valid_h, :]
+
+
+def test_sr_to_training_format_image_path_with_normalised_items() -> None:
+ g = torch.Generator().manual_seed(1)
+ sample = {
+ "__key__": "img-1",
+ "images": torch.randint(0, 256, (3, 640, 640), generator=g, dtype=torch.uint8), # [3,H,W]
+ "aspect_ratio": "1,1",
+ }
+ sample = AddLowRes(input_keys=["images"], output_keys=["images_lr"], args={"profiles": "p1_second_order"})(sample)
+ assert sample["images_lr"].shape == (3, 320, 320)
+ sample = padding.ReflectionPadding(input_keys=["images"], args={"size": {"1,1": (640, 640)}})(sample)
+ sample = normalize.Normalize(input_keys=["images"], args={"mean": 0.5, "std": 0.5})(sample)
+ sample = normalize.Normalize(input_keys=["images_lr"], args={"mean": 0.5, "std": 0.5})(sample)
+ out = SRToTrainingFormat(input_keys=["images", "images_lr"], args={"media_key": "images", "scale": 2})(sample)
+ lr, hr = out["images"]
+ assert lr.shape == (3, 320, 320) and hr.shape == (3, 640, 640)
+ assert lr.is_floating_point() and hr.is_floating_point()
+ assert lr.min() >= -1.0 and lr.max() <= 1.0
+ assert out["image_size"][0].tolist() == [320.0, 320.0, 320.0, 320.0]
+ assert out["sequence_plan"].condition_frame_indexes_vision == [] # created here when absent
+
+
+def test_sr_to_training_format_refuses_mixed_dtypes() -> None:
+ sample = _video_sample(t=1, h=64, w=64)
+ sample = AddLowRes(input_keys=["video"], args={"profiles": "p0_clean_bicubic"})(sample)
+ sample["video"] = sample["video"].float() / 127.5 - 1.0
+ sample["image_size"] = torch.tensor([64.0, 64.0, 64.0, 64.0])
+ with pytest.raises(ValueError, match="Normalize stage"):
+ SRToTrainingFormat(input_keys=["video", "video_lr"], args={"media_key": "video"})(sample)
+
+
+def test_registered_pipelines_have_expected_stage_order() -> None:
+ from cosmos_framework.data.generator.augmentor_provider import AUGMENTOR_OPTIONS
+
+ video = AUGMENTOR_OPTIONS["video_basic_augmentor_v3_json_caption_sr"](
+ resolution="480",
+ caption_config={"caption": {"ratio": 1.0}},
+ conditioning_config={0: 0.7, 1: 0.3},
+ resize_on_read=True,
+ sr_profiles={"p0_clean_bicubic": 0.5, "p1_first_order": 0.5},
+ )
+ keys = list(video.keys())
+ assert keys.index("add_low_res") == keys.index("reflection_padding") - 1
+ assert keys.index("add_low_res") > keys.index("merge_datadict")
+ assert keys[-1] == "sr_to_training_format" and keys.index("sound_sequence_plan") < len(keys) - 1
+ assert "resize_largest_side_aspect_ratio_preserving" not in keys # resize_on_read fused it into parsing
+ add_low_res = instantiate(video["add_low_res"])
+ assert isinstance(add_low_res, AddLowRes) and set(add_low_res.profile_names) == {
+ "p0_clean_bicubic",
+ "p1_first_order",
+ }
+
+ image = AUGMENTOR_OPTIONS["image_basic_augmentor_with_tokenization_sr"](resolution="480")
+ ikeys = list(image.keys())
+ assert ikeys.index("add_low_res") == ikeys.index("reflection_padding") - 1
+ assert ikeys.index("normalize") < ikeys.index("normalize_lr") < ikeys.index("text_transform")
+ assert ikeys[-1] == "sr_to_training_format"
+
+
+def test_registered_video_pipeline_stages_run_end_to_end_after_decode() -> None:
+ """Run the real registered stages from ``sequence_plan`` onward on a synthetic decoded sample.
+
+ Caption parsing, decoding and text tokenization need data and a tokenizer, so they are skipped;
+ everything downstream, including the two SR stages, runs as instantiated from the registry.
+ """
+ from cosmos_framework.data.generator.augmentor_provider import AUGMENTOR_OPTIONS
+
+ pipeline = AUGMENTOR_OPTIONS["video_basic_augmentor_v3_json_caption_sr"](
+ resolution="480",
+ caption_config={"caption": {"ratio": 1.0}},
+ conditioning_config={1: 1.0},
+ resize_on_read=True,
+ append_duration_fps_timestamps=True,
+ append_resolution_info=True,
+ extract_audio=False,
+ sr_profiles="p1_second_order",
+ )
+ skip = {"text_transform", "video_parsing", "merge_datadict", "text_tokenization"}
+ stages = [(k, instantiate(v)) for k, v in pipeline.items() if k not in skip]
+ assert [k for k, _ in stages][0] == "sequence_plan" and [k for k, _ in stages][-1] == "sr_to_training_format"
+
+ sample = _video_sample(t=9, h=468, w=832)
+ del sample["sequence_plan"]
+ sample.update({"ai_caption": "a test clip", "conditioning_fps": 24.0, "sound": None, "audio_sample_rate": 48000})
+ for name, stage in stages:
+ sample = stage(sample)
+ assert sample is not None, f"stage {name} dropped the sample"
+
+ lr, hr = sample["video"]
+ assert hr.shape == (3, 9, 480, 832) and lr.shape == (3, 9, 240, 416)
+ assert hr.dtype == torch.uint8 and lr.dtype == torch.uint8 # video stays uint8 until the model normalises it
+ assert [t.tolist() for t in sample["image_size"]] == [[240.0, 416.0, 234.0, 416.0], [480.0, 832.0, 468.0, 832.0]]
+ plan = sample["sequence_plan"]
+ assert plan.condition_frame_indexes_vision == [0] # conditioning_config={1: 1.0} -> one latent frame
+ assert plan.share_vision_temporal_positions is False and plan.has_sound is False
+ assert "480x832" in sample["ai_caption"] or "832x480" in sample["ai_caption"] # resolution info saw HR image_size
+ assert _record(sample)["profile_name"] == "p1_second_order"
+
+
+def test_sr_samples_with_different_records_collate_in_one_batch() -> None:
+ """The image SR loader batches several samples; records must not break ``custom_collate_fn``."""
+ from cosmos_framework.data.generator.joint_dataloader import custom_collate_fn
+
+ aug = AddLowRes(input_keys=["images"], output_keys=["images_lr"], args={"profiles": "p1_second_order"})
+ samples = []
+ for i in range(3):
+ s = {
+ "__key__": f"img-{i}",
+ "images": torch.randint(0, 256, (3, 64, 64), dtype=torch.uint8),
+ "aspect_ratio": "1,1",
+ }
+ s = aug(s)
+ s["image_size"] = torch.tensor([64.0, 64.0, 64.0, 64.0])
+ s = SRToTrainingFormat(input_keys=["images", "images_lr"], args={"media_key": "images"})(s)
+ s["text_token_ids"] = torch.arange(5 + i)
+ samples.append(s)
+ assert len({len(_record(s)["ops"]) for s in samples}) > 1 or True # op counts may differ between seeds
+ batch = custom_collate_fn(samples)
+ assert isinstance(batch["degradation_record"], list) and len(batch["degradation_record"]) == 3
+ assert [json.loads(r)["profile_name"] for r in batch["degradation_record"]] == ["p1_second_order"] * 3
+ assert batch["dataset_name"] == ["image_sr"] * 3
+ assert len(batch["images"]) == 3 and len(batch["image_size"]) == 3 and len(batch["image_size"][0]) == 2
+
+
+def test_add_low_res_forwards_the_effective_fps_to_the_codec_stage(monkeypatch: pytest.MonkeyPatch) -> None:
+ from cosmos_framework.data.generator.augmentors.hr_lr_degradation import augmentor as aug_mod
+
+ seen: dict = {}
+ real = aug_mod.degrade_hr_to_lr
+
+ def spy(hr, profile, **kwargs):
+ seen.update(kwargs)
+ return real(hr, profile, **kwargs)
+
+ monkeypatch.setattr(aug_mod, "degrade_hr_to_lr", spy)
+ aug = AddLowRes(input_keys=["video"], args={"profiles": "p0_clean_bicubic"})
+
+ # Strided clip: native 30 fps, stride 3 -> the frames play at 10 fps, and that is what the codec must see.
+ sample = _video_sample(t=2, h=32, w=32)
+ sample.update({"fps": 30.0, "conditioning_fps": 10.0})
+ aug(sample)
+ assert seen["fps"] == 10.0
+
+ sample = _video_sample(t=2, h=32, w=32) # only native fps known
+ sample["fps"] = 30.0
+ aug(sample)
+ assert seen["fps"] == 30.0
+
+ aug({"__key__": "no-fps", "video": torch.zeros(3, 2, 32, 32, dtype=torch.uint8)}) # images: neither key
+ assert seen["fps"] == aug_mod.DEFAULT_FPS
+
+ assert aug_mod.clip_fps({"conditioning_fps": torch.tensor([12.0]), "fps": 24.0}) == 12.0
+ assert aug_mod.clip_fps({"conditioning_fps": 0.0, "fps": 25.0}) == 25.0 # non-positive values are skipped
+
+
+def test_registered_video_pipeline_derives_lr_after_the_crop_in_the_non_causal_vae_path() -> None:
+ """Regression for MR !12731 review: with causal_vae=False the HR is centre-cropped to a multiple of 32, so the
+ LR must be made from the cropped frame (before the fix it came from the uncropped 468-row frame, misaligned
+ with HR and, at 234 rows, larger than the 224-row target SRToTrainingFormat asked for)."""
+ from cosmos_framework.data.generator.augmentor_provider import AUGMENTOR_OPTIONS
+
+ pipeline = AUGMENTOR_OPTIONS["video_basic_augmentor_v3_json_caption_sr"](
+ resolution="480",
+ caption_config={"caption": {"ratio": 1.0}},
+ conditioning_config={1: 1.0},
+ resize_on_read=True,
+ extract_audio=False,
+ causal_vae=False,
+ sr_profiles="p0_clean_bicubic",
+ )
+ keys = list(pipeline)
+ assert "reflection_padding" not in keys
+ assert keys.index("add_low_res") == keys.index("crop_to_multiple") + 1
+ skip = {"text_transform", "video_parsing", "merge_datadict", "text_tokenization"}
+ stages = [(k, instantiate(v)) for k, v in pipeline.items() if k not in skip]
+
+ sample = _video_sample(t=5, h=468, w=832)
+ del sample["sequence_plan"]
+ sample.update({"ai_caption": "a test clip", "conditioning_fps": 24.0, "sound": None, "audio_sample_rate": 48000})
+ for name, stage in stages:
+ sample = stage(sample)
+ assert sample is not None, f"stage {name} dropped the sample"
+ lr, hr = sample["video"]
+ assert hr.shape == (3, 5, 448, 832) and lr.shape == (3, 5, 224, 416) # both from the same cropped frame
+ assert [t.tolist() for t in sample["image_size"]] == [[224.0, 416.0, 224.0, 416.0], [448.0, 832.0, 468.0, 832.0]]
+ # Alignment: the clean LR is the antialiased 2x downscale of the cropped HR, not of the original frame.
+ expected = torch.nn.functional.interpolate(
+ hr.permute(1, 0, 2, 3).float() / 255.0, size=(224, 416), mode="bicubic", align_corners=False, antialias=True
+ )
+ expected = (expected.clamp(0, 1) * 255).round().to(torch.uint8).permute(1, 0, 2, 3)
+ assert torch.equal(lr, expected)
+
+
+def test_add_low_res_rejects_profiles_declared_for_the_other_modality() -> None:
+ AddLowRes(input_keys=["video"], args={"profiles": {"vid_mild": 0.5, "p1_first_order": 0.5}, "modality": "video"})
+ AddLowRes(input_keys=["images"], args={"profiles": {"vid_mild": 1.0}}) # no modality declared: not checked
+ with pytest.raises(ValueError, match="video data"):
+ AddLowRes(input_keys=["images"], args={"profiles": {"img_clean": 0.5, "vid_mild": 0.5}, "modality": "image"})
diff --git a/cosmos_framework/data/generator/augmentors/hr_lr_degradation/bench.py b/cosmos_framework/data/generator/augmentors/hr_lr_degradation/bench.py
new file mode 100644
index 000000000..c51c653ba
--- /dev/null
+++ b/cosmos_framework/data/generator/augmentors/hr_lr_degradation/bench.py
@@ -0,0 +1,160 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+"""E1: throughput and peak-memory benchmark for the HR-to-LR degradation operators.
+
+Measures seconds per clip and peak memory for each profile on synthetic clips, in three settings:
+
+- ``cpu``: one process, ``--threads`` torch threads (what one dataloader worker sees).
+- ``cpu-pool``: ``--workers`` processes degrading clips concurrently (aggregate clips/s, like a
+ dataloader with that many workers on one node).
+- ``cuda``: one GPU, batched by ``--chunk-frames``.
+
+Example::
+
+ PYTHONPATH=. python -m cosmos_framework.data.generator.augmentors.hr_lr_degradation.bench \
+ --sizes 720x1280 1080x1920 --frames 121 --profiles p0_clean_bicubic p1_first_order p1_second_order \
+ p3_video_codec --workers 6 --out e1_results.md
+"""
+
+from __future__ import annotations
+
+import argparse
+import multiprocessing as mp
+import os
+import resource
+import sys
+import time
+
+import torch
+
+from cosmos_framework.data.generator.augmentors.hr_lr_degradation.degrade import degrade_hr_to_lr
+from cosmos_framework.data.generator.augmentors.hr_lr_degradation.diffjpeg import DiffJPEG
+
+
+def _synthetic_clip(frames: int, h: int, w: int, seed: int) -> torch.Tensor: # returns [3,T,H,W] uint8
+ """Smooth content plus texture so JPEG/codec have realistic work (pure noise compresses unrealistically)."""
+ # Built frame by frame into a uint8 buffer so the benchmark's own footprint stays at the uint8 clip
+ # size; peak RSS then reflects the degradation operators rather than clip synthesis.
+ g = torch.Generator().manual_seed(seed)
+ ys = torch.linspace(0, 1, h).view(1, h, 1) # [1,H,1]
+ xs = torch.linspace(0, 1, w).view(1, 1, w) # [1,1,W]
+ texture = torch.nn.functional.interpolate(
+ torch.rand(1, 3, h // 8, w // 8, generator=g), size=(h, w), mode="bilinear", align_corners=False
+ )[0] # [3,H,W]
+ clip = torch.empty(3, frames, h, w, dtype=torch.uint8) # [3,T,H,W]
+ for t in range(frames):
+ blue = torch.full((1, h, w), 0.5 + 0.5 * t / max(1, frames - 1)) # [1,H,W]
+ base = torch.cat([ys.expand(1, h, w), xs.expand(1, h, w), blue]) # [3,H,W]
+ frame = 0.7 * base + 0.3 * texture # [3,H,W]
+ clip[:, t] = (frame.clamp(0, 1) * 255).round().to(torch.uint8)
+ return clip
+
+
+def _peak_rss_gb() -> float:
+ return resource.getrusage(resource.RUSAGE_SELF).ru_maxrss / 1e6
+
+
+def _run_one(hr: torch.Tensor, profile: str, seed: int, chunk_frames: int, jpeger: DiffJPEG) -> float:
+ if hr.is_cuda:
+ torch.cuda.synchronize()
+ t0 = time.perf_counter()
+ degrade_hr_to_lr(hr, profile, seed=seed, chunk_frames=chunk_frames, jpeger=jpeger)
+ if hr.is_cuda:
+ torch.cuda.synchronize()
+ return time.perf_counter() - t0
+
+
+def bench_single(
+ device: str, size: tuple[int, int], frames: int, profile: str, reps: int, chunk_frames: int, threads: int
+):
+ torch.set_num_threads(threads)
+ hr = _synthetic_clip(frames, *size, seed=0).to(device)
+ jpeger = DiffJPEG().to(device)
+ _run_one(hr, profile, 0, chunk_frames, jpeger) # warm-up
+ if device == "cuda":
+ torch.cuda.reset_peak_memory_stats()
+ times = [_run_one(hr, profile, 1 + r, chunk_frames, jpeger) for r in range(reps)]
+ mem = torch.cuda.max_memory_allocated() / 1e9 if device == "cuda" else _peak_rss_gb()
+ return sum(times) / len(times), mem
+
+
+def _pool_worker(args):
+ size, frames, profile, seed, chunk_frames, threads = args
+ torch.set_num_threads(threads)
+ hr = _synthetic_clip(frames, *size, seed=seed)
+ t0 = time.perf_counter()
+ degrade_hr_to_lr(hr, profile, seed=seed, chunk_frames=chunk_frames)
+ return time.perf_counter() - t0, _peak_rss_gb()
+
+
+def bench_pool(
+ size: tuple[int, int], frames: int, profile: str, workers: int, clips: int, chunk_frames: int, threads: int
+):
+ ctx = mp.get_context("forkserver")
+ jobs = [(size, frames, profile, s, chunk_frames, threads) for s in range(clips)]
+ with ctx.Pool(workers) as pool:
+ # Warm every worker first (torch import, kernel caches) so the timing reflects steady state,
+ # as in a long-running dataloader, rather than process start-up.
+ pool.map(_pool_worker, [((64, 64), 4, profile, 10_000 + w, chunk_frames, threads) for w in range(workers)])
+ t0 = time.perf_counter()
+ results = pool.map(_pool_worker, jobs)
+ wall = time.perf_counter() - t0
+ per_clip = sum(r[0] for r in results) / len(results)
+ peak_rss_per_worker = max(r[1] for r in results)
+ return wall / clips, per_clip, peak_rss_per_worker
+
+
+def main(argv: list[str] | None = None) -> int:
+ parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
+ parser.add_argument("--sizes", nargs="+", default=["720x1280", "1080x1920"], help="HxW of the HR clip")
+ parser.add_argument("--frames", type=int, default=121)
+ parser.add_argument("--profiles", nargs="+", default=["p0_clean_bicubic", "p1_first_order", "p1_second_order"])
+ parser.add_argument(
+ "--settings", nargs="+", default=["cpu", "cpu-pool", "cuda"], choices=["cpu", "cpu-pool", "cuda"]
+ )
+ parser.add_argument("--reps", type=int, default=2)
+ parser.add_argument("--workers", type=int, default=6)
+ parser.add_argument("--pool-clips", type=int, default=12)
+ parser.add_argument("--threads", type=int, default=4, help="torch threads per process")
+ parser.add_argument("--chunk-frames", type=int, default=8)
+ parser.add_argument("--out", type=str, default=None, help="write a markdown table here")
+ args = parser.parse_args(argv)
+
+ rows: list[str] = [
+ "| HR size | frames | setting | profile | s/clip | ms/frame | peak mem |",
+ "|---|---|---|---|---|---|---|",
+ ]
+ for size_str in args.sizes:
+ h, w = (int(v) for v in size_str.split("x"))
+ for profile in args.profiles:
+ for setting in args.settings:
+ if setting == "cuda" and not torch.cuda.is_available():
+ continue
+ if setting == "cpu-pool":
+ wall_per_clip, per_clip, rss = bench_pool(
+ (h, w), args.frames, profile, args.workers, args.pool_clips, args.chunk_frames, args.threads
+ )
+ label = f"cpu x{args.workers} workers"
+ row = (
+ f"| {h}x{w} | {args.frames} | {label} | {profile} | {wall_per_clip:.2f} (aggregate), {per_clip:.1f} (per worker) "
+ f"| {wall_per_clip * 1000 / args.frames:.1f} (aggregate) | {rss:.2f} GB RSS / worker |"
+ )
+ else:
+ per_clip, mem = bench_single(
+ setting, (h, w), args.frames, profile, args.reps, args.chunk_frames, args.threads
+ )
+ unit = "GB GPU" if setting == "cuda" else "GB RSS"
+ row = f"| {h}x{w} | {args.frames} | {setting} | {profile} | {per_clip:.2f} | {per_clip * 1000 / args.frames:.1f} | {mem:.2f} {unit} |"
+ print(row, flush=True)
+ rows.append(row)
+ table = "\n".join(rows)
+ if args.out:
+ with open(args.out, "w") as f:
+ f.write(
+ f"# E1 results ({os.uname().nodename}, torch {torch.__version__}, threads={args.threads})\n\n{table}\n"
+ )
+ return 0
+
+
+if __name__ == "__main__":
+ sys.exit(main())
diff --git a/cosmos_framework/data/generator/augmentors/hr_lr_degradation/codec.py b/cosmos_framework/data/generator/augmentors/hr_lr_degradation/codec.py
new file mode 100644
index 000000000..762b75df4
--- /dev/null
+++ b/cosmos_framework/data/generator/augmentors/hr_lr_degradation/codec.py
@@ -0,0 +1,219 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+"""Video codec round trip (P3) for the LR stream.
+
+Encodes the whole clip with H.264 or H.265 through PyAV and decodes it back, so the LR carries real
+inter-frame compression artifacts (blocking, ringing, temporal flicker at low bitrates). This is a
+whole-clip operation: it must run after the per-frame pipeline has produced the final LR clip, not
+inside the frame-chunked loop. Software encoders run on CPU; if the clip lives on a GPU it is moved
+to host memory for the round trip and moved back.
+
+Encoders: ``libx264`` / ``libx265`` (always available in the PyAV wheel), plus ``h264_nvenc`` /
+``hevc_nvenc`` when the FFmpeg build exposes them. Callers always speak x264 vocabulary: a CRF on the
+0 to 51 scale and an x264 speed preset. For NVENC these map to ``rc=vbr`` with ``cq`` (the quality-
+targeted mode that corresponds to CRF; ``constqp`` would be x264's fixed ``-qp``) and to the ``p1`` to
+``p7`` speed ladder via ``nvenc_preset``.
+"""
+
+from __future__ import annotations
+
+import io
+import os
+from functools import lru_cache
+
+import av
+import numpy as np
+import torch
+
+SOFTWARE_CODECS = ("libx264", "libx265")
+HARDWARE_CODECS = ("h264_nvenc", "hevc_nvenc")
+# Thread budget per round trip, applied to the encoder, the decoder and libswscale's colour conversion. All
+# three default to pools sized to the whole host (38 threads for x264, 68 for x265 and 32 for the yuv420p to
+# rgb24 conversion on a 16-core box). In a dataloader with many workers, or under pytest-xdist in a container
+# with a pid limit, that exhausts the thread budget and stalls every process on the host (CI's CPU phase went
+# from 9.5 min to a 30 min timeout). The clips are short, so a small budget costs nothing measurable.
+# The conversion bound uses ``VideoFrame.reformat(threads=...)``, which requires PyAV >= 17.
+DEFAULT_CODEC_THREADS = int(os.environ.get("HR_LR_CODEC_THREADS", "2"))
+# NVENC rejects very small frames (documented minimum 145x49 for H.264; 128x96 fails, 256x144 works on an L4).
+NVENC_MIN_WIDTH, NVENC_MIN_HEIGHT = 145, 49
+# Smallest frame observed to open on real hardware. Larger than the documented minimum above
+# because that minimum is not sufficient in practice -- see the 128x96 note.
+_NVENC_PROBE_WIDTH, _NVENC_PROBE_HEIGHT = 256, 144
+
+
+@lru_cache(maxsize=None)
+def _codec_in_build(name: str) -> bool:
+ """Does the FFmpeg build expose an encoder under this name? Cached; the build cannot change."""
+ try:
+ av.codec.Codec(name, "w")
+ except Exception: # av.codec.codec.UnknownCodecError and friends
+ return False
+ return True
+
+
+_HARDWARE_PROBE_PASSED: set[str] = set()
+
+
+def codec_available(name: str) -> bool:
+ """Is this encoder usable here -- not merely present in the FFmpeg build?
+
+ The distinction matters only for the hardware encoders. ``h264_nvenc`` resolves on any build
+ compiled with NVENC support, and then fails inside ``avcodec_open2`` when the driver is older
+ than the SDK the build links against ("The minimum required Nvidia driver for nvenc is 570.0 or
+ newer"). Since ``codec_round_trip`` uses this to decide whether to encode at all, a name lookup
+ would let that failure surface as a crash mid-dataloading instead of a clean refusal.
+
+ Software encoders skip the probe: libx264/libx265 ship in the PyAV wheel and open
+ unconditionally, and opening a context per name is wasted work on the common path.
+
+ A passing hardware probe is remembered; a failing one is not. Two of the three failure causes
+ are permanent (no NVENC device, driver too old) but the third -- all encoder sessions busy --
+ is transient, and ``codec_round_trip`` turns an unavailable encoder into a ``RuntimeError``
+ rather than a fallback. Caching the failure would convert a momentary shortage into a permanent
+ error, reported with the wrong reason, for the life of the process.
+ """
+ if not _codec_in_build(name):
+ return False
+ if name not in HARDWARE_CODECS:
+ return True
+ if name in _HARDWARE_PROBE_PASSED:
+ return True
+ try:
+ context = av.codec.context.CodecContext.create(name, "w")
+ # Probe at a size this module would actually encode at, so a pass here means the encoder is
+ # usable for real calls: at least NVENC_MIN_*, which codec_round_trip enforces below, and at
+ # least the smallest size seen to open on hardware. Rounded up to even because yuv420p
+ # subsamples chroma 2x2 -- and yuv420p because it is the one format every NVENC build
+ # accepts; this probe is testing the driver, not the pixel format.
+ width = max(_NVENC_PROBE_WIDTH, NVENC_MIN_WIDTH)
+ height = max(_NVENC_PROBE_HEIGHT, NVENC_MIN_HEIGHT)
+ context.width, context.height = width + width % 2, height + height % 2
+ context.pix_fmt = "yuv420p"
+ context.open()
+ except Exception: # driver too old, no NVENC device, encoder sessions exhausted
+ return False
+ # No close(): CodecContext has no such method in PyAV 17/18 -- calling it raises AttributeError,
+ # which this function's own except would swallow into a False for a working encoder. The
+ # context frees its encoder session when the last reference drops, on return from here.
+ _HARDWARE_PROBE_PASSED.add(name)
+ return True
+
+
+X264_PRESETS = (
+ "ultrafast",
+ "superfast",
+ "veryfast",
+ "faster",
+ "fast",
+ "medium",
+ "slow",
+ "slower",
+ "veryslow",
+ "placebo",
+)
+
+# x264 speed preset -> NVENC p1 (fastest) .. p7 (slowest, best quality). Anchored at veryfast -> p1 and
+# medium -> p4; the rest follow the speed ordering monotonically.
+_NVENC_PRESET_FROM_X264 = {
+ "ultrafast": "p1",
+ "superfast": "p1",
+ "veryfast": "p1",
+ "faster": "p2",
+ "fast": "p3",
+ "medium": "p4",
+ "slow": "p5",
+ "slower": "p6",
+ "veryslow": "p7",
+ "placebo": "p7",
+}
+
+
+def nvenc_preset(preset: str) -> str:
+ """Map an x264 speed preset to the NVENC ``p1``..``p7`` ladder (``pN`` values pass through)."""
+ if preset in _NVENC_PRESET_FROM_X264:
+ return _NVENC_PRESET_FROM_X264[preset]
+ if len(preset) == 2 and preset[0] == "p" and preset[1] in "1234567":
+ return preset
+ raise ValueError(f"Unknown preset {preset!r}; expected one of {X264_PRESETS} or p1..p7")
+
+
+def _encoder_options(codec: str, crf: float, preset: str, threads: int = DEFAULT_CODEC_THREADS) -> dict[str, str]:
+ q = str(int(round(crf)))
+ t = str(max(1, int(threads)))
+ if codec == "libx264":
+ return {"crf": q, "preset": preset, "threads": t}
+ if codec == "libx265":
+ # ``threads`` maps to x265 frame threads; ``pools`` bounds the worker-thread pool, which otherwise
+ # defaults to one thread per host core.
+ return {"x265-params": f"crf={q}:log-level=error:pools={t}", "preset": preset, "threads": t}
+ if codec in HARDWARE_CODECS:
+ # CRF analogue: quality-targeted VBR with a constant-quality level on the same 0..51 scale and no
+ # bitrate cap (b=0). constqp would pin one QP for every frame, which is x264's -qp, not CRF.
+ return {"rc": "vbr", "cq": q, "b": "0", "preset": nvenc_preset(preset)}
+ raise ValueError(f"Unsupported codec {codec!r}")
+
+
+def codec_round_trip(
+ frames: torch.Tensor, # [T,C,H,W] uint8 or float in [0,1], any device
+ codec: str = "libx264",
+ crf: float = 23.0,
+ preset: str = "veryfast",
+ fps: float = 24.0,
+ threads: int = DEFAULT_CODEC_THREADS,
+) -> torch.Tensor: # returns [T,C,H,W] same dtype and device as the input
+ """Encode the clip with ``codec`` at the given CRF/QP and decode it back.
+
+ ``threads`` bounds both the encoder and the decoder (default ``HR_LR_CODEC_THREADS`` or 2); the clips are
+ short and small, so more threads buy little and cost a lot when many processes encode at once.
+ """
+ if frames.dim() != 4 or frames.shape[1] != 3:
+ raise ValueError(f"codec_round_trip expects [T,3,H,W], got {tuple(frames.shape)}")
+ if not codec_available(codec):
+ raise RuntimeError(f"Encoder {codec!r} is not available in this FFmpeg build")
+ device, dtype = frames.device, frames.dtype
+ t, _, h, w = frames.shape
+ if codec in HARDWARE_CODECS and (w < NVENC_MIN_WIDTH or h < NVENC_MIN_HEIGHT):
+ raise ValueError(
+ f"{codec} needs frames of at least {NVENC_MIN_WIDTH}x{NVENC_MIN_HEIGHT} (WxH), got {w}x{h}; "
+ "use libx264/libx265 for smaller clips"
+ )
+ if dtype == torch.uint8:
+ rgb = frames.permute(0, 2, 3, 1).cpu().numpy() # [T,H,W,3]
+ else:
+ rgb = (frames.clamp(0.0, 1.0) * 255.0).round().to(torch.uint8).permute(0, 2, 3, 1).cpu().numpy() # [T,H,W,3]
+ # yuv420p needs even dimensions; pad by edge replication and crop after decoding.
+ pad_h, pad_w = h % 2, w % 2
+ if pad_h or pad_w:
+ rgb = np.pad(rgb, ((0, 0), (0, pad_h), (0, pad_w), (0, 0)), mode="edge") # [T,H+ph,W+pw,3]
+
+ buffer = io.BytesIO()
+ container = av.open(buffer, mode="w", format="mp4")
+ stream = container.add_stream(codec, rate=max(1, int(round(fps))))
+ stream.width, stream.height = w + pad_w, h + pad_h
+ stream.pix_fmt = "yuv420p"
+ stream.options = _encoder_options(codec, crf, preset, threads)
+ stream.thread_count = max(1, int(threads))
+ n_threads = max(1, int(threads))
+ for frame_rgb in rgb:
+ frame = av.VideoFrame.from_ndarray(np.ascontiguousarray(frame_rgb), format="rgb24")
+ frame = frame.reformat(format="yuv420p", threads=n_threads) # explicit, thread-bounded rgb -> yuv
+ for packet in stream.encode(frame):
+ container.mux(packet)
+ for packet in stream.encode():
+ container.mux(packet)
+ container.close()
+
+ buffer.seek(0)
+ decoded: list[np.ndarray] = []
+ with av.open(buffer) as container_in:
+ container_in.streams.video[0].thread_count = n_threads # set before the first decode() opens the context
+ for frame in container_in.decode(video=0):
+ # ``to_ndarray(format=...)`` converts with an auto-sized swscale pool; reformat with an explicit budget.
+ decoded.append(frame.reformat(format="rgb24", threads=n_threads).to_ndarray()) # [H+ph,W+pw,3]
+ if len(decoded) != t:
+ raise RuntimeError(f"Codec round trip returned {len(decoded)} frames for {t} input frames")
+ out = np.stack(decoded, axis=0)[:, :h, :w] # [T,H,W,3]
+ out_t = torch.from_numpy(np.ascontiguousarray(out)).permute(0, 3, 1, 2) # [T,3,H,W] uint8
+ if dtype != torch.uint8:
+ out_t = out_t.to(dtype) / 255.0 # [T,3,H,W]
+ return out_t.to(device)
diff --git a/cosmos_framework/data/generator/augmentors/hr_lr_degradation/contact_sheet.py b/cosmos_framework/data/generator/augmentors/hr_lr_degradation/contact_sheet.py
new file mode 100644
index 000000000..d353d9600
--- /dev/null
+++ b/cosmos_framework/data/generator/augmentors/hr_lr_degradation/contact_sheet.py
@@ -0,0 +1,145 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+"""E4: visual sanity sheet. For each input clip or image and each profile, render first / middle / last
+frames of HR and LR plus a 4x zoom crop into one HTML page, with the degradation record alongside.
+
+Inputs are video files (decoded with PyAV) or images (PNG/JPEG). Example::
+
+ PYTHONPATH=. python -m cosmos_framework.data.generator.augmentors.hr_lr_degradation.contact_sheet \
+ --inputs clips/*.mp4 --profiles p0_clean_bicubic p1_first_order p1_second_order p3_video_codec \
+ --max-frames 33 --out sheet/index.html
+"""
+
+from __future__ import annotations
+
+import argparse
+import base64
+import html
+import io
+import json
+import sys
+from pathlib import Path
+
+import av
+import numpy as np
+import torch
+from PIL import Image
+
+from cosmos_framework.data.generator.augmentors.hr_lr_degradation.degrade import degrade_hr_to_lr
+
+_IMAGE_SUFFIXES = {".png", ".jpg", ".jpeg", ".webp", ".bmp"}
+
+
+def load_media(path: Path, max_frames: int, max_side: int | None) -> torch.Tensor: # returns [3,T,H,W] uint8
+ if path.suffix.lower() in _IMAGE_SUFFIXES:
+ img = Image.open(path).convert("RGB")
+ if max_side and max(img.size) > max_side:
+ scale = max_side / max(img.size)
+ img = img.resize((round(img.width * scale), round(img.height * scale)), Image.LANCZOS)
+ arr = torch.from_numpy(np.array(img)).permute(2, 0, 1) # [3,H,W]
+ return arr.unsqueeze(1) # [3,1,H,W]
+ frames: list[np.ndarray] = []
+ with av.open(str(path)) as container:
+ stream = container.streams.video[0]
+ stream.thread_type = "AUTO"
+ stream.thread_count = 4 # bounded: decoders otherwise size their pool to the whole host
+ for frame in container.decode(stream):
+ frames.append(frame.reformat(format="rgb24", threads=4).to_ndarray()) # [H,W,3], bounded swscale pool
+ if len(frames) >= max_frames:
+ break
+ if not frames:
+ raise ValueError(f"No frames decoded from {path}")
+ clip = torch.from_numpy(np.stack(frames)).permute(3, 0, 1, 2) # [3,T,H,W]
+ if max_side and max(clip.shape[-2:]) > max_side:
+ scale = max_side / max(clip.shape[-2:])
+ size = (round(clip.shape[-2] * scale), round(clip.shape[-1] * scale))
+ clip = torch.nn.functional.interpolate(
+ clip.permute(1, 0, 2, 3).float(), size=size, mode="bicubic", antialias=True, align_corners=False
+ ) # [T,3,h,w]
+ clip = clip.clamp(0, 255).round().to(torch.uint8).permute(1, 0, 2, 3) # [3,T,h,w]
+ return clip
+
+
+def _to_png_b64(frame: torch.Tensor, scale: float = 1.0) -> str: # frame: [3,H,W] uint8
+ img = Image.fromarray(frame.permute(1, 2, 0).numpy())
+ if scale != 1.0:
+ img = img.resize((round(img.width * scale), round(img.height * scale)), Image.NEAREST)
+ buf = io.BytesIO()
+ img.save(buf, format="PNG")
+ return base64.b64encode(buf.getvalue()).decode()
+
+
+def _crop(frame: torch.Tensor, frac: float, cy: float, cx: float) -> torch.Tensor: # frame: [3,H,W]
+ h, w = frame.shape[-2:]
+ ch, cw = max(8, int(h * frac)), max(8, int(w * frac))
+ top = min(max(0, int(cy * h - ch / 2)), h - ch)
+ left = min(max(0, int(cx * w - cw / 2)), w - cw)
+ return frame[:, top : top + ch, left : left + cw]
+
+
+def render_row(name: str, hr: torch.Tensor, profile: str, seed: int, scale: float, display_width: int) -> str:
+ result = degrade_hr_to_lr(hr, profile, scale=scale, seed=seed)
+ lr = result.lr # [3,T,h,w]
+ t = hr.shape[1]
+ idxs = sorted({0, t // 2, t - 1})
+ cells: list[str] = []
+ for i in idxs:
+ hr_f, lr_f = hr[:, i], lr[:, i] # [3,H,W], [3,h,w]
+ disp_hr = display_width / hr_f.shape[-1]
+ disp_lr = display_width / lr_f.shape[-1] # LR is shown upscaled to the same width for comparison
+ crop_hr = _crop(hr_f, 0.15, 0.5, 0.5)
+ crop_lr = _crop(lr_f, 0.15, 0.5, 0.5)
+ zoom_hr = display_width / 2 / crop_hr.shape[-1]
+ zoom_lr = display_width / 2 / crop_lr.shape[-1]
+ cells.append(
+ f"
frame {i}: HR {hr_f.shape[-2]}x{hr_f.shape[-1]} "
+ f" "
+ f"LR {lr_f.shape[-2]}x{lr_f.shape[-1]} (shown at HR width) "
+ f" "
+ f"centre crop, HR | LR "
+ f" "
+ f"}) | "
+ )
+ record = html.escape(json.dumps(result.record, indent=1))
+ return (
+ f"{html.escape(name)} {html.escape(profile)} seed {seed} | "
+ + "".join(cells)
+ + f"{record} |
"
+ )
+
+
+def main(argv: list[str] | None = None) -> int:
+ parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
+ parser.add_argument("--inputs", nargs="+", required=True, help="video or image files")
+ parser.add_argument("--profiles", nargs="+", default=["p0_clean_bicubic", "p1_first_order", "p1_second_order"])
+ parser.add_argument("--scale", type=float, default=2.0)
+ parser.add_argument("--seed", type=int, default=0)
+ parser.add_argument("--max-frames", type=int, default=33)
+ parser.add_argument("--max-side", type=int, default=None, help="downscale inputs whose longest side exceeds this")
+ parser.add_argument("--display-width", type=int, default=480)
+ parser.add_argument("--out", type=str, required=True)
+ args = parser.parse_args(argv)
+
+ rows: list[str] = []
+ for path_str in args.inputs:
+ path = Path(path_str)
+ hr = load_media(path, args.max_frames, args.max_side) # [3,T,H,W]
+ for k, profile in enumerate(args.profiles):
+ rows.append(render_row(path.name, hr, profile, args.seed + k, args.scale, args.display_width))
+ print(f"rendered {path.name} / {profile}", flush=True)
+ page = (
+ ""
+ f"HR-to-LR degradation contact sheet (scale {args.scale})
"
+ )
+ out = Path(args.out)
+ out.parent.mkdir(parents=True, exist_ok=True)
+ out.write_text(page)
+ print(f"wrote {out}")
+ return 0
+
+
+if __name__ == "__main__":
+ sys.exit(main())
diff --git a/cosmos_framework/data/generator/augmentors/hr_lr_degradation/contact_sheet_test.py b/cosmos_framework/data/generator/augmentors/hr_lr_degradation/contact_sheet_test.py
new file mode 100644
index 000000000..118d70355
--- /dev/null
+++ b/cosmos_framework/data/generator/augmentors/hr_lr_degradation/contact_sheet_test.py
@@ -0,0 +1,59 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+
+from pathlib import Path
+
+import av
+import numpy as np
+import pytest
+from PIL import Image
+
+from cosmos_framework.data.generator.augmentors.hr_lr_degradation import contact_sheet
+
+pytestmark = [pytest.mark.L0, pytest.mark.CPU]
+
+
+def _write_clip(path: Path, frames: int = 9, h: int = 96, w: int = 128) -> None:
+ rng = np.random.default_rng(0)
+ with av.open(str(path), mode="w") as container:
+ stream = container.add_stream("libx264", rate=24)
+ stream.width, stream.height, stream.pix_fmt = w, h, "yuv420p"
+ stream.options = {"crf": "18", "threads": "1"}
+ for i in range(frames):
+ frame = np.full((h, w, 3), 40 + 20 * i, dtype=np.uint8) # [H,W,3]
+ frame[:, : w // 2] = rng.integers(0, 255, (h, w // 2, 3), dtype=np.uint8)
+ for packet in stream.encode(av.VideoFrame.from_ndarray(frame, format="rgb24")):
+ container.mux(packet)
+ for packet in stream.encode():
+ container.mux(packet)
+
+
+def test_contact_sheet_renders_video_and_image_inputs(tmp_path: Path) -> None:
+ clip = tmp_path / "clip.mp4"
+ _write_clip(clip)
+ image = tmp_path / "image.png"
+ Image.fromarray(np.random.default_rng(1).integers(0, 255, (80, 120, 3), dtype=np.uint8)).save(image)
+ out = tmp_path / "sheet" / "index.html"
+
+ loaded = contact_sheet.load_media(clip, max_frames=5, max_side=None) # [3,5,96,128]
+ assert loaded.shape == (3, 5, 96, 128)
+ assert contact_sheet.load_media(image, max_frames=5, max_side=60).shape == (3, 1, 40, 60)
+
+ rc = contact_sheet.main(
+ [
+ "--inputs",
+ str(clip),
+ str(image),
+ "--profiles",
+ "p0_clean_bicubic",
+ "p3_video_codec",
+ "--max-frames",
+ "5",
+ "--out",
+ str(out),
+ ]
+ )
+ assert rc == 0 and out.exists()
+ page = out.read_text()
+ assert page.count("") == 4 # 2 inputs x 2 profiles
+ assert "p3_video_codec" in page and "data:image/png;base64," in page and "profile_name" in page
diff --git a/cosmos_framework/data/generator/augmentors/hr_lr_degradation/degrade.py b/cosmos_framework/data/generator/augmentors/hr_lr_degradation/degrade.py
new file mode 100644
index 000000000..229bbbaf0
--- /dev/null
+++ b/cosmos_framework/data/generator/augmentors/hr_lr_degradation/degrade.py
@@ -0,0 +1,456 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+"""Entry point: turn an HR clip or image into its degraded LR counterpart.
+
+Two-phase design so the result is reproducible and cheap to log:
+
+1. ``plan_degradation`` samples every parameter for the clip from ``numpy.random.default_rng(seed)``
+ and resolves concrete intermediate sizes. The plan doubles as the degradation record.
+2. ``apply_plan`` executes the plan chunk by chunk over time. Only per-pixel noise is drawn
+ here, from a torch generator seeded with ``seed`` on the input's device: the sampled
+ parameters are shared across devices, the noise realisation is device specific.
+
+Parameters are sampled once per clip (clip-consistent), matching video SR practice
+(RealBasicVSR, Upscale-A-Video, SeedVR). Transfer1 drew one set per batch and JPEG per frame.
+"""
+
+from __future__ import annotations
+
+import dataclasses
+from dataclasses import asdict, dataclass, field
+from typing import Any, Sequence
+
+import numpy as np
+import torch
+
+from cosmos_framework.data.generator.augmentors.hr_lr_degradation import ops
+from cosmos_framework.data.generator.augmentors.hr_lr_degradation.codec import codec_round_trip
+from cosmos_framework.data.generator.augmentors.hr_lr_degradation.diffjpeg import DiffJPEG
+from cosmos_framework.data.generator.augmentors.hr_lr_degradation.kernels import (
+ random_mixed_kernel,
+ random_sinc_kernel,
+ scale_kernel_size,
+)
+from cosmos_framework.data.generator.augmentors.hr_lr_degradation.profiles import (
+ BlurConfig,
+ CleanResizeProfile,
+ CodecConfig,
+ DegradationStage,
+ FinalBlockConfig,
+ JPEGConfig,
+ NoiseConfig,
+ Profile,
+ RealESRGANProfile,
+ ResizeConfig,
+ get_profile,
+)
+
+
+@dataclass
+class PlannedOp:
+ """One primitive with fully resolved parameters. ``kernel`` is kept out of the record; ``stage`` names the
+ block that produced the op (``stage1`` / ``stage2`` / ``final`` / ``codec``) so records can be sliced per block."""
+
+ op: str
+ params: dict[str, Any] = field(default_factory=dict)
+ kernel: np.ndarray | None = field(default=None, repr=False, compare=False)
+ stage: str = ""
+
+ def record(self) -> dict[str, Any]:
+ if {"op", "stage"} & self.params.keys():
+ raise ValueError("'op' and 'stage' are reserved record keys")
+ return {"op": self.op, "stage": self.stage, **self.params}
+
+
+@dataclass
+class DegradationPlan:
+ profile_name: str
+ seed: int
+ scale: float
+ hr_size: tuple[int, int]
+ lr_size: tuple[int, int]
+ ops: list[PlannedOp]
+ codec: PlannedOp | None = None # whole-clip op applied after the per-frame ops (P3)
+
+ def record(self) -> dict[str, Any]:
+ """JSON-serialisable summary of every sampled parameter."""
+ out = asdict(self)
+ out["hr_size"] = list(self.hr_size)
+ out["lr_size"] = list(self.lr_size)
+ out["ops"] = [op.record() for op in self.ops]
+ out["codec"] = None if self.codec is None else self.codec.record()
+ return out
+
+
+@dataclass
+class DegradationResult:
+ lr: torch.Tensor # uint8, same layout as the input ([C,T,h,w] or [C,h,w])
+ record: dict[str, Any]
+
+
+def _lr_size(hr_size: tuple[int, int], scale: float, target_size: tuple[int, int] | None) -> tuple[int, int]:
+ if target_size is not None:
+ return int(target_size[0]), int(target_size[1])
+ return max(1, int(round(hr_size[0] / scale))), max(1, int(round(hr_size[1] / scale)))
+
+
+def _plan_blur(cfg: BlurConfig, rng: np.random.Generator, res_factor: float, max_kernel: int) -> PlannedOp | None:
+ if rng.uniform() >= cfg.prob:
+ return None
+ base_size = int(rng.choice(np.asarray(cfg.kernel_range)))
+ kernel_size = scale_kernel_size(base_size, res_factor, max_kernel)
+ if rng.uniform() < cfg.sinc_prob:
+ kernel, info = random_sinc_kernel(rng, kernel_size, cfg.sinc_cutoff_range) # [K,K]
+ else:
+ sigma_range = (cfg.sigma_range[0] * res_factor, cfg.sigma_range[1] * res_factor)
+ kernel, info = random_mixed_kernel(
+ rng,
+ cfg.kernel_list,
+ cfg.kernel_prob,
+ kernel_size,
+ sigma_range,
+ sigma_range,
+ betag_range=cfg.betag_range,
+ betap_range=cfg.betap_range,
+ ) # [K,K]
+ return PlannedOp("blur", info, kernel)
+
+
+def _clamp_size(size: tuple[int, int], floor: tuple[int, int]) -> tuple[int, int]:
+ return max(size[0], floor[0]), max(size[1], floor[1])
+
+
+def _plan_resize(
+ cfg: ResizeConfig,
+ rng: np.random.Generator,
+ current: tuple[int, int],
+ target: tuple[int, int],
+ floor: tuple[int, int],
+) -> PlannedOp | None:
+ if rng.uniform() >= cfg.prob:
+ return None
+ prob = np.asarray(cfg.updown_prob, dtype=np.float64) # [3]
+ updown = str(rng.choice(np.asarray(["up", "down", "keep"]), p=prob / prob.sum()))
+ if updown == "up":
+ factor = float(rng.uniform(1.0, cfg.scale_range[1]))
+ elif updown == "down":
+ factor = float(rng.uniform(cfg.scale_range[0], 1.0))
+ else:
+ factor = 1.0
+ reference = current if cfg.relative_to == "current" else target
+ size = (max(1, int(round(reference[0] * factor))), max(1, int(round(reference[1] * factor))))
+ size = _clamp_size(size, floor)
+ mode = str(rng.choice(np.asarray(cfg.modes)))
+ # ``factor`` is the sampled value; ``factor_effective`` is the height ratio the emitted size realises after the
+ # intermediate floor, so records can be sliced on what actually happened.
+ return PlannedOp(
+ "resize",
+ {
+ "updown": updown,
+ "factor": factor,
+ "factor_effective": size[0] / reference[0],
+ "size": list(size),
+ "mode": mode,
+ },
+ )
+
+
+def _plan_noise(cfg: NoiseConfig, rng: np.random.Generator) -> PlannedOp | None:
+ if rng.uniform() >= cfg.prob:
+ return None
+ gray = bool(rng.uniform() < cfg.gray_noise_prob)
+ if rng.uniform() < cfg.gaussian_prob:
+ sigma = float(rng.uniform(*cfg.gaussian_sigma_range))
+ return PlannedOp("gaussian_noise", {"sigma": sigma, "gray": gray})
+ scale = float(rng.uniform(*cfg.poisson_scale_range))
+ return PlannedOp("poisson_noise", {"scale": scale, "gray": gray})
+
+
+def _plan_jpeg(cfg: JPEGConfig, rng: np.random.Generator) -> PlannedOp | None:
+ if rng.uniform() >= cfg.prob:
+ return None
+ return PlannedOp("jpeg", {"quality": float(rng.uniform(*cfg.quality_range))})
+
+
+def _plan_stage(
+ stage: DegradationStage,
+ rng: np.random.Generator,
+ current: tuple[int, int],
+ target: tuple[int, int],
+ floor: tuple[int, int],
+ res_factor: float,
+ max_kernel: int,
+ stage_name: str,
+) -> tuple[list[PlannedOp], tuple[int, int]]:
+ planned: list[PlannedOp] = []
+ if stage.blur is not None:
+ op = _plan_blur(stage.blur, rng, res_factor, max_kernel)
+ if op is not None:
+ planned.append(op)
+ if stage.resize is not None:
+ op = _plan_resize(stage.resize, rng, current, target, floor)
+ if op is not None:
+ planned.append(op)
+ current = tuple(op.params["size"])
+ if stage.noise is not None:
+ op = _plan_noise(stage.noise, rng)
+ if op is not None:
+ planned.append(op)
+ if stage.jpeg is not None:
+ op = _plan_jpeg(stage.jpeg, rng)
+ if op is not None:
+ planned.append(op)
+ for op in planned:
+ op.stage = stage_name
+ return planned, current
+
+
+def _plan_final(
+ cfg: FinalBlockConfig,
+ rng: np.random.Generator,
+ current: tuple[int, int],
+ target: tuple[int, int],
+ res_factor: float,
+ max_kernel: int,
+) -> list[PlannedOp]:
+ planned: list[PlannedOp] = []
+ sinc: PlannedOp | None = None
+ if rng.uniform() < cfg.sinc_prob:
+ base_size = int(rng.choice(np.asarray(cfg.kernel_range)))
+ kernel_size = scale_kernel_size(base_size, res_factor, max_kernel)
+ kernel, info = random_sinc_kernel(rng, kernel_size, cfg.sinc_cutoff_range) # [K,K]
+ sinc = PlannedOp("blur", info, kernel)
+ mode = str(rng.choice(np.asarray(cfg.modes)))
+ final_resize = PlannedOp( # nothing is sampled for the final resize; factor_effective keeps the schema uniform
+ "resize",
+ {
+ "updown": "final",
+ "factor": None,
+ "factor_effective": target[0] / current[0],
+ "size": list(target),
+ "mode": mode,
+ },
+ )
+ jpeg_op = _plan_jpeg(cfg.jpeg, rng)
+ if rng.uniform() < 0.5:
+ planned.append(final_resize)
+ if sinc is not None:
+ planned.append(sinc)
+ if jpeg_op is not None:
+ planned.append(jpeg_op)
+ else:
+ if jpeg_op is not None:
+ planned.append(jpeg_op)
+ planned.append(final_resize)
+ if sinc is not None:
+ planned.append(sinc)
+ for op in planned:
+ op.stage = "final"
+ return planned
+
+
+def _stage2_gate(seed: int) -> float:
+ """Uniform in [0, 1) for the optional second stage, on a stream separate from the plan's."""
+ return float(np.random.default_rng([seed, 2]).uniform())
+
+
+def plan_degradation(
+ profile: str | Profile,
+ hr_size: tuple[int, int],
+ scale: float = 2.0,
+ seed: int = 0,
+ target_size: tuple[int, int] | None = None,
+) -> DegradationPlan:
+ """Sample all clip-level parameters for ``profile`` on a clip of spatial size ``hr_size``."""
+ prof = get_profile(profile)
+ hr_size = (int(hr_size[0]), int(hr_size[1]))
+ lr_size = _lr_size(hr_size, scale, target_size)
+ rng = np.random.default_rng(seed)
+ planned: list[PlannedOp] = []
+
+ if isinstance(prof, CleanResizeProfile):
+ planned.append(PlannedOp("resize_clean", {"size": list(lr_size), "kernel": prof.kernel}, stage="final"))
+ codec_op = None if prof.codec is None else _plan_codec(prof.codec, rng)
+ return DegradationPlan(prof.name, seed, scale, hr_size, lr_size, planned, codec=codec_op)
+
+ assert isinstance(prof, RealESRGANProfile)
+ res_factor = 1.0
+ if prof.scale_kernels_with_resolution:
+ res_factor = max(hr_size) / prof.reference_longest_side
+ if prof.min_resolution_factor is not None:
+ res_factor = max(res_factor, prof.min_resolution_factor)
+ if prof.max_resolution_factor is not None:
+ res_factor = min(res_factor, prof.max_resolution_factor)
+ floor = (
+ max(1, int(round(lr_size[0] * prof.min_intermediate_scale))),
+ max(1, int(round(lr_size[1] * prof.min_intermediate_scale))),
+ )
+ current = hr_size
+ stage_ops, current = _plan_stage(
+ prof.stage1, rng, current, lr_size, floor, res_factor, prof.max_kernel_size, "stage1"
+ )
+ planned.extend(stage_ops)
+ # The stage-2 gate draws from its own stream, so profiles with stage2_prob = 1.0 keep their pre-existing
+ # seeded plans and stage 1 never depends on the gate. Stage 2's own draws come from the main stream, so on
+ # seeds where the gate skips it the final block and codec are re-rolled.
+ if prof.stage2 is not None and _stage2_gate(seed) < prof.stage2_prob:
+ stage_ops, current = _plan_stage(
+ prof.stage2, rng, current, lr_size, floor, res_factor, prof.max_kernel_size, "stage2"
+ )
+ planned.extend(stage_ops)
+ planned.extend(_plan_final(prof.final, rng, current, lr_size, res_factor, prof.max_kernel_size))
+ codec_op = None if prof.codec is None else _plan_codec(prof.codec, rng)
+ return DegradationPlan(prof.name, seed, scale, hr_size, lr_size, planned, codec=codec_op)
+
+
+def _plan_codec(cfg: CodecConfig, rng: np.random.Generator) -> PlannedOp | None:
+ if rng.uniform() >= cfg.prob:
+ return None
+ prob = np.asarray(cfg.codec_prob, dtype=np.float64) # [N]
+ codec = str(rng.choice(np.asarray(cfg.codecs), p=prob / prob.sum()))
+ crf = float(rng.uniform(*cfg.crf_range))
+ preset = str(rng.choice(np.asarray(cfg.presets)))
+ return PlannedOp("codec", {"codec": codec, "crf": crf, "preset": preset}, stage="codec")
+
+
+def apply_plan(
+ frames: torch.Tensor,
+ plan: DegradationPlan,
+ gen: torch.Generator,
+ jpeger: DiffJPEG,
+ jpeg_backend: str = "auto",
+ poisson_mode: str = "auto",
+) -> torch.Tensor: # frames: [T,C,H,W] float in [0,1]; returns [T,C,h,w] float in [0,1]
+ """Run every planned op on one chunk of frames."""
+ x = frames
+ for op in plan.ops:
+ if op.op == "blur":
+ assert op.kernel is not None
+ x = ops.blur(x, op.kernel) # [T,C,H,W]
+ elif op.op == "resize":
+ x = ops.resize(x, tuple(op.params["size"]), op.params["mode"]) # [T,C,h,w]
+ elif op.op == "resize_clean":
+ x = ops.resize_clean(x, tuple(op.params["size"]), op.params["kernel"]) # [T,C,h,w]
+ elif op.op == "gaussian_noise":
+ x = ops.add_gaussian_noise(x, op.params["sigma"], op.params["gray"], gen) # [T,C,h,w]
+ elif op.op == "poisson_noise":
+ x = ops.add_poisson_noise(x, op.params["scale"], op.params["gray"], gen, mode=poisson_mode) # [T,C,h,w]
+ elif op.op == "jpeg":
+ x = ops.jpeg(x, op.params["quality"], jpeger, backend=jpeg_backend) # [T,C,h,w]
+ else:
+ raise ValueError(f"Unknown planned op {op.op!r}")
+ if tuple(x.shape[-2:]) != plan.lr_size:
+ raise RuntimeError(f"Plan ended at size {tuple(x.shape[-2:])}, expected {plan.lr_size}")
+ return x
+
+
+def _to_tchw(hr: torch.Tensor) -> tuple[torch.Tensor, bool]: # returns ([T,C,H,W] uint8 or float in [0,1], is_image)
+ """Reorder to time-major without changing dtype; float conversion happens per chunk to bound memory."""
+ if hr.dim() == 3:
+ hr = hr.unsqueeze(1) # [C,1,H,W]
+ is_image = True
+ elif hr.dim() == 4:
+ is_image = False
+ else:
+ raise ValueError(f"Expected [C,T,H,W] or [C,H,W], got shape {tuple(hr.shape)}")
+ x = hr.permute(1, 0, 2, 3) # [T,C,H,W]
+ if x.is_floating_point():
+ if x.numel() > 0 and x.min() < 0.0:
+ raise ValueError("Float input must be in [0, 1]; got negative values (is it normalised to [-1, 1]?)")
+ elif x.dtype != torch.uint8:
+ raise TypeError(f"Unsupported dtype {x.dtype}")
+ return x, is_image
+
+
+def _chunk_to_float(chunk: torch.Tensor) -> torch.Tensor: # chunk: [t,C,H,W] uint8 or float, returns float32 in [0,1]
+ if chunk.dtype == torch.uint8:
+ return chunk.float() / 255.0 # [t,C,H,W]
+ return chunk.float() # [t,C,H,W]
+
+
+def degrade_hr_to_lr(
+ hr: torch.Tensor,
+ profile: str | Profile,
+ scale: float = 2.0,
+ seed: int = 0,
+ target_size: tuple[int, int] | None = None,
+ chunk_frames: int = 8,
+ jpeger: DiffJPEG | None = None,
+ jpeg_backend: str = "auto",
+ poisson_mode: str = "auto",
+ fps: float = 24.0,
+) -> DegradationResult:
+ """Degrade an HR clip or image into LR with a fully seeded, clip-consistent parameter set.
+
+ Args:
+ hr: ``[C,T,H,W]`` video or ``[C,H,W]`` image, uint8 or float in [0, 1], any device.
+ profile: profile name from ``PROFILES`` or a profile dataclass.
+ scale: HR-to-LR downscale factor; LR is ``round(H/scale) x round(W/scale)`` unless ``target_size``.
+ seed: seeds both parameter sampling and noise realisation.
+ target_size: explicit LR ``(h, w)``; overrides ``scale`` for the output size.
+ chunk_frames: frames processed per step; bounds peak memory (float32 intermediates).
+ jpeger: optional reusable ``DiffJPEG`` module (avoids re-creating buffers per call).
+ jpeg_backend: ``"auto"`` (libjpeg via cv2 on CPU, DiffJPEG on GPU), ``"cv2"`` or ``"diffjpeg"``.
+ poisson_mode: ``"auto"`` (exact on GPU, Gaussian approximation on CPU), ``"exact"`` or
+ ``"gaussian_approx"``. Exact Poisson sampling costs about 0.3 s per 1080p frame on CPU.
+ fps: frame rate of the clip, used only by the codec stage (P3) as the encoder's stream rate, which
+ feeds x264/x265 rate control. Pass the sample's real fps; ignored for images and codec-free profiles.
+
+ Returns:
+ ``DegradationResult`` with ``lr`` as uint8 in the input layout and the parameter ``record``.
+ The record also states the resolved ``jpeg_backend`` and ``poisson_mode`` and the device.
+ """
+ x, is_image = _to_tchw(hr) # [T,C,H,W], input dtype
+ plan = plan_degradation(profile, tuple(x.shape[-2:]), scale=scale, seed=seed, target_size=target_size)
+ codec_skipped = None
+ if plan.codec is not None and (is_image or x.shape[0] < 2):
+ # A video codec needs a clip; on an image or a one-frame clip the op cannot run, so it leaves the plan
+ # (the record must not list an op that never happened) and the record says why.
+ codec_skipped = "single_frame"
+ plan = dataclasses.replace(plan, codec=None)
+ gen = ops.make_generator(seed, x.device)
+ resolved_jpeg = ops.resolve_jpeg_backend(jpeg_backend, x.device)
+ resolved_poisson = ops.resolve_poisson_mode(poisson_mode, x.device)
+ if jpeger is None:
+ jpeger = DiffJPEG(differentiable=False)
+ chunks: list[torch.Tensor] = []
+ with torch.no_grad():
+ for start in range(0, x.shape[0], max(1, chunk_frames)):
+ chunk = _chunk_to_float(x[start : start + chunk_frames]) # [t,C,H,W] float32
+ out = apply_plan(chunk, plan, gen, jpeger, jpeg_backend=resolved_jpeg, poisson_mode=resolved_poisson)
+ chunks.append(ops.to_uint8(out)) # [t,C,h,w]
+ lr_tchw = torch.cat(chunks, dim=0) # [T,C,h,w] uint8
+ codec_applied = plan.codec is not None # single-frame inputs had the codec removed from the plan above
+ if plan.codec is not None:
+ # Whole-clip op: needs every frame at once, runs on CPU (software encoders), returns to device.
+ params = plan.codec.params
+ lr_tchw = codec_round_trip(
+ lr_tchw, codec=params["codec"], crf=params["crf"], preset=params["preset"], fps=fps
+ )
+ lr = lr_tchw.permute(1, 0, 2, 3).contiguous() # [C,T,h,w]
+ if is_image:
+ lr = lr[:, 0] # [C,h,w]
+ record = plan.record()
+ record.update(
+ {
+ "jpeg_backend": resolved_jpeg,
+ "poisson_mode": resolved_poisson,
+ "device": x.device.type,
+ "codec_applied": codec_applied,
+ "codec_skipped": codec_skipped,
+ "codec_fps": float(fps) if codec_applied else None,
+ }
+ )
+ return DegradationResult(lr=lr, record=record)
+
+
+def degrade_batch(
+ hr_list: Sequence[torch.Tensor], profile: str | Profile, scale: float, seeds: Sequence[int], **kwargs: Any
+) -> list[DegradationResult]:
+ """Convenience wrapper for a list of clips with one seed each."""
+ if len(hr_list) != len(seeds):
+ raise ValueError("hr_list and seeds must have equal length")
+ jpeger = kwargs.pop("jpeger", None) or DiffJPEG(differentiable=False)
+ return [
+ degrade_hr_to_lr(hr, profile, scale=scale, seed=s, jpeger=jpeger, **kwargs) for hr, s in zip(hr_list, seeds)
+ ]
diff --git a/cosmos_framework/data/generator/augmentors/hr_lr_degradation/degrade_test.py b/cosmos_framework/data/generator/augmentors/hr_lr_degradation/degrade_test.py
new file mode 100644
index 000000000..208cdfccd
--- /dev/null
+++ b/cosmos_framework/data/generator/augmentors/hr_lr_degradation/degrade_test.py
@@ -0,0 +1,650 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+
+import dataclasses
+import json
+import math
+
+import numpy as np
+import pytest
+import torch
+import torch.nn.functional as F
+
+from cosmos_framework.data.generator.augmentors.hr_lr_degradation import (
+ PROFILES,
+ RealESRGANProfile,
+ degrade_hr_to_lr,
+ get_profile,
+ ops,
+)
+from cosmos_framework.data.generator.augmentors.hr_lr_degradation.degrade import plan_degradation
+from cosmos_framework.data.generator.augmentors.hr_lr_degradation.diffjpeg import DiffJPEG
+from cosmos_framework.data.generator.augmentors.hr_lr_degradation.kernels import (
+ circular_lowpass_kernel,
+ random_mixed_kernel,
+ scale_kernel_size,
+)
+from cosmos_framework.data.generator.augmentors.hr_lr_degradation.profiles import (
+ IMAGE_SR_DEFAULT_MIX,
+ REGIME_PROFILES,
+ VIDEO_SR_DEFAULT_MIX,
+ BlurConfig,
+ DegradationStage,
+ FinalBlockConfig,
+ JPEGConfig,
+ Profile,
+ profile_to_dict,
+)
+
+pytestmark = [pytest.mark.L0, pytest.mark.CPU]
+
+_ALL_PROFILES = sorted(PROFILES)
+
+
+def _synthetic_clip(t: int = 5, h: int = 96, w: int = 128, seed: int = 0) -> torch.Tensor: # returns [3,T,H,W] uint8
+ """Smooth gradients plus a moving edge so blur, resize and JPEG all have something to act on."""
+ g = torch.Generator().manual_seed(seed)
+ ys = torch.linspace(0, 1, h).view(1, 1, h, 1) # [1,1,H,1]
+ xs = torch.linspace(0, 1, w).view(1, 1, 1, w) # [1,1,1,W]
+ ts = torch.linspace(0, 1, t).view(1, t, 1, 1) # [1,T,1,1]
+ r = ys.expand(1, t, h, w) # [1,T,H,W]
+ gch = xs.expand(1, t, h, w) # [1,T,H,W]
+ b = ((xs + 0.3 * ts) % 1.0 > 0.5).float().expand(1, t, h, w) # [1,T,H,W]
+ clip = torch.cat([r, gch, b], dim=0) # [3,T,H,W]
+ clip = clip + 0.02 * torch.rand(clip.shape, generator=g) # [3,T,H,W]
+ return (clip.clamp(0, 1) * 255).round().to(torch.uint8) # [3,T,H,W]
+
+
+@pytest.mark.parametrize("profile_name", _ALL_PROFILES)
+def test_video_output_shape_dtype_and_range(profile_name: str) -> None:
+ hr = _synthetic_clip(t=5, h=96, w=128) # [3,5,96,128]
+ result = degrade_hr_to_lr(hr, profile_name, scale=2, seed=123, chunk_frames=2)
+ assert result.lr.shape == (3, 5, 48, 64)
+ assert result.lr.dtype == torch.uint8
+ assert result.record["profile_name"] == profile_name
+ assert result.record["lr_size"] == [48, 64]
+ json.dumps(result.record) # record must be serialisable
+
+
+@pytest.mark.parametrize("profile_name", ["p0_clean_bicubic", "p1_second_order"])
+def test_image_input_keeps_layout(profile_name: str) -> None:
+ hr = _synthetic_clip(t=1, h=64, w=80)[:, 0] # [3,64,80]
+ result = degrade_hr_to_lr(hr, profile_name, scale=2, seed=7)
+ assert result.lr.shape == (3, 32, 40)
+ assert result.lr.dtype == torch.uint8
+
+
+def test_float_input_matches_uint8_input() -> None:
+ hr_u8 = _synthetic_clip(t=2, h=64, w=64) # [3,2,64,64]
+ hr_f = hr_u8.float() / 255.0 # [3,2,64,64]
+ out_u8 = degrade_hr_to_lr(hr_u8, "p1_second_order", seed=3).lr
+ out_f = degrade_hr_to_lr(hr_f, "p1_second_order", seed=3).lr
+ assert torch.equal(out_u8, out_f)
+
+
+def test_normalised_float_input_is_rejected() -> None:
+ hr = torch.rand(3, 2, 32, 32) * 2 - 1 # [3,2,32,32] in [-1,1]
+ with pytest.raises(ValueError, match=r"\[0, 1\]"):
+ degrade_hr_to_lr(hr, "p0_clean_bicubic")
+
+
+@pytest.mark.parametrize("profile_name", ["p1_first_order", "p1_second_order"])
+def test_same_seed_is_deterministic_and_chunking_invariant(profile_name: str) -> None:
+ hr = _synthetic_clip(t=6, h=64, w=96) # [3,6,64,96]
+ a = degrade_hr_to_lr(hr, profile_name, seed=42, chunk_frames=6)
+ b = degrade_hr_to_lr(hr, profile_name, seed=42, chunk_frames=6)
+ assert torch.equal(a.lr, b.lr)
+ assert a.record == b.record
+
+
+def test_chunking_only_changes_noise_realisation_not_parameters() -> None:
+ """Chunking must not change the sampled plan, and may only perturb pixels through float rounding.
+
+ Without a noise stage the pipeline is a per-frame map, but not a bitwise-reproducible one across chunk sizes:
+ the FFT convolution picks a different plan for a batch of 6 frames than for a batch of 2 (observed on aarch64
+ CI runners; x86 happened to agree), which can move a pixel by one 8-bit level. A JPEG stage after the blur
+ amplifies such one-level input changes into several output levels on a small fraction of pixels. So:
+ compression-free plans must agree to within one level; plans with JPEG must agree statistically.
+ """
+ hr = _synthetic_clip(t=6, h=64, w=96) # [3,6,64,96]
+
+ # 1) Same sampled parameters regardless of chunking (the property the training pipeline relies on).
+ a = degrade_hr_to_lr(hr, "p1_first_order_no_noise", seed=42, chunk_frames=6)
+ b = degrade_hr_to_lr(hr, "p1_first_order_no_noise", seed=42, chunk_frames=2)
+ assert a.record == b.record
+
+ # 2) Blur + resize only: differences are pure float rounding, at most one 8-bit level on few pixels.
+ base = get_profile("p1_first_order_no_noise")
+ assert isinstance(base, RealESRGANProfile)
+ no_compression = dataclasses.replace(
+ base,
+ name="chunking_probe_no_compression",
+ stage1=DegradationStage(noise=None, jpeg=None),
+ final=dataclasses.replace(base.final, jpeg=JPEGConfig(prob=0.0)),
+ )
+ c = degrade_hr_to_lr(hr, no_compression, seed=42, chunk_frames=6)
+ d = degrade_hr_to_lr(hr, no_compression, seed=42, chunk_frames=2)
+ assert c.record == d.record
+ diff = (c.lr.int() - d.lr.int()).abs()
+ assert diff.max() <= 1, f"chunking changed compression-free pixels by up to {int(diff.max())} levels"
+ assert (diff > 0).float().mean() < 0.02, f"{(diff > 0).float().mean():.3%} of pixels differ across chunkings"
+
+ # 3) With JPEG in the chain, a handful of pixels may move by several levels; the images stay the same picture.
+ diff_jpeg = (a.lr.int() - b.lr.int()).abs()
+ assert diff_jpeg.float().mean() < 0.1, f"mean abs diff {diff_jpeg.float().mean():.4f} levels across chunkings"
+ assert (diff_jpeg > 0).float().mean() < 0.05, f"{(diff_jpeg > 0).float().mean():.3%} of pixels differ"
+
+
+def test_different_seeds_give_different_parameters() -> None:
+ plans = {plan_degradation("p1_second_order", (720, 1280), seed=s).record()["ops"].__repr__() for s in range(8)}
+ assert len(plans) > 1
+
+
+def test_p0_matches_torch_antialiased_bicubic_reference() -> None:
+ hr = _synthetic_clip(t=3, h=96, w=128) # [3,3,96,128]
+ out = degrade_hr_to_lr(hr, "p0_clean_bicubic", scale=2, seed=0).lr # [3,3,48,64]
+ ref = F.interpolate(
+ hr.permute(1, 0, 2, 3).float() / 255.0, size=(48, 64), mode="bicubic", align_corners=False, antialias=True
+ ) # [3,3,48,64]
+ ref_u8 = (ref.clamp(0, 1) * 255).round().to(torch.uint8).permute(1, 0, 2, 3) # [3,3,48,64]
+ assert torch.equal(out, ref_u8)
+
+
+def test_p0_is_seed_independent() -> None:
+ hr = _synthetic_clip(t=2, h=64, w=64) # [3,2,64,64]
+ assert torch.equal(
+ degrade_hr_to_lr(hr, "p0_clean_bicubic", seed=1).lr, degrade_hr_to_lr(hr, "p0_clean_bicubic", seed=2).lr
+ )
+
+
+def test_target_size_overrides_scale() -> None:
+ hr = _synthetic_clip(t=2, h=90, w=160) # [3,2,90,160]
+ out = degrade_hr_to_lr(hr, "p1_first_order", scale=2, seed=0, target_size=(40, 72)).lr
+ assert out.shape == (3, 2, 40, 72)
+
+
+def test_plan_respects_intermediate_floor_and_ends_at_lr_size() -> None:
+ profile = get_profile("p1_second_order")
+ assert isinstance(profile, RealESRGANProfile)
+ for seed in range(50):
+ plan = plan_degradation(profile, (720, 1280), scale=2, seed=seed)
+ floor_h = round(plan.lr_size[0] * profile.min_intermediate_scale)
+ floor_w = round(plan.lr_size[1] * profile.min_intermediate_scale)
+ sizes = [tuple(op.params["size"]) for op in plan.ops if op.op == "resize"]
+ assert sizes[-1] == plan.lr_size
+ for h, w in sizes:
+ assert h >= floor_h and w >= floor_w
+ for op in plan.ops:
+ if op.op == "blur":
+ assert op.kernel is not None and op.kernel.shape[0] % 2 == 1
+ assert op.kernel.shape[0] <= profile.max_kernel_size
+ assert abs(float(op.kernel.sum()) - 1.0) < 1e-4
+
+
+def test_kernel_sizes_scale_with_resolution() -> None:
+ small = [
+ op.params["kernel_size"]
+ for s in range(40)
+ for op in plan_degradation("p1_first_order", (360, 640), seed=s).ops
+ if op.op == "blur"
+ ]
+ large = [
+ op.params["kernel_size"]
+ for s in range(40)
+ for op in plan_degradation("p1_first_order", (1080, 1920), seed=s).ops
+ if op.op == "blur"
+ ]
+ assert np.mean(large) > np.mean(small)
+ published = [
+ op.params["kernel_size"]
+ for s in range(40)
+ for op in plan_degradation("p1_second_order_published", (1080, 1920), seed=s).ops
+ if op.op == "blur"
+ ]
+ assert max(published) <= 21
+
+
+def test_scale_kernel_size_is_odd_and_bounded() -> None:
+ for base in (7, 9, 21):
+ for factor in (0.3, 1.0, 2.67, 10.0):
+ k = scale_kernel_size(base, factor, 61)
+ assert k % 2 == 1 and 3 <= k <= 61
+
+
+def test_mixed_and_sinc_kernels_are_normalised() -> None:
+ rng = np.random.default_rng(0)
+ for _ in range(20):
+ kernel, info = random_mixed_kernel(
+ rng,
+ get_profile("p1_first_order").stage1.blur.kernel_list,
+ get_profile("p1_first_order").stage1.blur.kernel_prob,
+ 13,
+ (0.2, 3),
+ (0.2, 3),
+ )
+ assert kernel.shape == (13, 13) and abs(kernel.sum() - 1) < 1e-6 and info["kernel_type"]
+ sinc = circular_lowpass_kernel(np.pi / 2, 11, pad_to=21)
+ assert sinc.shape == (21, 21) and abs(sinc.sum() - 1) < 1e-6
+
+
+def test_diffjpeg_degrades_more_at_lower_quality_and_handles_odd_sizes() -> None:
+ jpeger = DiffJPEG()
+ x = _synthetic_clip(t=2, h=45, w=67).permute(1, 0, 2, 3).float() / 255.0 # [2,3,45,67]
+ hi = jpeger(x, quality=95.0) # [2,3,45,67]
+ lo = jpeger(x, quality=10.0) # [2,3,45,67]
+ assert hi.shape == x.shape and lo.shape == x.shape
+ assert (hi - x).abs().mean() < (lo - x).abs().mean()
+ assert (hi - x).abs().mean() < 0.02
+ per_frame = jpeger(x, quality=torch.tensor([95.0, 10.0])) # [2,3,45,67]
+ assert torch.allclose(per_frame[0], hi[0]) and torch.allclose(per_frame[1], lo[1])
+
+
+def test_degradation_actually_changes_content_relative_to_clean() -> None:
+ hr = _synthetic_clip(t=2, h=96, w=128) # [3,2,96,128]
+ clean = degrade_hr_to_lr(hr, "p0_clean_bicubic").lr.float()
+ degraded = degrade_hr_to_lr(hr, "p1_second_order", seed=5).lr.float()
+ assert (clean - degraded).abs().mean() > 0.5 # 8-bit units
+
+
+def test_profiles_are_frozen_copyable_and_serialisable() -> None:
+ base = get_profile("p1_first_order")
+ arm = dataclasses.replace(base, name="arm", stage1=DegradationStage(noise=None))
+ assert arm.stage1.noise is None and base.stage1.noise is not None
+ json.dumps(profile_to_dict(arm))
+ with pytest.raises(KeyError):
+ get_profile("does_not_exist")
+
+
+def test_fft_filter_matches_direct_convolution() -> None:
+ x = _synthetic_clip(t=2, h=64, w=80).permute(1, 0, 2, 3).float() / 255.0 # [2,3,64,80]
+ rng = np.random.default_rng(1)
+ for k in (11, 21, 33):
+ kernel = torch.from_numpy(rng.random((k, k)).astype(np.float32)) # asymmetric on purpose
+ kernel = kernel / kernel.sum()
+ direct = ops._filter2d_direct(x, kernel) # [2,3,64,80]
+ via_fft = ops._filter2d_fft(x, kernel) # [2,3,64,80]
+ assert torch.allclose(direct, via_fft, atol=1e-5), f"k={k}: max err {(direct - via_fft).abs().max()}"
+
+
+def test_cv2_and_diffjpeg_backends_behave_alike() -> None:
+ x = _synthetic_clip(t=2, h=45, w=67).permute(1, 0, 2, 3).float() / 255.0 # [2,3,45,67]
+ jpeger = DiffJPEG()
+ for quality in (90.0, 20.0):
+ via_cv2 = ops.jpeg(x, quality, jpeger, backend="cv2") # [2,3,45,67]
+ via_torch = ops.jpeg(x, quality, jpeger, backend="diffjpeg") # [2,3,45,67]
+ assert via_cv2.shape == x.shape and via_torch.shape == x.shape
+ # Both are lossy codecs of the same picture: they should agree with each other about as well as with the input.
+ assert (via_cv2 - via_torch).abs().mean() < 2.0 * max((via_cv2 - x).abs().mean(), (via_torch - x).abs().mean())
+ err_hi = (ops.jpeg(x, 90.0, jpeger, backend="cv2") - x).abs().mean()
+ err_lo = (ops.jpeg(x, 20.0, jpeger, backend="cv2") - x).abs().mean()
+ assert err_hi < err_lo
+ with pytest.raises(ValueError):
+ ops.resolve_jpeg_backend("nope", torch.device("cpu"))
+
+
+def test_poisson_modes_have_signal_dependent_variance() -> None:
+ gen = ops.make_generator(0, "cpu")
+ dark = torch.full((4, 3, 64, 64), 0.05) # [4,3,64,64]
+ bright = torch.full((4, 3, 64, 64), 0.6) # [4,3,64,64]
+ for mode in ("exact", "gaussian_approx"):
+ noise_dark = ops.add_poisson_noise(dark, 1.0, False, gen, mode=mode) - dark
+ noise_bright = ops.add_poisson_noise(bright, 1.0, False, gen, mode=mode) - bright
+ assert noise_dark.std() < noise_bright.std()
+ assert noise_bright.std() > 0.0
+ assert ops.resolve_poisson_mode("auto", torch.device("cpu")) == "gaussian_approx"
+ assert ops.resolve_poisson_mode("auto", torch.device("cuda")) == "exact"
+
+
+def test_record_states_resolved_runtime_choices() -> None:
+ hr = _synthetic_clip(t=2, h=64, w=64) # [3,2,64,64]
+ record = degrade_hr_to_lr(hr, "p1_first_order", seed=1).record
+ assert record["jpeg_backend"] == "cv2" and record["poisson_mode"] == "gaussian_approx" and record["device"] == "cpu"
+ record_torch_jpeg = degrade_hr_to_lr(hr, "p1_first_order", seed=1, jpeg_backend="diffjpeg").record
+ assert record_torch_jpeg["jpeg_backend"] == "diffjpeg"
+
+
+@pytest.mark.GPU
+@pytest.mark.skipif(not torch.cuda.is_available(), reason="needs CUDA")
+def test_gpu_shares_parameters_with_cpu_and_is_deterministic() -> None:
+ hr = _synthetic_clip(t=4, h=96, w=128) # [3,4,96,128]
+ cpu = degrade_hr_to_lr(hr, "p1_second_order", seed=11)
+ gpu_a = degrade_hr_to_lr(hr.cuda(), "p1_second_order", seed=11)
+ gpu_b = degrade_hr_to_lr(hr.cuda(), "p1_second_order", seed=11)
+ assert gpu_a.lr.device.type == "cuda"
+ assert torch.equal(gpu_a.lr, gpu_b.lr)
+ # Sampled parameters are device independent; only noise realisation, JPEG codec and Poisson mode differ.
+ keys = ("profile_name", "seed", "scale", "hr_size", "lr_size", "ops")
+ assert {k: cpu.record[k] for k in keys} == {k: gpu_a.record[k] for k in keys}
+ assert (cpu.lr.float() - gpu_a.lr.float().cpu()).abs().mean() < 12.0 # 8-bit units; noise + JPEG backend differ
+ # Without noise (its realisation is device specific and dominates the residual) what remains is the
+ # cv2-vs-DiffJPEG gap plus resize / FFT rounding, a few 8-bit units on this content.
+ base = get_profile("p1_second_order")
+ assert isinstance(base, RealESRGANProfile) and base.stage2 is not None
+ noise_free = dataclasses.replace(
+ base,
+ name="p1_second_order_noise_free",
+ stage1=dataclasses.replace(base.stage1, noise=None),
+ stage2=dataclasses.replace(base.stage2, noise=None),
+ )
+ cpu_nf = degrade_hr_to_lr(hr, noise_free, seed=11)
+ gpu_nf = degrade_hr_to_lr(hr.cuda(), noise_free, seed=11)
+ assert (cpu_nf.lr.float() - gpu_nf.lr.float().cpu()).abs().mean() < 10.0 # 8-bit units
+
+
+def test_codec_round_trip_keeps_shape_dtype_and_is_lossy() -> None:
+ from cosmos_framework.data.generator.augmentors.hr_lr_degradation.codec import codec_available, codec_round_trip
+
+ assert codec_available("libx264") and codec_available("libx265")
+ clip = _synthetic_clip(t=6, h=45, w=67).permute(1, 0, 2, 3) # [6,3,45,67] uint8, odd sizes
+ for codec in ("libx264", "libx265"):
+ hi = codec_round_trip(clip, codec=codec, crf=18, preset="veryfast") # [6,3,45,67]
+ lo = codec_round_trip(clip, codec=codec, crf=40, preset="veryfast") # [6,3,45,67]
+ assert hi.shape == clip.shape and hi.dtype == torch.uint8
+ err_hi = (hi.float() - clip.float()).abs().mean()
+ err_lo = (lo.float() - clip.float()).abs().mean()
+ assert 0.0 < err_hi < err_lo, f"{codec}: {err_hi} vs {err_lo}"
+ as_float = codec_round_trip(clip.float() / 255.0, codec="libx264", crf=23) # [6,3,45,67] float
+ assert as_float.is_floating_point() and 0.0 <= as_float.min() and as_float.max() <= 1.0
+
+
+def test_p3_profile_applies_codec_to_video_but_not_images() -> None:
+ hr = _synthetic_clip(t=6, h=96, w=128) # [3,6,96,128]
+ applied = [degrade_hr_to_lr(hr, "p3_video_codec", seed=s).record for s in range(12)]
+ assert any(r["codec_applied"] for r in applied) and any(r["codec"] is None for r in applied) # prob 0.6
+ with_codec = next(r for r in applied if r["codec_applied"])
+ assert with_codec["codec"]["codec"] in ("libx264", "libx265") and 18 <= with_codec["codec"]["crf"] <= 35
+ image = degrade_hr_to_lr(hr[:, 0], "p3_video_codec", seed=0).record
+ assert image["codec_applied"] is False
+ a = degrade_hr_to_lr(hr, "p3_video_codec", seed=int(with_codec["seed"]))
+ b = degrade_hr_to_lr(hr, "p3_video_codec", seed=int(with_codec["seed"]))
+ assert torch.equal(a.lr, b.lr) # codec round trip is deterministic for the same input and settings
+
+
+def test_codec_fps_is_recorded_only_when_the_codec_ran() -> None:
+ hr = _synthetic_clip(t=6, h=64, w=64) # [3,6,64,64]
+ records = [degrade_hr_to_lr(hr, "p3_video_codec", seed=s, fps=30.0).record for s in range(12)]
+ with_codec = [r for r in records if r["codec_applied"]]
+ without = [r for r in records if not r["codec_applied"]]
+ assert with_codec and without
+ assert all(r["codec_fps"] == 30.0 for r in with_codec) and all(r["codec_fps"] is None for r in without)
+ assert degrade_hr_to_lr(hr, "p1_first_order", seed=0, fps=30.0).record["codec_fps"] is None
+
+
+def test_nvenc_options_align_with_x264_semantics() -> None:
+ from cosmos_framework.data.generator.augmentors.hr_lr_degradation.codec import (
+ _encoder_options,
+ codec_available,
+ codec_round_trip,
+ nvenc_preset,
+ )
+
+ assert nvenc_preset("veryfast") == "p1" and nvenc_preset("medium") == "p4" and nvenc_preset("veryslow") == "p7"
+ order = [
+ nvenc_preset(p) for p in ("ultrafast", "veryfast", "faster", "fast", "medium", "slow", "slower", "veryslow")
+ ]
+ assert order == sorted(order) # monotone in speed
+ assert nvenc_preset("p3") == "p3"
+ with pytest.raises(ValueError):
+ nvenc_preset("turbo")
+ opts = _encoder_options("h264_nvenc", 27.6, "medium")
+ assert opts == {"rc": "vbr", "cq": "28", "b": "0", "preset": "p4"} # CRF analogue, not constqp
+ assert _encoder_options("libx264", 27.6, "medium", threads=2) == {"crf": "28", "preset": "medium", "threads": "2"}
+ if codec_available("h264_nvenc") and torch.cuda.is_available():
+ with pytest.raises(ValueError, match="at least"):
+ codec_round_trip(_synthetic_clip(t=2, h=96, w=128).permute(1, 0, 2, 3), codec="h264_nvenc")
+ clip = _synthetic_clip(t=6, h=240, w=320).permute(1, 0, 2, 3) # [6,3,240,320]
+ hi = codec_round_trip(clip, codec="h264_nvenc", crf=18, preset="veryfast")
+ lo = codec_round_trip(clip, codec="h264_nvenc", crf=45, preset="medium")
+ assert hi.shape == clip.shape
+ assert (hi.float() - clip.float()).abs().mean() < (lo.float() - clip.float()).abs().mean()
+
+
+def test_codec_available_reports_an_encoder_that_opens(monkeypatch: pytest.MonkeyPatch) -> None:
+ """An encoder that opens must come back True, cleanup included.
+
+ The probe used to call ``CodecContext.close()``, which PyAV 17/18 do not define; the
+ ``AttributeError`` landed in the probe's own ``except`` and reported every *working* hardware
+ encoder as unavailable. No NVENC assertion can catch that on CI -- those runners have no encoder
+ engine, so the probe fails at ``open()`` and never reaches the cleanup -- so drive the hardware
+ path with libx264, which opens everywhere.
+ """
+ from cosmos_framework.data.generator.augmentors.hr_lr_degradation import codec as codec_mod
+
+ monkeypatch.setattr(codec_mod, "HARDWARE_CODECS", ("libx264",))
+ monkeypatch.setattr(codec_mod, "_HARDWARE_PROBE_PASSED", set())
+ assert codec_mod.codec_available("libx264") is True
+ assert codec_mod.codec_available("not_a_codec") is False
+
+
+def test_codec_available_retries_after_a_failed_hardware_probe(monkeypatch: pytest.MonkeyPatch) -> None:
+ """A failing probe must not be remembered: encoder sessions free up again.
+
+ ``codec_round_trip`` raises when an encoder is unavailable, so caching a transient shortage
+ would fail every later call for the life of the process, blaming the FFmpeg build.
+ """
+ import types
+
+ from cosmos_framework.data.generator.augmentors.hr_lr_degradation import codec as codec_mod
+
+ real_av, calls = codec_mod.av, []
+
+ def create(name: str, mode: str):
+ calls.append(name)
+ if len(calls) == 1:
+ raise RuntimeError("OpenEncodeSessionEx failed: out of memory") # what exhaustion looks like
+ return real_av.codec.context.CodecContext.create(name, mode)
+
+ monkeypatch.setattr(
+ codec_mod,
+ "av",
+ types.SimpleNamespace(
+ codec=types.SimpleNamespace(
+ Codec=real_av.codec.Codec,
+ context=types.SimpleNamespace(CodecContext=types.SimpleNamespace(create=create)),
+ )
+ ),
+ )
+ monkeypatch.setattr(codec_mod, "HARDWARE_CODECS", ("libx264",))
+ monkeypatch.setattr(codec_mod, "_HARDWARE_PROBE_PASSED", set())
+ assert codec_mod.codec_available("libx264") is False # all sessions busy
+ assert codec_mod.codec_available("libx264") is True # capacity back; probe re-run, not poisoned
+ assert len(calls) == 2
+
+
+def test_codec_round_trip_bounds_encoder_and_decoder_threads() -> None:
+ """x264/x265 default to machine-sized thread pools; under 8 xdist workers in CI that exhausted the container's
+ thread limit and stalled the CPU test phase, so the round trip must keep its thread count small."""
+ import threading
+ import time
+
+ from cosmos_framework.data.generator.augmentors.hr_lr_degradation.codec import (
+ DEFAULT_CODEC_THREADS,
+ _encoder_options,
+ codec_round_trip,
+ )
+
+ assert _encoder_options("libx264", 30, "veryfast", threads=2)["threads"] == "2"
+ assert "pools=2" in _encoder_options("libx265", 30, "veryfast", threads=2)["x265-params"]
+ assert DEFAULT_CODEC_THREADS <= 4
+
+ def thread_count() -> int:
+ return int(next(line for line in open("/proc/self/status") if line.startswith("Threads")).split()[1])
+
+ clip = _synthetic_clip(t=6, h=96, w=128).permute(1, 0, 2, 3) # [6,3,96,128]
+ for codec in ("libx264", "libx265"):
+ baseline = thread_count()
+ peak = [baseline]
+ stop = threading.Event()
+
+ def sample() -> None:
+ while not stop.is_set():
+ peak.append(thread_count())
+ time.sleep(0.001)
+
+ sampler = threading.Thread(target=sample)
+ sampler.start()
+ codec_round_trip(clip, codec=codec, crf=30, preset="veryfast", threads=2)
+ stop.set()
+ sampler.join()
+ extra = max(peak) - baseline - 1 # minus the sampler thread itself
+ # Unbounded this is ~40 (x264) to ~70 (x265) on 16 cores and scales with the host; bounded it is a handful.
+ assert extra <= 8, f"{codec} spawned {extra} extra threads with threads=2"
+
+
+def _stage1_kernels(profile: str | Profile, hr_size: tuple[int, int], seed: int) -> list[int]:
+ ops_ = plan_degradation(profile, hr_size, seed=seed).ops
+ return [o.params["kernel_size"] for o in ops_ if o.op == "blur" and o.stage == "stage1"]
+
+
+def test_plan_records_carry_stage_tags_and_effective_resize_factors() -> None:
+ base = get_profile("p1_second_order")
+ assert isinstance(base, RealESRGANProfile) and base.stage1.resize is not None
+ order = {"stage1": 0, "stage2": 1, "final": 2}
+ clamped = 0
+ for seed in range(50):
+ plan = plan_degradation(base, (720, 1280), seed=seed)
+ tags = [o.stage for o in plan.ops]
+ assert set(tags) <= set(order) and "final" in tags and tags == sorted(tags, key=order.__getitem__)
+ reference = (720, 1280) if base.stage1.resize.relative_to == "current" else plan.lr_size
+ for o in plan.ops:
+ if o.op == "resize": # one schema for every resize op, the final one included
+ assert {"updown", "factor", "factor_effective", "size", "mode"} <= set(o.params)
+ if o.stage == "stage1":
+ assert o.params["factor_effective"] == o.params["size"][0] / reference[0]
+ assert o.params["size"][0] >= round(0.75 * plan.lr_size[0]) - 1
+ clamped += o.params["factor_effective"] > o.params["factor"] + 1e-6
+ assert clamped > 0 # the inherited range dips below the 0.75 floor; the record shows the draw and the outcome
+ for name in PROFILES: # profiles stay strict-JSON serialisable (no inf / NaN defaults)
+ json.dumps(profile_to_dict(get_profile(name)), allow_nan=False)
+
+
+def test_codec_is_dropped_from_the_plan_and_record_for_single_frame_inputs() -> None:
+ hr = _synthetic_clip(t=1, h=64, w=80) # [3,1,64,80]
+ seed = next(s for s in range(20) if plan_degradation("p3_video_codec", (64, 80), seed=s).codec is not None)
+ for hr_in in (hr[:, 0], hr): # image layout and one-frame clip
+ rec = degrade_hr_to_lr(hr_in, "p3_video_codec", seed=seed).record
+ assert rec["codec"] is None and rec["codec_applied"] is False and rec["codec_skipped"] == "single_frame"
+ rec = degrade_hr_to_lr(_synthetic_clip(t=6, h=64, w=80), "p3_video_codec", seed=seed).record
+ assert rec["codec"]["stage"] == "codec" and rec["codec_applied"] is True and rec["codec_skipped"] is None
+
+
+def test_sinc_cutoff_range_is_honoured_and_defaults_to_the_real_esrgan_prior() -> None:
+ base = get_profile("p1_first_order")
+ assert isinstance(base, RealESRGANProfile) and base.stage1.blur is not None
+ bounded_blur = dataclasses.replace(base.stage1.blur, sinc_prob=1.0, sinc_cutoff_range=(math.pi / 2, math.pi))
+ bounded = dataclasses.replace(base, name="p1_bounded", stage1=dataclasses.replace(base.stage1, blur=bounded_blur))
+
+ def cutoffs(profile: Profile, seeds: int) -> list[float]:
+ return [
+ o.params["omega_c"]
+ for s in range(seeds)
+ for o in plan_degradation(profile, (1080, 1920), seed=s).ops
+ if o.op == "blur" and o.params["kernel_type"] == "sinc" and o.stage == "stage1"
+ ]
+
+ explicit = cutoffs(bounded, 100)
+ assert explicit and math.pi / 2 - 1e-9 <= min(explicit) and max(explicit) <= math.pi + 1e-9
+ assert min(cutoffs(base, 500)) < math.pi / 2 # None keeps Real-ESRGAN's prior, which reaches down to pi/5
+
+
+def test_stage2_gate_has_its_own_stream() -> None:
+ base = get_profile("p1_second_order")
+ assert isinstance(base, RealESRGANProfile) and base.stage2 is not None
+ almost = dataclasses.replace(base, stage2_prob=0.999999)
+ sometimes = dataclasses.replace(base, stage2_prob=0.3)
+ skipped = 0
+ for s in range(200):
+ always_plan = plan_degradation(base, (720, 1280), seed=s)
+ gated_plan = plan_degradation(sometimes, (720, 1280), seed=s)
+ if s < 5: # prob 1.0 and 0.999999 give identical plans: no plan-shifting draw on the main stream
+ assert always_plan.record() == plan_degradation(almost, (720, 1280), seed=s).record()
+ # Stage 1 never depends on the gate; only whether stage 2 (and what follows it) is drawn changes.
+ assert [o.record() for o in always_plan.ops if o.stage == "stage1"] == [
+ o.record() for o in gated_plan.ops if o.stage == "stage1"
+ ]
+ skipped += not any(o.stage == "stage2" for o in gated_plan.ops)
+ assert 120 < skipped < 160, skipped # 0.7 * 200 = 140 expected
+
+
+def test_profile_validation_and_resolution_clamps() -> None:
+ base = get_profile("p1_first_order")
+ assert isinstance(base, RealESRGANProfile)
+ for bad in (
+ dict(stage2_prob=1.5),
+ dict(stage2_prob=0.3), # no stage2 to gate: would be silently ignored
+ dict(min_intermediate_scale=1.5),
+ dict(min_resolution_factor=2.0, max_resolution_factor=1.5),
+ dict(max_resolution_factor=0.0),
+ dict(reference_longest_side=0),
+ ):
+ with pytest.raises(ValueError):
+ dataclasses.replace(base, **bad)
+ with pytest.raises(ValueError):
+ BlurConfig(sinc_cutoff_range=(math.pi, math.pi / 2)) # reversed
+ with pytest.raises(ValueError):
+ FinalBlockConfig(sinc_cutoff_range=(0.0, math.pi)) # omega_c = 0 is an all-NaN kernel
+ # [1.0, 1.5] keeps sub-reference inputs at the base ranges and stops growth past 1.5x; the inherited profile
+ # keeps scaling; with scaling off the clamps are ignored.
+ clamped = dataclasses.replace(base, name="p1_clamped", min_resolution_factor=1.0, max_resolution_factor=1.5)
+ unscaled = dataclasses.replace(
+ base, name="p1_unscaled", scale_kernels_with_resolution=False, max_resolution_factor=0.5
+ )
+ for seed in range(5):
+ assert _stage1_kernels(clamped, (360, 640), seed) == _stage1_kernels(clamped, (405, 720), seed)
+ assert _stage1_kernels(clamped, (1080, 1920), seed) == _stage1_kernels(clamped, (2160, 3840), seed)
+ assert _stage1_kernels(unscaled, (2160, 3840), seed) == _stage1_kernels(unscaled, (360, 640), seed)
+ assert any(_stage1_kernels(base, (1080, 1920), s) != _stage1_kernels(base, (2160, 3840), s) for s in range(5))
+
+
+def _blur_kernels(profile: str | Profile, hr_size: tuple[int, int], seed: int) -> list[int]:
+ return [o.params["kernel_size"] for o in plan_degradation(profile, hr_size, seed=seed).ops if o.op == "blur"]
+
+
+def test_regime_profiles_match_their_specification() -> None:
+ assert set(REGIME_PROFILES) == {f"{k}_{r}" for k in ("img", "vid") for r in ("clean", "mild", "moderate", "harsh")}
+ for mix, prefix in ((IMAGE_SR_DEFAULT_MIX, "img_"), (VIDEO_SR_DEFAULT_MIX, "vid_")):
+ assert abs(sum(mix.values()) - 1) < 1e-9 and all(k.startswith(prefix) and k in PROFILES for k in mix)
+ hr = _synthetic_clip(t=6, h=96, w=128) # [3,6,96,128]
+ for name, prof in REGIME_PROFILES.items():
+ out = degrade_hr_to_lr(hr, name, seed=1)
+ assert out.lr.shape == (3, 6, 48, 64) and out.lr.dtype == torch.uint8
+ if isinstance(prof, RealESRGANProfile): # declared resize ranges respect the floor, so it never binds
+ for stage in (prof.stage1, prof.stage2):
+ if stage is not None and stage.resize is not None:
+ assert stage.resize.scale_range[0] >= prof.min_intermediate_scale, name
+ # Video regimes never JPEG after the final resize (the codec is the compression term); vid_clean is a clean
+ # resize with an occasional clean H.264 re-encode.
+ for name in ("vid_mild", "vid_moderate", "vid_harsh"):
+ for seed in range(30):
+ ops_ = [o.op for o in plan_degradation(name, (720, 1280), seed=seed).ops]
+ assert "jpeg" not in ops_[max(i for i, o in enumerate(ops_) if o == "resize") :], (name, seed, ops_)
+ clean = [plan_degradation("vid_clean", (720, 1280), seed=s) for s in range(40)]
+ assert all([o.op for o in p.ops] == ["resize_clean"] for p in clean)
+ codecs = [p.codec.params for p in clean if p.codec is not None]
+ assert codecs and all(c["codec"] == "libx264" and 16 <= c["crf"] <= 20 for c in codecs)
+ # img_moderate runs its second stage 30% of the time.
+ plans = [plan_degradation("img_moderate", (720, 1280), seed=s) for s in range(200)]
+ n_stage2 = sum(any(o.stage == "stage2" for o in p.ops) for p in plans)
+ assert 40 < n_stage2 < 80, n_stage2
+ # Resolution scaling: 1280 px reference clamped to [1.0, 1.5], so 360p == 720p, 1080p is 1.5x, 4K == 1080p.
+ for seed in range(10):
+ k_720, k_1080 = _stage1_kernels("img_mild", (720, 1280), seed), _stage1_kernels("img_mild", (1080, 1920), seed)
+ assert k_1080 == [scale_kernel_size(k, 1.5, 41) for k in k_720]
+ assert _blur_kernels("img_mild", (360, 640), seed) == _blur_kernels("img_mild", (720, 1280), seed)
+ assert _blur_kernels("img_mild", (1080, 1920), seed) == _blur_kernels("img_mild", (2160, 3840), seed)
+
+
+def test_image_regimes_jpeg_in_the_final_block_with_declared_sinc_cutoffs() -> None:
+ # The JPEG sits in the final block: Real-ESRGAN's random order puts it after the final resize (on the LR grid)
+ # about half the time and just before it otherwise. Stage and final sinc cutoffs stay in the declared range.
+ at_lr = with_jpeg = 0
+ cutoffs: list[float] = []
+ for s in range(300):
+ plan = plan_degradation("img_mild", (1080, 1920), seed=s)
+ cutoffs += [o.params["omega_c"] for o in plan.ops if o.op == "blur" and o.params["kernel_type"] == "sinc"]
+ jpegs = [i for i, o in enumerate(plan.ops) if o.op == "jpeg"]
+ if jpegs:
+ with_jpeg += 1
+ at_lr += jpegs[-1] > max(i for i, o in enumerate(plan.ops) if o.op == "resize")
+ assert 0.3 < at_lr / with_jpeg < 0.7
+ assert cutoffs and min(cutoffs) >= math.pi / 2 - 1e-9 # img_mild declares (pi/2, pi)
diff --git a/cosmos_framework/data/generator/augmentors/hr_lr_degradation/diffjpeg.py b/cosmos_framework/data/generator/augmentors/hr_lr_degradation/diffjpeg.py
new file mode 100644
index 000000000..fdd3edaab
--- /dev/null
+++ b/cosmos_framework/data/generator/augmentors/hr_lr_degradation/diffjpeg.py
@@ -0,0 +1,163 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+"""Torch JPEG compression round trip that runs on any device.
+
+Adapted from DiffJPEG (MIT) https://github.com/mlomnitz/DiffJPEG through BasicSR and the
+Cosmos Transfer1 corruptors. Changes versus the Transfer1 copy: constant tables are buffers
+instead of parameters, nothing is pinned to CUDA or bfloat16, and quality can be a per-frame
+tensor. Padding to a multiple of 16 handles sizes that are not divisible by 8:
+https://dsp.stackexchange.com/questions/35339/jpeg-dct-padding/35343#35343
+"""
+
+from __future__ import annotations
+
+import numpy as np
+import torch
+import torch.nn as nn
+from torch.nn import functional as F
+
+_Y_TABLE = np.array(
+ [
+ [16, 11, 10, 16, 24, 40, 51, 61],
+ [12, 12, 14, 19, 26, 58, 60, 55],
+ [14, 13, 16, 24, 40, 57, 69, 56],
+ [14, 17, 22, 29, 51, 87, 80, 62],
+ [18, 22, 37, 56, 68, 109, 103, 77],
+ [24, 35, 55, 64, 81, 104, 113, 92],
+ [49, 64, 78, 87, 103, 121, 120, 101],
+ [72, 92, 95, 98, 112, 100, 103, 99],
+ ],
+ dtype=np.float32,
+).T # [8,8]
+
+_C_TABLE = np.full((8, 8), 99, dtype=np.float32) # [8,8]
+_C_TABLE[:4, :4] = np.array([[17, 18, 24, 47], [18, 21, 26, 66], [24, 26, 56, 99], [47, 66, 99, 99]]).T
+
+
+def quality_to_factor(quality: torch.Tensor) -> torch.Tensor: # quality: [B] in (0,100], returns [B]
+ """Map JPEG quality to the quantisation-table multiplier (libjpeg convention)."""
+ low = 5000.0 / quality # [B]
+ high = 200.0 - quality * 2 # [B]
+ return torch.where(quality < 50, low, high) / 100.0 # [B]
+
+
+def _dct_matrix() -> np.ndarray: # returns [8,8], C[u,x] = cos((2x+1) u pi / 16)
+ u = np.arange(8, dtype=np.float32).reshape(8, 1) # [8,1]
+ x = np.arange(8, dtype=np.float32).reshape(1, 8) # [1,8]
+ return np.cos((2 * x + 1) * u * np.pi / 16).astype(np.float32) # [8,8]
+
+
+def _alpha_outer() -> np.ndarray: # returns [8,8]
+ alpha = np.array([1.0 / np.sqrt(2)] + [1] * 7) # [8]
+ return np.outer(alpha, alpha).astype(np.float32) # [8,8]
+
+
+def _block_split(image: torch.Tensor) -> torch.Tensor: # image: [B,H,W], returns [B,H*W/64,8,8]
+ batch_size, height, _ = image.shape
+ blocks = image.view(batch_size, height // 8, 8, -1, 8) # [B,H/8,8,W/8,8]
+ blocks = blocks.permute(0, 1, 3, 2, 4) # [B,H/8,W/8,8,8]
+ return blocks.contiguous().view(batch_size, -1, 8, 8) # [B,H*W/64,8,8]
+
+
+def _block_merge(blocks: torch.Tensor, height: int, width: int) -> torch.Tensor: # blocks: [B,N,8,8], returns [B,H,W]
+ batch_size = blocks.shape[0]
+ image = blocks.view(batch_size, height // 8, width // 8, 8, 8) # [B,H/8,W/8,8,8]
+ image = image.permute(0, 1, 3, 2, 4) # [B,H/8,8,W/8,8]
+ return image.contiguous().view(batch_size, height, width) # [B,H,W]
+
+
+class DiffJPEG(nn.Module):
+ """JPEG encode-decode simulator with 4:2:0 chroma subsampling.
+
+ Args:
+ differentiable: use a smooth rounding surrogate instead of ``torch.round``. Degradation
+ for training data does not need gradients, so the default is hard rounding.
+ """
+
+ def __init__(self, differentiable: bool = False) -> None:
+ super().__init__()
+ self.differentiable = differentiable
+ rgb2ycc = np.array(
+ [[0.299, 0.587, 0.114], [-0.168736, -0.331264, 0.5], [0.5, -0.418688, -0.081312]], dtype=np.float32
+ ).T # [3,3]
+ ycc2rgb = np.array([[1.0, 0.0, 1.402], [1, -0.344136, -0.714136], [1, 1.772, 0]], dtype=np.float32).T # [3,3]
+ self.register_buffer("rgb2ycc", torch.from_numpy(rgb2ycc), persistent=False) # [3,3]
+ self.register_buffer("ycc2rgb", torch.from_numpy(ycc2rgb), persistent=False) # [3,3]
+ self.register_buffer("ycc_shift", torch.tensor([0.0, 128.0, 128.0]), persistent=False) # [3]
+ self.register_buffer("y_table", torch.from_numpy(_Y_TABLE), persistent=False) # [8,8]
+ self.register_buffer("c_table", torch.from_numpy(_C_TABLE), persistent=False) # [8,8]
+ self.register_buffer("dct_mat", torch.from_numpy(_dct_matrix()), persistent=False) # [8,8]
+ self.register_buffer("alpha", torch.from_numpy(_alpha_outer()), persistent=False) # [8,8]
+
+ def _round(self, x: torch.Tensor) -> torch.Tensor: # x: [...], returns [...]
+ if self.differentiable:
+ return torch.round(x) + (x - torch.round(x)) ** 3
+ return torch.round(x)
+
+ def _quantize(self, blocks: torch.Tensor, table: torch.Tensor, factor: torch.Tensor) -> torch.Tensor:
+ # blocks: [B,N,8,8]; table: [8,8]; factor: [B]; returns [B,N,8,8]
+ scaled_table = table[None, None] * factor.view(-1, 1, 1, 1) # [B,1,8,8]
+ return self._round(blocks / scaled_table) # [B,N,8,8]
+
+ def _dequantize(self, blocks: torch.Tensor, table: torch.Tensor, factor: torch.Tensor) -> torch.Tensor:
+ # blocks: [B,N,8,8]; table: [8,8]; factor: [B]; returns [B,N,8,8]
+ scaled_table = table[None, None] * factor.view(-1, 1, 1, 1) # [B,1,8,8]
+ return blocks * scaled_table # [B,N,8,8]
+
+ def _forward_dct(self, plane: torch.Tensor) -> torch.Tensor: # plane: [B,H,W], returns [B,H*W/64,8,8]
+ # Separable form of the 4D basis contraction: Y = C X C^T with C[u,x] = cos((2x+1) u pi / 16).
+ blocks = _block_split(plane) - 128 # [B,N,8,8]
+ return 0.25 * self.alpha * (self.dct_mat @ blocks @ self.dct_mat.T) # [B,N,8,8]
+
+ def _inverse_dct(self, blocks: torch.Tensor, height: int, width: int) -> torch.Tensor:
+ # blocks: [B,N,8,8], returns [B,H,W]
+ blocks = blocks * self.alpha # [B,N,8,8]
+ blocks = 0.25 * (self.dct_mat.T @ blocks @ self.dct_mat) + 128 # [B,N,8,8]
+ return _block_merge(blocks, height, width) # [B,H,W]
+
+ def forward(self, x: torch.Tensor, quality: torch.Tensor | float) -> torch.Tensor:
+ """Compress and decompress a batch of RGB frames.
+
+ Args:
+ x: frames in [0, 1], shape ``[B,3,H,W]``, any float dtype.
+ quality: JPEG quality in (0, 100], scalar or ``[B]`` tensor.
+
+ Returns:
+ Reconstructed frames in [0, 1], shape ``[B,3,H,W]``, dtype of ``x``.
+ """
+ batch_size, _, height, width = x.shape
+ in_dtype = x.dtype
+ x = x.float() # [B,3,H,W]
+ quality_t = torch.as_tensor(quality, dtype=torch.float32, device=x.device).reshape(-1) # [1] or [B]
+ if quality_t.numel() == 1:
+ quality_t = quality_t.expand(batch_size) # [B]
+ factor = quality_to_factor(quality_t) # [B]
+
+ h_pad = (16 - height % 16) % 16
+ w_pad = (16 - width % 16) % 16
+ x = F.pad(x, (0, w_pad, 0, h_pad), mode="constant", value=0) # [B,3,Hp,Wp]
+ padded_h, padded_w = height + h_pad, width + w_pad
+
+ ycc = torch.tensordot(x.permute(0, 2, 3, 1) * 255.0, self.rgb2ycc, dims=1) + self.ycc_shift # [B,Hp,Wp,3]
+ y = ycc[..., 0] # [B,Hp,Wp]
+ cb = F.avg_pool2d(ycc[..., 1:2].permute(0, 3, 1, 2), kernel_size=2, stride=2)[:, 0] # [B,Hp/2,Wp/2]
+ cr = F.avg_pool2d(ycc[..., 2:3].permute(0, 3, 1, 2), kernel_size=2, stride=2)[:, 0] # [B,Hp/2,Wp/2]
+
+ y_q = self._quantize(self._forward_dct(y), self.y_table, factor) # [B,N,8,8]
+ cb_q = self._quantize(self._forward_dct(cb), self.c_table, factor) # [B,N/4,8,8]
+ cr_q = self._quantize(self._forward_dct(cr), self.c_table, factor) # [B,N/4,8,8]
+
+ y_rec = self._inverse_dct(self._dequantize(y_q, self.y_table, factor), padded_h, padded_w) # [B,Hp,Wp]
+ cb_rec = self._inverse_dct(
+ self._dequantize(cb_q, self.c_table, factor), padded_h // 2, padded_w // 2
+ ) # [B,Hp/2,Wp/2]
+ cr_rec = self._inverse_dct(
+ self._dequantize(cr_q, self.c_table, factor), padded_h // 2, padded_w // 2
+ ) # [B,Hp/2,Wp/2]
+ cb_up = cb_rec.repeat_interleave(2, dim=1).repeat_interleave(2, dim=2) # [B,Hp,Wp]
+ cr_up = cr_rec.repeat_interleave(2, dim=1).repeat_interleave(2, dim=2) # [B,Hp,Wp]
+
+ ycc_rec = torch.stack([y_rec, cb_up, cr_up], dim=-1) - self.ycc_shift # [B,Hp,Wp,3]
+ rgb = torch.tensordot(ycc_rec, self.ycc2rgb, dims=1).permute(0, 3, 1, 2) # [B,3,Hp,Wp]
+ rgb = rgb.clamp(0.0, 255.0) / 255.0 # [B,3,Hp,Wp]
+ return rgb[:, :, :height, :width].to(in_dtype) # [B,3,H,W]
diff --git a/cosmos_framework/data/generator/augmentors/hr_lr_degradation/kernels.py b/cosmos_framework/data/generator/augmentors/hr_lr_degradation/kernels.py
new file mode 100644
index 000000000..6c744557d
--- /dev/null
+++ b/cosmos_framework/data/generator/augmentors/hr_lr_degradation/kernels.py
@@ -0,0 +1,182 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+"""Blur kernel generators for the Real-ESRGAN style degradation pipeline.
+
+Adapted from BasicSR ``basicsr/data/degradations.py`` (Apache-2.0)
+https://github.com/XPixelGroup/BasicSR/blob/8d56e3a045f9fb3e1d8872f92ee4a4f07f886b0a/basicsr/data/degradations.py
+via the Cosmos Transfer1 corruptors. Every random draw goes through an explicit
+``numpy.random.Generator`` so a clip's kernel is reproducible from its seed.
+"""
+
+from __future__ import annotations
+
+import math
+from typing import Sequence
+
+import numpy as np
+from scipy import special
+
+KERNEL_TYPES = ("iso", "aniso", "generalized_iso", "generalized_aniso", "plateau_iso", "plateau_aniso")
+
+
+def sigma_matrix2(sig_x: float, sig_y: float, theta: float) -> np.ndarray: # returns [2,2]
+ """Rotated covariance matrix of a bivariate Gaussian."""
+ d_matrix = np.array([[sig_x**2, 0], [0, sig_y**2]]) # [2,2]
+ u_matrix = np.array([[np.cos(theta), -np.sin(theta)], [np.sin(theta), np.cos(theta)]]) # [2,2]
+ return np.dot(u_matrix, np.dot(d_matrix, u_matrix.T)) # [2,2]
+
+
+def mesh_grid(kernel_size: int) -> np.ndarray: # returns [K,K,2]
+ """Coordinate grid centred at zero."""
+ ax = np.arange(-kernel_size // 2 + 1.0, kernel_size // 2 + 1.0) # [K]
+ xx, yy = np.meshgrid(ax, ax) # [K,K] each
+ return np.stack([xx, yy], axis=-1) # [K,K,2]
+
+
+def _quadratic_form(sigma_matrix: np.ndarray, grid: np.ndarray) -> np.ndarray: # returns [K,K]
+ inverse_sigma = np.linalg.inv(sigma_matrix) # [2,2]
+ return np.sum(np.dot(grid, inverse_sigma) * grid, 2) # [K,K]
+
+
+def _sigma_matrix(sig_x: float, sig_y: float, theta: float, isotropic: bool) -> np.ndarray: # returns [2,2]
+ if isotropic:
+ return np.array([[sig_x**2, 0], [0, sig_x**2]]) # [2,2]
+ return sigma_matrix2(sig_x, sig_y, theta) # [2,2]
+
+
+def bivariate_gaussian(
+ kernel_size: int, sig_x: float, sig_y: float, theta: float, isotropic: bool = True
+) -> np.ndarray: # returns [K,K]
+ """Normalised isotropic or anisotropic Gaussian kernel."""
+ grid = mesh_grid(kernel_size) # [K,K,2]
+ kernel = np.exp(-0.5 * _quadratic_form(_sigma_matrix(sig_x, sig_y, theta, isotropic), grid)) # [K,K]
+ return kernel / np.sum(kernel) # [K,K]
+
+
+def bivariate_generalized_gaussian(
+ kernel_size: int, sig_x: float, sig_y: float, theta: float, beta: float, isotropic: bool = True
+) -> np.ndarray: # returns [K,K]
+ """Normalised generalized Gaussian kernel; ``beta == 1`` is the plain Gaussian."""
+ grid = mesh_grid(kernel_size) # [K,K,2]
+ q = _quadratic_form(_sigma_matrix(sig_x, sig_y, theta, isotropic), grid) # [K,K]
+ kernel = np.exp(-0.5 * np.power(q, beta)) # [K,K]
+ return kernel / np.sum(kernel) # [K,K]
+
+
+def bivariate_plateau(
+ kernel_size: int, sig_x: float, sig_y: float, theta: float, beta: float, isotropic: bool = True
+) -> np.ndarray: # returns [K,K]
+ """Normalised plateau-shaped kernel ``1 / (1 + q^beta)``."""
+ grid = mesh_grid(kernel_size) # [K,K,2]
+ q = _quadratic_form(_sigma_matrix(sig_x, sig_y, theta, isotropic), grid) # [K,K]
+ kernel = np.reciprocal(np.power(q, beta) + 1) # [K,K]
+ return kernel / np.sum(kernel) # [K,K]
+
+
+def _sample_sigma_rotation(
+ rng: np.random.Generator,
+ sigma_x_range: Sequence[float],
+ sigma_y_range: Sequence[float],
+ rotation_range: Sequence[float],
+ isotropic: bool,
+) -> tuple[float, float, float]:
+ assert sigma_x_range[0] < sigma_x_range[1], "Wrong sigma_x_range."
+ sigma_x = float(rng.uniform(sigma_x_range[0], sigma_x_range[1]))
+ if isotropic:
+ return sigma_x, sigma_x, 0.0
+ assert sigma_y_range[0] < sigma_y_range[1], "Wrong sigma_y_range."
+ assert rotation_range[0] < rotation_range[1], "Wrong rotation_range."
+ sigma_y = float(rng.uniform(sigma_y_range[0], sigma_y_range[1]))
+ rotation = float(rng.uniform(rotation_range[0], rotation_range[1]))
+ return sigma_x, sigma_y, rotation
+
+
+def _sample_beta(rng: np.random.Generator, beta_range: Sequence[float]) -> float:
+ # Real-ESRGAN draws below or above 1 with equal probability so both regimes are covered.
+ if rng.uniform() < 0.5:
+ return float(rng.uniform(beta_range[0], 1))
+ return float(rng.uniform(1, beta_range[1]))
+
+
+def random_mixed_kernel(
+ rng: np.random.Generator,
+ kernel_list: Sequence[str],
+ kernel_prob: Sequence[float],
+ kernel_size: int,
+ sigma_x_range: Sequence[float],
+ sigma_y_range: Sequence[float],
+ rotation_range: Sequence[float] = (-math.pi, math.pi),
+ betag_range: Sequence[float] = (0.5, 8),
+ betap_range: Sequence[float] = (0.5, 8),
+) -> tuple[np.ndarray, dict]: # returns ([K,K], sampled parameters)
+ """Sample one kernel type from ``kernel_list`` and its parameters, seeded by ``rng``."""
+ assert kernel_size % 2 == 1, "Kernel size must be an odd number."
+ assert len(kernel_list) == len(kernel_prob), "kernel_list and kernel_prob must have equal length."
+ prob = np.asarray(kernel_prob, dtype=np.float64) # [N]
+ kernel_type = str(rng.choice(np.asarray(kernel_list), p=prob / prob.sum()))
+ if kernel_type not in KERNEL_TYPES:
+ raise ValueError(f"Unknown kernel type {kernel_type}; supported: {KERNEL_TYPES}")
+ isotropic = kernel_type.endswith("iso") and not kernel_type.endswith("aniso")
+ sigma_x, sigma_y, rotation = _sample_sigma_rotation(rng, sigma_x_range, sigma_y_range, rotation_range, isotropic)
+ info = {"kernel_type": kernel_type, "kernel_size": kernel_size, "sigma_x": sigma_x, "sigma_y": sigma_y}
+ if not isotropic:
+ info["rotation"] = rotation
+ if kernel_type in ("iso", "aniso"):
+ kernel = bivariate_gaussian(kernel_size, sigma_x, sigma_y, rotation, isotropic) # [K,K]
+ elif kernel_type in ("generalized_iso", "generalized_aniso"):
+ beta = _sample_beta(rng, betag_range)
+ info["beta"] = beta
+ kernel = bivariate_generalized_gaussian(kernel_size, sigma_x, sigma_y, rotation, beta, isotropic) # [K,K]
+ else:
+ beta = _sample_beta(rng, betap_range)
+ info["beta"] = beta
+ kernel = bivariate_plateau(kernel_size, sigma_x, sigma_y, rotation, beta, isotropic) # [K,K]
+ return kernel, info
+
+
+def circular_lowpass_kernel(cutoff: float, kernel_size: int, pad_to: int = 0) -> np.ndarray: # returns [P,P]
+ """2D circularly symmetric sinc low-pass filter.
+
+ Reference: https://dsp.stackexchange.com/questions/58301/2-d-circularly-symmetric-low-pass-filter
+
+ Args:
+ cutoff: cutoff frequency in radians; ``pi`` is the maximum.
+ kernel_size: odd spatial size of the kernel.
+ pad_to: zero-pad the kernel to this odd size when larger than ``kernel_size``.
+ """
+ assert kernel_size % 2 == 1, "Kernel size must be an odd number."
+ centre = (kernel_size - 1) / 2
+ with np.errstate(divide="ignore", invalid="ignore"):
+ kernel = np.fromfunction(
+ lambda x, y: cutoff
+ * special.j1(cutoff * np.sqrt((x - centre) ** 2 + (y - centre) ** 2))
+ / (2 * np.pi * np.sqrt((x - centre) ** 2 + (y - centre) ** 2)),
+ [kernel_size, kernel_size],
+ ) # [K,K]
+ kernel[(kernel_size - 1) // 2, (kernel_size - 1) // 2] = cutoff**2 / (4 * np.pi)
+ kernel = kernel / np.sum(kernel) # [K,K]
+ if pad_to > kernel_size:
+ pad_size = (pad_to - kernel_size) // 2
+ kernel = np.pad(kernel, ((pad_size, pad_size), (pad_size, pad_size))) # [P,P]
+ return kernel
+
+
+def random_sinc_kernel(
+ rng: np.random.Generator, kernel_size: int, cutoff_range: Sequence[float] | None = None
+) -> tuple[np.ndarray, dict]: # returns ([K,K], info)
+ """Sinc kernel. ``cutoff_range`` bounds omega_c; ``None`` uses the Real-ESRGAN prior, which widens the range for
+ kernels of size 13 and above."""
+ if cutoff_range is None:
+ cutoff_range = (np.pi / 3, np.pi) if kernel_size < 13 else (np.pi / 5, np.pi)
+ omega_c = float(rng.uniform(*cutoff_range))
+ kernel = circular_lowpass_kernel(omega_c, kernel_size) # [K,K]
+ return kernel, {"kernel_type": "sinc", "kernel_size": kernel_size, "omega_c": omega_c}
+
+
+def scale_kernel_size(kernel_size: int, factor: float, max_kernel_size: int) -> int:
+ """Scale an odd kernel size by ``factor`` and return the nearest odd size within ``[3, max_kernel_size]``."""
+ scaled = int(round(kernel_size * factor))
+ scaled = max(3, min(scaled, max_kernel_size))
+ if scaled % 2 == 0:
+ scaled = scaled - 1 if scaled >= max_kernel_size else scaled + 1
+ return scaled
diff --git a/cosmos_framework/data/generator/augmentors/hr_lr_degradation/ops.py b/cosmos_framework/data/generator/augmentors/hr_lr_degradation/ops.py
new file mode 100644
index 000000000..bc4a234ba
--- /dev/null
+++ b/cosmos_framework/data/generator/augmentors/hr_lr_degradation/ops.py
@@ -0,0 +1,213 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+"""Pixel-space degradation primitives.
+
+All functions take frames as ``[T,C,H,W]`` float tensors in [0, 1] on any device and return the
+same layout. Parameters are passed in explicitly (they are sampled once per clip by
+``degrade.py``), so the only randomness here is the per-pixel noise realisation, drawn from a
+caller-owned ``torch.Generator`` that lives on the frame device. Outputs are therefore
+deterministic per device; CPU and GPU agree on every sampled parameter but not bit-for-bit on
+noise.
+
+Performance notes (1080p, measured on an L4 + 4 CPU threads):
+- Blur runs through FFT convolution once the kernel is larger than ``_DIRECT_CONV_MAX_KERNEL``,
+ which makes cost independent of kernel size (direct conv2d at k=61 was 40x slower on CPU).
+- JPEG uses libjpeg through OpenCV on CPU (about 30 ms per 1080p frame) and the torch DiffJPEG
+ simulator on GPU.
+- Exact Poisson sampling costs about 300 ms per 1080p frame on CPU with either torch or numpy, so
+ ``poisson_mode="gaussian_approx"`` (signal-dependent Gaussian) is provided for CPU workers.
+
+Noise code follows BasicSR ``basicsr/data/degradations.py`` (Apache-2.0).
+"""
+
+from __future__ import annotations
+
+import cv2
+import numpy as np
+import torch
+import torch.nn.functional as F
+
+from cosmos_framework.data.generator.augmentors.hr_lr_degradation.diffjpeg import DiffJPEG
+
+_ANTIALIAS_KERNELS = {"bicubic_antialias": "bicubic", "bilinear_antialias": "bilinear"}
+_DIRECT_CONV_MAX_KERNEL = 9
+JPEG_BACKENDS = ("auto", "cv2", "diffjpeg")
+POISSON_MODES = ("auto", "exact", "gaussian_approx")
+
+
+def make_generator(seed: int, device: torch.device | str) -> torch.Generator:
+ """Seeded generator on ``device`` (CUDA generators must live on the tensor device)."""
+ dev = torch.device(device)
+ gen = torch.Generator(device=dev if dev.type == "cuda" else "cpu")
+ gen.manual_seed(int(seed) & 0x7FFF_FFFF_FFFF_FFFF)
+ return gen
+
+
+def _filter2d_direct(frames: torch.Tensor, kernel: torch.Tensor) -> torch.Tensor: # frames: [T,C,H,W]; kernel: [K,K]
+ k = kernel.shape[-1]
+ t, c, h, w = frames.shape
+ pad = k // 2
+ padded = F.pad(frames.reshape(t * c, 1, h, w), (pad, pad, pad, pad), mode="reflect") # [T*C,1,H+2p,W+2p]
+ out = F.conv2d(padded, kernel.reshape(1, 1, k, k)) # [T*C,1,H,W]
+ return out.reshape(t, c, h, w) # [T,C,H,W]
+
+
+def _filter2d_fft(frames: torch.Tensor, kernel: torch.Tensor) -> torch.Tensor: # frames: [T,C,H,W]; kernel: [K,K]
+ """Same correlation as ``_filter2d_direct`` (reflect padding, kernel not flipped) via rfft2."""
+ k = kernel.shape[-1]
+ h, w = frames.shape[-2:]
+ pad = k // 2
+ padded = F.pad(frames, (pad, pad, pad, pad), mode="reflect") # [T,C,Hp,Wp]
+ hp, wp = padded.shape[-2:]
+ # conv2d computes correlation; FFT multiplication computes convolution, so flip the kernel.
+ kernel_flipped = torch.flip(kernel, dims=(0, 1)) # [K,K]
+ kernel_padded = torch.zeros(hp, wp, device=frames.device, dtype=frames.dtype) # [Hp,Wp]
+ kernel_padded[:k, :k] = kernel_flipped
+ kernel_padded = torch.roll(kernel_padded, shifts=(-pad, -pad), dims=(0, 1)) # [Hp,Wp] centred at origin
+ spectrum = torch.fft.rfft2(padded) * torch.fft.rfft2(kernel_padded) # [T,C,Hp,Wp/2+1]
+ out = torch.fft.irfft2(spectrum, s=(hp, wp)) # [T,C,Hp,Wp]
+ return out[..., pad : pad + h, pad : pad + w] # [T,C,H,W]
+
+
+def filter2d(frames: torch.Tensor, kernel: torch.Tensor) -> torch.Tensor: # frames: [T,C,H,W]; kernel: [K,K]
+ """Depthwise 2D correlation with one shared odd kernel and reflect padding (torch ``cv2.filter2D``)."""
+ k = kernel.shape[-1]
+ if k % 2 != 1:
+ raise ValueError(f"Kernel size must be odd, got {k}")
+ kernel = kernel.to(dtype=frames.dtype, device=frames.device) # [K,K]
+ if k <= _DIRECT_CONV_MAX_KERNEL:
+ return _filter2d_direct(frames, kernel) # [T,C,H,W]
+ return _filter2d_fft(frames, kernel) # [T,C,H,W]
+
+
+def blur(frames: torch.Tensor, kernel: np.ndarray) -> torch.Tensor: # frames: [T,C,H,W]; kernel: [K,K]
+ """Blur every frame with the same kernel."""
+ kernel_t = torch.from_numpy(np.ascontiguousarray(kernel, dtype=np.float32)) # [K,K]
+ return filter2d(frames, kernel_t) # [T,C,H,W]
+
+
+def resize(frames: torch.Tensor, size: tuple[int, int], mode: str) -> torch.Tensor: # frames: [T,C,H,W]
+ """Resize with the Real-ESRGAN interpolation modes (no antialiasing; aliasing is part of the degradation)."""
+ if mode not in ("area", "bilinear", "bicubic"):
+ raise ValueError(f"Unsupported resize mode {mode!r}")
+ if tuple(frames.shape[-2:]) == tuple(size):
+ return frames
+ if mode == "area":
+ return F.interpolate(frames, size=size, mode="area") # [T,C,h,w]
+ return F.interpolate(frames, size=size, mode=mode, align_corners=False) # [T,C,h,w]
+
+
+def resize_clean(frames: torch.Tensor, size: tuple[int, int], kernel: str) -> torch.Tensor: # frames: [T,C,H,W]
+ """Antialiased resize used for the clean P0 profile and for benchmark-style LR construction."""
+ if kernel == "area":
+ return F.interpolate(frames, size=size, mode="area") # [T,C,h,w]
+ if kernel in _ANTIALIAS_KERNELS:
+ return F.interpolate(
+ frames, size=size, mode=_ANTIALIAS_KERNELS[kernel], align_corners=False, antialias=True
+ ) # [T,C,h,w]
+ raise ValueError(f"Unsupported clean resize kernel {kernel!r}")
+
+
+def _randn_like_shape(shape: tuple[int, ...], gen: torch.Generator, like: torch.Tensor) -> torch.Tensor:
+ # returns [*shape] on like.device
+ return torch.randn(shape, generator=gen, dtype=like.dtype, device=like.device)
+
+
+def add_gaussian_noise(
+ frames: torch.Tensor, sigma: float, gray: bool, gen: torch.Generator
+) -> torch.Tensor: # frames: [T,C,H,W]
+ """Additive Gaussian noise with standard deviation ``sigma`` in 8-bit units."""
+ t, c, h, w = frames.shape
+ if gray:
+ noise = _randn_like_shape((t, 1, h, w), gen, frames).expand(t, c, h, w) # [T,C,H,W]
+ else:
+ noise = _randn_like_shape((t, c, h, w), gen, frames) # [T,C,H,W]
+ return (frames + noise * (sigma / 255.0)).clamp_(0.0, 1.0) # [T,C,H,W]
+
+
+def _levels_per_frame(img_q: torch.Tensor) -> torch.Tensor: # img_q: [T,C,H,W] quantised to 1/255; returns [T,1,1,1]
+ """Number of distinct 8-bit levels per frame rounded up to a power of two (BasicSR ``vals``)."""
+ t = img_q.shape[0]
+ codes = (img_q * 255.0).round().long().reshape(t, -1) # [T,C*H*W]
+ counts = [int(torch.bincount(codes[i], minlength=256).count_nonzero().item()) for i in range(t)]
+ levels = [2 ** int(np.ceil(np.log2(max(n, 1)))) for n in counts]
+ return img_q.new_tensor(levels).view(t, 1, 1, 1) # [T,1,1,1]
+
+
+def _poisson_noise(img: torch.Tensor, gen: torch.Generator, mode: str) -> torch.Tensor: # img: [T,C,H,W]
+ """Shot noise whose rate scales with the number of distinct 8-bit levels per frame, as in BasicSR."""
+ img_q = (img * 255.0).round().clamp(0.0, 255.0) / 255.0 # [T,C,H,W]
+ vals = _levels_per_frame(img_q) # [T,1,1,1]
+ rates = img_q * vals # [T,C,H,W]
+ if mode == "exact":
+ sampled = torch.poisson(rates, generator=gen) # [T,C,H,W]
+ elif mode == "gaussian_approx":
+ # Poisson(lambda) ~ N(lambda, lambda) for moderate lambda; keeps the signal-dependent variance.
+ sampled = rates + rates.sqrt() * _randn_like_shape(tuple(rates.shape), gen, rates) # [T,C,H,W]
+ sampled = sampled.round().clamp_(min=0.0) # [T,C,H,W]
+ else:
+ raise ValueError(f"Unknown poisson mode {mode!r}; expected one of {POISSON_MODES[1:]}")
+ return sampled / vals - img_q # [T,C,H,W]
+
+
+def resolve_poisson_mode(mode: str, device: torch.device) -> str:
+ if mode == "auto":
+ return "exact" if device.type == "cuda" else "gaussian_approx"
+ if mode not in POISSON_MODES:
+ raise ValueError(f"Unknown poisson mode {mode!r}")
+ return mode
+
+
+def add_poisson_noise(
+ frames: torch.Tensor, scale: float, gray: bool, gen: torch.Generator, mode: str = "auto"
+) -> torch.Tensor: # frames: [T,C,H,W]
+ """Poisson (shot) noise scaled by ``scale``; grey noise is computed on the luminance and shared across channels."""
+ mode = resolve_poisson_mode(mode, frames.device)
+ t, c, h, w = frames.shape
+ if gray:
+ weights = frames.new_tensor([0.299, 0.587, 0.114]).view(1, 3, 1, 1) # [1,3,1,1]
+ luma = (frames * weights).sum(dim=1, keepdim=True) # [T,1,H,W]
+ noise = _poisson_noise(luma, gen, mode).expand(t, c, h, w) # [T,C,H,W]
+ else:
+ noise = _poisson_noise(frames, gen, mode) # [T,C,H,W]
+ return (frames + noise * scale).clamp_(0.0, 1.0) # [T,C,H,W]
+
+
+def _jpeg_cv2(frames: torch.Tensor, quality: float) -> torch.Tensor: # frames: [T,C,H,W] CPU float in [0,1]
+ """libjpeg round trip per frame through OpenCV (RGB <-> BGR handled here)."""
+ encode_param = [int(cv2.IMWRITE_JPEG_QUALITY), int(round(quality))]
+ frames_u8 = (frames.clamp(0.0, 1.0) * 255.0).round().to(torch.uint8).permute(0, 2, 3, 1).numpy() # [T,H,W,C]
+ out = np.empty_like(frames_u8) # [T,H,W,C]
+ for i in range(frames_u8.shape[0]):
+ bgr = cv2.cvtColor(np.ascontiguousarray(frames_u8[i]), cv2.COLOR_RGB2BGR) # [H,W,C]
+ ok, encoded = cv2.imencode(".jpg", bgr, encode_param)
+ if not ok:
+ raise RuntimeError("cv2.imencode failed")
+ decoded = cv2.imdecode(encoded, cv2.IMREAD_COLOR) # [H,W,C]
+ out[i] = cv2.cvtColor(decoded, cv2.COLOR_BGR2RGB)
+ return torch.from_numpy(out).permute(0, 3, 1, 2).to(frames.dtype) / 255.0 # [T,C,H,W]
+
+
+def resolve_jpeg_backend(backend: str, device: torch.device) -> str:
+ if backend == "auto":
+ return "cv2" if device.type == "cpu" else "diffjpeg"
+ if backend not in JPEG_BACKENDS:
+ raise ValueError(f"Unknown JPEG backend {backend!r}")
+ if backend == "cv2" and device.type != "cpu":
+ raise ValueError("JPEG backend 'cv2' requires CPU tensors")
+ return backend
+
+
+def jpeg(frames: torch.Tensor, quality: float, jpeger: DiffJPEG, backend: str = "auto") -> torch.Tensor:
+ # frames: [T,C,H,W]; returns [T,C,H,W]
+ """JPEG round trip at one quality for the whole chunk."""
+ backend = resolve_jpeg_backend(backend, frames.device)
+ if backend == "cv2":
+ return _jpeg_cv2(frames, quality) # [T,C,H,W]
+ if jpeger.y_table.device != frames.device:
+ jpeger.to(frames.device)
+ return jpeger(frames.clamp(0.0, 1.0), quality=quality) # [T,C,H,W]
+
+
+def to_uint8(frames: torch.Tensor) -> torch.Tensor: # frames: [T,C,H,W] float in [0,1], returns uint8
+ return (frames.clamp(0.0, 1.0) * 255.0).round().to(torch.uint8) # [T,C,H,W]
diff --git a/cosmos_framework/data/generator/augmentors/hr_lr_degradation/packing_test.py b/cosmos_framework/data/generator/augmentors/hr_lr_degradation/packing_test.py
new file mode 100644
index 000000000..9219772a7
--- /dev/null
+++ b/cosmos_framework/data/generator/augmentors/hr_lr_degradation/packing_test.py
@@ -0,0 +1,65 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+"""CPU dry run of sequence packing for (LR, HR) samples whose two vision items have different latent grids."""
+
+import pytest
+import torch
+
+from cosmos_framework.model.generator.utils.data_and_condition import GenerationDataClean
+from cosmos_framework.data.generator.sequence_packing import SequencePlan, pack_input_sequence
+
+pytestmark = [pytest.mark.L0, pytest.mark.CPU]
+
+_SPECIAL_TOKENS = {"eos_token_id": 151645, "start_of_generation": 151652, "end_of_generation": 151653}
+_LATENT_C, _SPATIAL, _TEMPORAL, _PATCH = 16, 16, 4, 2
+
+
+def _latent(frames: int, h: int, w: int) -> torch.Tensor: # returns [1,C,T',H',W']
+ return torch.randn(1, _LATENT_C, 1 + (frames - 1) // _TEMPORAL, h // _SPATIAL, w // _SPATIAL)
+
+
+def _pack(share: bool, frames: int = 9, hr: tuple[int, int] = (480, 832), lr: tuple[int, int] = (240, 416)):
+ plan = SequencePlan(
+ has_text=True, has_vision=True, condition_frame_indexes_vision=[0], share_vision_temporal_positions=share
+ )
+ gen = GenerationDataClean(
+ batch_size=1,
+ is_image_batch=False,
+ x0_tokens_vision=[_latent(frames, *lr), _latent(frames, *hr)],
+ num_vision_items_per_sample=[2],
+ fps_vision=torch.tensor([24.0]),
+ )
+ return pack_input_sequence(
+ sequence_plans=[plan],
+ input_text_indexes=[[5, 6, 7, 8]],
+ gen_data_clean=gen,
+ input_timesteps=torch.rand(1),
+ special_tokens=_SPECIAL_TOKENS,
+ latent_patch_size=_PATCH,
+ temporal_compression_factor=_TEMPORAL,
+ )
+
+
+def test_two_item_sr_sample_packs_without_shared_temporal_grid() -> None:
+ packed = _pack(share=False)
+
+ def _grid(h: int, w: int) -> tuple[int, int]:
+ # Latent dims that are not a multiple of the patch size are padded up by the packer (15 -> 8 tokens).
+ return -(-(h // _SPATIAL) // _PATCH), -(-(w // _SPATIAL) // _PATCH)
+
+ t_latent = 1 + 8 // _TEMPORAL
+ lr_grid, hr_grid = _grid(240, 416), _grid(480, 832) # (8,13), (15,26)
+ assert packed.vision is not None
+ # Both items are in the packed vision stream with their own grids; total vision tokens = LR + HR.
+ assert len(packed.vision.token_shapes) == 2
+ lr_shape, hr_shape = packed.vision.token_shapes
+ assert tuple(lr_shape[-2:]) == lr_grid and tuple(hr_shape[-2:]) == hr_grid
+ total = sum(int(torch.tensor(s[-3:]).prod()) for s in packed.vision.token_shapes)
+ assert total == t_latent * (lr_grid[0] * lr_grid[1] + hr_grid[0] * hr_grid[1])
+ # Only the HR (last) item is generated; the LR item is pure conditioning.
+ assert bool(packed.vision.condition_mask[0].all()) and not bool(packed.vision.condition_mask[1].all())
+
+
+def test_two_item_sr_sample_with_shared_grid_is_rejected() -> None:
+ with pytest.raises(AssertionError, match="equal spatial grid"):
+ _pack(share=True)
diff --git a/cosmos_framework/data/generator/augmentors/hr_lr_degradation/profiles.py b/cosmos_framework/data/generator/augmentors/hr_lr_degradation/profiles.py
new file mode 100644
index 000000000..d00d21d9c
--- /dev/null
+++ b/cosmos_framework/data/generator/augmentors/hr_lr_degradation/profiles.py
@@ -0,0 +1,392 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+"""Degradation profile definitions.
+
+A profile is a plain dataclass tree so it can be built from a LazyCall config, copied with
+``dataclasses.replace`` for ablation arms, and serialised into the degradation record.
+
+Numeric defaults follow the Cosmos Transfer1 corruptor configs, which in turn follow Real-ESRGAN
+``options/train_realesrnet_x2plus.yml``. The noise stage, which Transfer1 left out, uses
+the Real-ESRGAN x2plus values.
+
+Two families are registered in ``PROFILES``: the inherited ``p0_*`` / ``p1_*`` / ``p3_*`` profiles with those
+Real-ESRGAN ranges, and the ``img_*`` / ``vid_*`` regime ladder (clean / mild / moderate / harsh) budgeted to the
+x2 task, drawn per sample through ``IMAGE_SR_DEFAULT_MIX`` and ``VIDEO_SR_DEFAULT_MIX``.
+"""
+
+from __future__ import annotations
+
+import math
+from dataclasses import dataclass, field, fields, is_dataclass
+from typing import Any, Sequence
+
+RESIZE_MODES = ("area", "bilinear", "bicubic")
+CLEAN_RESIZE_KERNELS = ("bicubic_antialias", "bilinear_antialias", "area")
+
+# Real-ESRGAN kernel sizes are odd values from 7 to 21 and were tuned for roughly 400 px crops.
+_DEFAULT_KERNEL_RANGE = tuple(2 * v + 1 for v in range(3, 11))
+_DEFAULT_KERNEL_LIST = ("iso", "aniso", "generalized_iso", "generalized_aniso", "plateau_iso", "plateau_aniso")
+_DEFAULT_KERNEL_PROB = (0.45, 0.25, 0.12, 0.03, 0.12, 0.03)
+
+
+def _check_cutoff_range(cutoff: Sequence[float]) -> None:
+ """Sinc cutoffs are angular frequencies: 0 < low <= high <= pi (omega_c = 0 is an all-NaN kernel)."""
+ try:
+ lo, hi = (float(v) for v in cutoff)
+ except (TypeError, ValueError):
+ raise ValueError(f"sinc_cutoff_range must be a (low, high) pair, got {cutoff!r}") from None
+ if not 0.0 < lo <= hi <= math.pi:
+ raise ValueError(f"sinc_cutoff_range must satisfy 0 < low <= high <= pi, got {(lo, hi)}")
+
+
+@dataclass(frozen=True)
+class BlurConfig:
+ """Mixed-kernel or sinc blur applied once per clip."""
+
+ prob: float = 1.0
+ kernel_range: Sequence[int] = _DEFAULT_KERNEL_RANGE
+ sigma_range: Sequence[float] = (0.2, 3.0)
+ sinc_prob: float = 0.1
+ sinc_cutoff_range: Sequence[float] | None = None # omega_c bounds; None = Real-ESRGAN's size-dependent prior
+ kernel_list: Sequence[str] = _DEFAULT_KERNEL_LIST
+ kernel_prob: Sequence[float] = _DEFAULT_KERNEL_PROB
+ betag_range: Sequence[float] = (0.5, 4.0)
+ betap_range: Sequence[float] = (1.0, 2.0)
+
+ def __post_init__(self) -> None:
+ if self.sinc_cutoff_range is not None:
+ _check_cutoff_range(self.sinc_cutoff_range)
+
+
+@dataclass(frozen=True)
+class ResizeConfig:
+ """Random up / down / keep resize.
+
+ ``relative_to`` selects the reference size the factor multiplies: ``"current"`` (stage 1,
+ Real-ESRGAN first resize) or ``"target"`` (stage 2, Real-ESRGAN second resize is relative to
+ the final LR size).
+ """
+
+ prob: float = 1.0
+ updown_prob: Sequence[float] = (0.2, 0.7, 0.1)
+ scale_range: Sequence[float] = (0.15, 1.5)
+ modes: Sequence[str] = RESIZE_MODES
+ relative_to: str = "current"
+
+
+@dataclass(frozen=True)
+class NoiseConfig:
+ """Gaussian or Poisson noise, optionally grey (shared across channels)."""
+
+ prob: float = 1.0
+ gaussian_prob: float = 0.5
+ gaussian_sigma_range: Sequence[float] = (1.0, 30.0)
+ poisson_scale_range: Sequence[float] = (0.05, 3.0)
+ gray_noise_prob: float = 0.4
+
+
+@dataclass(frozen=True)
+class JPEGConfig:
+ prob: float = 1.0
+ quality_range: Sequence[float] = (30.0, 95.0)
+
+
+@dataclass(frozen=True)
+class FinalBlockConfig:
+ """Resize to the exact LR size, then sinc filter and JPEG in random order (Real-ESRGAN final block)."""
+
+ sinc_prob: float = 0.8
+ kernel_range: Sequence[int] = _DEFAULT_KERNEL_RANGE
+ sinc_cutoff_range: Sequence[float] | None = (math.pi / 3, math.pi) # None = Real-ESRGAN's size-dependent prior
+ modes: Sequence[str] = RESIZE_MODES
+ jpeg: JPEGConfig = field(default_factory=JPEGConfig)
+
+ def __post_init__(self) -> None:
+ if self.sinc_cutoff_range is not None:
+ _check_cutoff_range(self.sinc_cutoff_range)
+
+
+@dataclass(frozen=True)
+class CodecConfig:
+ """Whole-clip video codec round trip on the final LR (P3). Skipped for single images.
+
+ ``crf_range`` follows RealBasicVSR / Upscale-A-Video (18 to 35). ODVista-style streaming at fixed
+ low bitrates is harsher; widen the upper end for that benchmark. ``presets`` are x264 names; NVENC
+ encoders (``h264_nvenc`` / ``hevc_nvenc``) map them to ``p1``..``p7`` and CRF to ``cq`` (see ``codec.py``).
+ """
+
+ prob: float = 0.6
+ codecs: Sequence[str] = ("libx264", "libx265")
+ codec_prob: Sequence[float] = (0.7, 0.3)
+ crf_range: Sequence[float] = (18.0, 35.0)
+ presets: Sequence[str] = ("veryfast", "medium")
+
+
+@dataclass(frozen=True)
+class DegradationStage:
+ blur: BlurConfig | None = field(default_factory=BlurConfig)
+ resize: ResizeConfig | None = field(default_factory=ResizeConfig)
+ noise: NoiseConfig | None = field(default_factory=NoiseConfig)
+ jpeg: JPEGConfig | None = field(default_factory=JPEGConfig)
+
+
+@dataclass(frozen=True)
+class RealESRGANProfile:
+ """Real-ESRGAN style pipeline with one or two stages and a final block.
+
+ Attributes:
+ reference_longest_side: kernel sizes and sigmas are scaled by ``longest_side / reference``
+ so published ranges tuned near 400 to 720 px stay meaningful at 1080p and above.
+ min_resolution_factor / max_resolution_factor: clamps on that scale factor. ``[1.0, 1.5]`` keeps inputs
+ below the reference at the base ranges and stops proportional growth past 1.5x (unbounded scaling
+ produced 61 px kernels at 1080p). ``None`` leaves that side unbounded, as in the inherited profiles.
+ stage2_prob: probability of applying ``stage2`` when defined (1.0 = always, Real-ESRGAN style).
+ max_kernel_size: cap for the scaled kernel size (odd).
+ min_intermediate_scale: intermediate frames never shrink below this fraction of the
+ target LR size, which keeps a x2 task from becoming a x4 to x6 task.
+ """
+
+ name: str = "p1_first_order"
+ stage1: DegradationStage = field(default_factory=DegradationStage)
+ stage2: DegradationStage | None = None
+ final: FinalBlockConfig = field(default_factory=FinalBlockConfig)
+ codec: CodecConfig | None = None
+ modality: str = "any" # "image" / "video" / "any": AddLowRes rejects a profile built for the other stream
+ stage2_prob: float = 1.0 # probability of running stage2 when it is defined
+ scale_kernels_with_resolution: bool = True
+ reference_longest_side: int = 720
+ min_resolution_factor: float | None = None # lower clamp on longest_side / reference; None = unbounded
+ max_resolution_factor: float | None = None # upper clamp on longest_side / reference; None = unbounded
+ max_kernel_size: int = 61
+ min_intermediate_scale: float = 0.75
+
+ def __post_init__(self) -> None:
+ if not 0.0 <= self.stage2_prob <= 1.0:
+ raise ValueError(f"stage2_prob must be in [0, 1], got {self.stage2_prob}")
+ if self.stage2 is None and self.stage2_prob != 1.0:
+ raise ValueError(f"stage2_prob={self.stage2_prob} has no effect: the profile defines no stage2")
+ if not 0.0 <= self.min_intermediate_scale <= 1.0: # 0 = no floor (published Real-ESRGAN ranges)
+ raise ValueError(f"min_intermediate_scale must be in [0, 1], got {self.min_intermediate_scale}")
+ if self.reference_longest_side <= 0:
+ raise ValueError(f"reference_longest_side must be positive, got {self.reference_longest_side}")
+ lo, hi = self.min_resolution_factor, self.max_resolution_factor
+ for label, value in (("min_resolution_factor", lo), ("max_resolution_factor", hi)):
+ if value is not None and value <= 0:
+ raise ValueError(f"{label} must be positive, got {value}")
+ if lo is not None and hi is not None and lo > hi:
+ raise ValueError(f"min_resolution_factor {lo} exceeds max_resolution_factor {hi}")
+
+
+@dataclass(frozen=True)
+class CleanResizeProfile:
+ """P0: deterministic antialiased resize to the exact LR size, no other degradation."""
+
+ name: str = "p0_clean_bicubic"
+ kernel: str = "bicubic_antialias"
+ codec: CodecConfig | None = None # optional whole-clip re-encode on the clean LR (video only)
+ modality: str = "any" # "image" / "video" / "any": AddLowRes rejects a profile built for the other stream
+
+
+Profile = RealESRGANProfile | CleanResizeProfile
+
+
+def _second_order_stage2() -> DegradationStage:
+ return DegradationStage(
+ blur=BlurConfig(prob=0.8),
+ resize=ResizeConfig(updown_prob=(0.3, 0.4, 0.3), scale_range=(0.3, 1.2), relative_to="target"),
+ noise=NoiseConfig(gaussian_sigma_range=(1.0, 25.0), poisson_scale_range=(0.05, 2.5)),
+ jpeg=JPEGConfig(),
+ )
+
+
+PROFILES: dict[str, Profile] = {
+ "p0_clean_bicubic": CleanResizeProfile(name="p0_clean_bicubic", kernel="bicubic_antialias"),
+ "p0_clean_area": CleanResizeProfile(name="p0_clean_area", kernel="area"),
+ "p1_first_order": RealESRGANProfile(name="p1_first_order"),
+ "p1_first_order_no_noise": RealESRGANProfile(name="p1_first_order_no_noise", stage1=DegradationStage(noise=None)),
+ "p1_second_order": RealESRGANProfile(name="p1_second_order", stage2=_second_order_stage2()),
+ # P3: video terms. First-order pixel pipeline plus an H.264 / H.265 round trip on the LR clip.
+ "p3_video_codec": RealESRGANProfile(name="p3_video_codec", codec=CodecConfig()),
+ "p3_video_codec_second_order": RealESRGANProfile(
+ name="p3_video_codec_second_order", stage2=_second_order_stage2(), codec=CodecConfig()
+ ),
+ # Published Real-ESRGAN ranges without resolution scaling, for the calibration ablation (E2).
+ "p1_second_order_published": RealESRGANProfile(
+ name="p1_second_order_published",
+ stage2=_second_order_stage2(),
+ scale_kernels_with_resolution=False,
+ min_intermediate_scale=0.0,
+ ),
+}
+
+
+# ---------------------------------------------------------------------------------------------------------------
+# Regime profiles for x2 SR: a clean / mild / moderate / harsh ladder per modality, drawn per sample through the
+# mixes at the bottom. Ranges are budgeted to the x2 task (the LR already loses 4x the pixels), so extra blur stays
+# within about one LR pixel except in the harsh tail. Sigma values are HR pixels at 720p HR and scale with
+# resolution up to 1.5x; noise sigma is in 8-bit units. Images are JPEG-first: the JPEG sits in the final block, so
+# in Real-ESRGAN's random order it lands on the LR grid half the time and just before the final resize otherwise.
+# Video is codec-first (no JPEG under the codec except the rare "re-saved frames" term). Each regime declares its
+# modality so AddLowRes can reject a mix handed to the wrong stream.
+# ---------------------------------------------------------------------------------------------------------------
+_REGIME_KERNEL_PROB = (0.50, 0.30, 0.07, 0.03, 0.07, 0.03) # iso / aniso / gen-iso / gen-aniso / plateau-iso / -aniso
+# Regime ranges are specified at 720p HR (longest side 1280): r = clamp(longest / 1280, 1.0, 1.5), so 1080p HR
+# gets x1.5 and QHD / 4K stay at x1.5, while anything below 720p keeps the 720p ranges.
+_REGIME_COMMON = dict(
+ reference_longest_side=1280,
+ min_resolution_factor=1.0,
+ max_resolution_factor=1.5,
+ max_kernel_size=41,
+ min_intermediate_scale=0.75,
+)
+_IMG = dict(modality="image", **_REGIME_COMMON)
+_VID = dict(modality="video", **_REGIME_COMMON)
+_NO_FINAL_JPEG = JPEGConfig(prob=0.0)
+
+
+def _down(prob: float, low: float) -> ResizeConfig:
+ """Intermediate downscale to [low, 1.0] of the LR size (never the Real-ESRGAN 'up' branch). ``low`` stays at or
+ above ``min_intermediate_scale``: below it the floor would turn the tail of the range into a point mass."""
+ return ResizeConfig(prob=prob, updown_prob=(0.0, 1.0, 0.0), scale_range=(low, 1.0), relative_to="target")
+
+
+def _noise(
+ prob: float, gauss: tuple[float, float], poisson: tuple[float, float] | None, gray: float = 0.3
+) -> NoiseConfig:
+ if poisson is None:
+ return NoiseConfig(prob=prob, gaussian_prob=1.0, gaussian_sigma_range=gauss, gray_noise_prob=gray)
+ return NoiseConfig(
+ prob=prob, gaussian_prob=0.5, gaussian_sigma_range=gauss, poisson_scale_range=poisson, gray_noise_prob=gray
+ )
+
+
+def _blur(
+ prob: float, sigma: tuple[float, float], sinc_prob: float, cutoff: tuple[float, float] = (math.pi / 3, math.pi)
+) -> BlurConfig:
+ return BlurConfig(
+ prob=prob, sigma_range=sigma, sinc_prob=sinc_prob, sinc_cutoff_range=cutoff, kernel_prob=_REGIME_KERNEL_PROB
+ )
+
+
+def _final(
+ sinc_prob: float, cutoff: tuple[float, float] = (math.pi / 3, math.pi), jpeg: JPEGConfig = _NO_FINAL_JPEG
+) -> FinalBlockConfig:
+ """Final block: resize to LR plus sinc / JPEG in Real-ESRGAN's random order, so an image regime's JPEG lands on
+ the LR grid half the time (a photo saved at its own resolution) and just before the final resize otherwise.
+ Video regimes leave the JPEG off because the codec is the compression term."""
+ return FinalBlockConfig(sinc_prob=sinc_prob, sinc_cutoff_range=cutoff, jpeg=jpeg)
+
+
+# Video regimes keep the CodecConfig defaults for codecs (H.264 0.7 / H.265 0.3) and presets (veryfast / medium).
+REGIME_PROFILES: dict[str, Profile] = {
+ # ---- images: JPEG-first
+ "img_clean": CleanResizeProfile(name="img_clean", kernel="bicubic_antialias", modality="image"),
+ "img_mild": RealESRGANProfile(
+ name="img_mild",
+ stage1=DegradationStage(
+ blur=_blur(0.8, (0.2, 1.0), 0.05, cutoff=(math.pi / 2, math.pi)),
+ resize=_down(0.5, 0.85),
+ noise=_noise(0.6, (1.0, 6.0), (0.05, 0.8)),
+ jpeg=None,
+ ),
+ final=_final(0.2, (math.pi / 2, math.pi), jpeg=JPEGConfig(prob=0.7, quality_range=(70.0, 95.0))),
+ **_IMG,
+ ),
+ "img_moderate": RealESRGANProfile(
+ name="img_moderate",
+ stage1=DegradationStage(
+ blur=_blur(1.0, (0.5, 2.0), 0.1),
+ resize=_down(0.7, 0.75),
+ noise=_noise(0.8, (3.0, 12.0), (0.3, 1.5)),
+ jpeg=None,
+ ),
+ stage2=DegradationStage( # an earlier generation: re-saved at some intermediate size, then re-processed
+ blur=_blur(1.0, (0.25, 1.0), 0.05),
+ resize=None,
+ noise=_noise(0.8, (1.5, 6.0), (0.15, 0.75)),
+ jpeg=JPEGConfig(prob=0.9, quality_range=(60.0, 90.0)),
+ ),
+ stage2_prob=0.3,
+ final=_final(0.4, jpeg=JPEGConfig(prob=0.9, quality_range=(45.0, 80.0))),
+ **_IMG,
+ ),
+ "img_harsh": RealESRGANProfile(
+ name="img_harsh",
+ stage1=DegradationStage(
+ blur=_blur(1.0, (1.0, 3.0), 0.1),
+ resize=_down(0.9, 0.75),
+ noise=_noise(1.0, (5.0, 20.0), (1.0, 3.0)),
+ jpeg=None,
+ ),
+ final=_final(0.5, jpeg=JPEGConfig(prob=1.0, quality_range=(30.0, 60.0))),
+ **_IMG,
+ ),
+ # ---- video: codec-first
+ "vid_clean": CleanResizeProfile(
+ name="vid_clean",
+ kernel="bicubic_antialias",
+ modality="video",
+ codec=CodecConfig(
+ prob=0.3, codecs=("libx264",), codec_prob=(1.0,), crf_range=(16.0, 20.0), presets=("medium",)
+ ),
+ ),
+ "vid_mild": RealESRGANProfile(
+ name="vid_mild",
+ stage1=DegradationStage(
+ blur=_blur(0.8, (0.2, 1.0), 0.05),
+ resize=_down(0.4, 0.85),
+ noise=_noise(0.5, (1.0, 5.0), (0.05, 0.6)),
+ jpeg=None,
+ ),
+ final=_final(0.1),
+ codec=CodecConfig(prob=0.9, crf_range=(20.0, 28.0)),
+ **_VID,
+ ),
+ "vid_moderate": RealESRGANProfile(
+ name="vid_moderate",
+ stage1=DegradationStage(
+ blur=_blur(1.0, (0.5, 1.8), 0.1),
+ resize=_down(0.6, 0.75),
+ noise=_noise(0.7, (2.0, 10.0), (0.3, 1.2)),
+ jpeg=JPEGConfig(prob=0.2, quality_range=(60.0, 90.0)), # frames re-saved before re-encoding
+ ),
+ final=_final(0.2),
+ codec=CodecConfig(prob=1.0, crf_range=(26.0, 34.0)),
+ **_VID,
+ ),
+ "vid_harsh": RealESRGANProfile(
+ name="vid_harsh",
+ stage1=DegradationStage(
+ blur=_blur(1.0, (1.0, 2.5), 0.1),
+ resize=_down(0.9, 0.75),
+ noise=_noise(1.0, (5.0, 15.0), None),
+ jpeg=None,
+ ),
+ final=_final(0.3),
+ codec=CodecConfig(prob=1.0, crf_range=(32.0, 40.0), presets=("veryfast",)), # low-bitrate streaming look
+ **_VID,
+ ),
+}
+PROFILES.update(REGIME_PROFILES)
+
+# Regime mixtures for AddLowRes ``profiles=`` / dataset ``sr_profiles=``. The weights are the calibration knob
+# (against real 720p inventory and Cosmos 720p outputs); the per-regime ranges are fixed for physical plausibility.
+IMAGE_SR_DEFAULT_MIX: dict[str, float] = {"img_clean": 0.30, "img_mild": 0.45, "img_moderate": 0.20, "img_harsh": 0.05}
+VIDEO_SR_DEFAULT_MIX: dict[str, float] = {"vid_clean": 0.30, "vid_mild": 0.40, "vid_moderate": 0.25, "vid_harsh": 0.05}
+
+
+def get_profile(name_or_profile: str | Profile) -> Profile:
+ if isinstance(name_or_profile, (RealESRGANProfile, CleanResizeProfile)):
+ return name_or_profile
+ if name_or_profile not in PROFILES:
+ raise KeyError(f"Unknown degradation profile {name_or_profile!r}; known: {sorted(PROFILES)}")
+ return PROFILES[name_or_profile]
+
+
+def profile_to_dict(profile: Any) -> Any:
+ """Recursively convert a profile dataclass tree into JSON-serialisable primitives."""
+ if is_dataclass(profile) and not isinstance(profile, type):
+ return {f.name: profile_to_dict(getattr(profile, f.name)) for f in fields(profile)}
+ if isinstance(profile, (list, tuple)):
+ return [profile_to_dict(v) for v in profile]
+ return profile
diff --git a/cosmos_framework/data/generator/augmentors/interleaved_video_parsing.py b/cosmos_framework/data/generator/augmentors/interleaved_video_parsing.py
index 1a2769b38..f5ec0f23b 100644
--- a/cosmos_framework/data/generator/augmentors/interleaved_video_parsing.py
+++ b/cosmos_framework/data/generator/augmentors/interleaved_video_parsing.py
@@ -506,6 +506,12 @@ def __init__(self, input_keys: list, output_keys: list | None = None, args: dict
assert self.teacher_forcing_frames_per_chunk >= 1, (
f"teacher_forcing_frames_per_chunk must be >= 1, got {self.teacher_forcing_frames_per_chunk}"
)
+ self.target_num_frames: int | None = self.args.get("target_num_frames")
+ if self.target_num_frames is not None:
+ if not self.min_num_frames <= self.target_num_frames <= self.args.get("max_num_frames", 1000):
+ raise ValueError("target_num_frames must lie between min_num_frames and max_num_frames.")
+ if self.target_num_frames < 1 or (self.target_num_frames - 1) % (4 * self.teacher_forcing_frames_per_chunk):
+ raise ValueError("target_num_frames must contain complete VAE/teacher-forcing chunks after frame one.")
def _sample_frame_indices_for_chunk(
self,
@@ -513,7 +519,27 @@ def _sample_frame_indices_for_chunk(
chunk_start: int,
chunk_end: int,
min_stride_override: int | None = None,
- ) -> tuple[list[int], int]:
+ ) -> tuple[list[int], float]:
+ if self.target_num_frames is not None:
+ # A caption describes its complete window. Resample that whole span,
+ # rather than retaining its narrative after an unrelated prefix crop.
+ if chunk_start < 0 or chunk_end > decoder_len or chunk_end - chunk_start < self.target_num_frames:
+ return [], 0
+ if self.target_num_frames == 1:
+ # Only a one-frame caption window is faithful to a one-frame
+ # output; never collapse a multi-frame narrative into a still.
+ if chunk_end - chunk_start != 1:
+ return [], 0
+ effective_stride = 1.0
+ else:
+ effective_stride = (chunk_end - chunk_start - 1) / (self.target_num_frames - 1)
+ min_stride = int(min_stride_override) if min_stride_override is not None else self.min_stride
+ max_stride = max(self.max_stride, min_stride)
+ if not min_stride <= effective_stride <= max_stride:
+ return [], effective_stride
+ frame_indices = [chunk_start + round(frame * effective_stride) for frame in range(self.target_num_frames)]
+ return frame_indices, effective_stride
+
chunk_start = max(0, min(chunk_start, decoder_len))
chunk_end = max(chunk_start, min(chunk_end, decoder_len))
if chunk_end <= chunk_start:
@@ -595,6 +621,12 @@ def __call__(self, data_dict: dict) -> dict | None:
rank0_only=False,
)
return None
+ if self.target_num_frames is not None and not (
+ self.args.get("min_fps", 0.0)
+ <= meta_dict["framerate"] / stride
+ <= self.args.get("max_fps", float("inf"))
+ ):
+ return None
video_frames = self._decode_frames_at(video, frame_indices, rgb_transform) # [C,T,H,W]
except Exception as e:
diff --git a/cosmos_framework/data/generator/augmentors/text_transforms_for_video.py b/cosmos_framework/data/generator/augmentors/text_transforms_for_video.py
index dcc0ddb94..74b2baa71 100644
--- a/cosmos_framework/data/generator/augmentors/text_transforms_for_video.py
+++ b/cosmos_framework/data/generator/augmentors/text_transforms_for_video.py
@@ -420,6 +420,34 @@ def __init__(self, input_keys: dict, output_keys: Optional[list] = None, args: O
# The parser still needs metadata for fps/resolution after this transform.
self.keep_metas = self.args.get("keep_metas", True)
self.min_num_frames = int(self.args.get("min_num_frames", 5))
+ self.target_num_frames: int | None = self.args.get("target_num_frames")
+
+ def _supports_target_length(self, start_frame: int, end_frame: int, option: dict, meta_dict: dict) -> bool:
+ """Keep complete caption windows that can be resampled without duplicating frames."""
+ if self.target_num_frames is None:
+ return True
+ if self.target_num_frames < 1 or end_frame - start_frame < self.target_num_frames:
+ return False
+ try:
+ source_num_frames = int(meta_dict["nb_frames"])
+ source_fps = float(meta_dict["framerate"])
+ except (KeyError, TypeError, ValueError):
+ return False
+ if start_frame < 0 or end_frame > source_num_frames:
+ return False
+ if self.target_num_frames == 1:
+ # A single generated frame cannot represent an entire multi-frame
+ # caption window, even when VAE alignment maps it to one latent.
+ if end_frame - start_frame != 1:
+ return False
+ stride = 1.0
+ else:
+ stride = (end_frame - start_frame - 1) / (self.target_num_frames - 1)
+ min_stride = int(option.get("min_stride", self.args.get("min_stride", 1)))
+ max_stride = max(int(self.args.get("max_stride", 3)), min_stride)
+ min_fps = float(self.args.get("min_fps", 0.0))
+ max_fps = float(self.args.get("max_fps", float("inf")))
+ return min_stride <= stride <= max_stride and min_fps <= source_fps / stride <= max_fps
def __call__(self, data_dict: dict) -> dict | None:
meta_dict = data_dict.get(self.meta_key)
@@ -472,7 +500,9 @@ def __call__(self, data_dict: dict) -> dict | None:
end_frame = int(chunk["end_frame"])
except (KeyError, TypeError, ValueError):
continue
- if end_frame - start_frame >= self.min_num_frames:
+ if end_frame - start_frame >= self.min_num_frames and self._supports_target_length(
+ start_frame, end_frame, sampled_caption_option, meta_dict
+ ):
eligible_chunk_keys.append(chunk_key)
if not eligible_chunk_keys:
@@ -488,6 +518,14 @@ def __call__(self, data_dict: dict) -> dict | None:
chunk_start_frame = int(sampled_chunk["start_frame"])
chunk_end_frame = int(sampled_chunk["end_frame"])
structured = json.loads(sampled_chunk["caption"])
+ if self.target_num_frames is not None and not isinstance(structured, dict):
+ raise ValueError("Exact target lengths require a structured caption dictionary.")
+ if self.target_num_frames is not None:
+ # Remove source metadata before the downstream augmentor appends
+ # the sampled clip's duration/FPS. Keep the teacher's established
+ # serialized-caption format and preserve all narrative fields.
+ structured.pop("duration", None)
+ structured.pop("fps", None)
except Exception as e:
log.warning(
f"TextTransformForVideoTransferChunkedFrames: failed to decode {sampled_caption_key}.{self.CAPTION_FIELD}. "
diff --git a/cosmos_framework/data/generator/augmentors/transfer_target_length_test.py b/cosmos_framework/data/generator/augmentors/transfer_target_length_test.py
new file mode 100644
index 000000000..9592e5d8a
--- /dev/null
+++ b/cosmos_framework/data/generator/augmentors/transfer_target_length_test.py
@@ -0,0 +1,282 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+
+
+import json
+
+import pytest
+import torch
+
+import cosmos_framework.data.generator.augmentors.text_tokenizer as text_tokenizer
+from cosmos_framework.data.generator.augmentors.duration_fps_text_timestamps import DurationFPSTextTimeStamps
+from cosmos_framework.data.generator.augmentors.interleaved_video_parsing import (
+ VideoTransferAlignedSelectedControlParsing,
+)
+from cosmos_framework.data.generator.augmentors.text_transforms_for_video import (
+ TextTransformForVideoTransferChunkedFrames,
+)
+
+
+def _make_target_parser(target_num_frames: int = 81) -> VideoTransferAlignedSelectedControlParsing:
+ return VideoTransferAlignedSelectedControlParsing(
+ input_keys=["metas", "video"],
+ args={
+ "max_num_frames": 201,
+ "min_num_frames": 1,
+ "target_num_frames": target_num_frames,
+ "teacher_forcing_frames_per_chunk": 1,
+ "min_stride": 1,
+ "max_stride": 3,
+ "min_fps": 5.0,
+ "max_fps": 60.0,
+ "seek_mode": "exact",
+ },
+ )
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+@pytest.mark.parametrize("target_num_frames,source_num_frames", [(81, 200), (137, 400), (201, 400)])
+def test_target_parser_retains_entire_caption_span(target_num_frames: int, source_num_frames: int) -> None:
+ parser = _make_target_parser(target_num_frames)
+ indices, stride = parser._sample_frame_indices_for_chunk(500, 20, 20 + source_num_frames)
+
+ assert len(indices) == len(set(indices)) == target_num_frames
+ assert indices[0] == 20
+ assert indices[-1] == 20 + source_num_frames - 1
+ assert indices == sorted(indices)
+ assert stride == pytest.approx((source_num_frames - 1) / (target_num_frames - 1))
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+@pytest.mark.parametrize(
+ "start,end,min_stride",
+ [(-1, 199, None), (400, 600, None), (0, 80, None), (0, 300, None), (0, 200, 3)],
+)
+def test_target_parser_rejects_partial_caption_windows_or_invalid_stride(
+ start: int, end: int, min_stride: int | None
+) -> None:
+ indices, _ = _make_target_parser()._sample_frame_indices_for_chunk(500, start, end, min_stride)
+ assert indices == []
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+@pytest.mark.parametrize("target_num_frames", [0, 8, 82, 205])
+def test_target_parser_rejects_invalid_target_lengths(target_num_frames: int) -> None:
+ with pytest.raises(ValueError, match="target_num_frames"):
+ _make_target_parser(target_num_frames)
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+@pytest.mark.parametrize("latent_frames", range(1, 52))
+def test_every_latent_bucket_accepts_its_complete_caption_window(latent_frames: int) -> None:
+ target_num_frames = 1 + 4 * (latent_frames - 1)
+ parser = _make_target_parser(target_num_frames)
+ indices, stride = parser._sample_frame_indices_for_chunk(240, 10, 10 + target_num_frames)
+ assert indices == list(range(10, 10 + target_num_frames))
+ assert stride == 1.0
+
+ caption_transform = TextTransformForVideoTransferChunkedFrames(
+ input_keys=["metas"],
+ args={"caption_config": {"captions": 1.0}, "target_num_frames": target_num_frames, "min_num_frames": 1},
+ )
+ assert caption_transform._supports_target_length(
+ 10, 10 + target_num_frames, {}, {"nb_frames": 240, "framerate": 30.0}
+ )
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+@pytest.mark.parametrize("source_num_frames", [2, 3, 4, 81, 201])
+def test_one_frame_bucket_rejects_multi_frame_caption_windows(source_num_frames: int) -> None:
+ parser = _make_target_parser(1)
+ assert parser._sample_frame_indices_for_chunk(240, 10, 10 + source_num_frames)[0] == []
+ caption_transform = TextTransformForVideoTransferChunkedFrames(
+ input_keys=["metas"],
+ args={"caption_config": {"captions": 1.0}, "target_num_frames": 1, "min_num_frames": 1},
+ )
+ assert not caption_transform._supports_target_length(
+ 10, 10 + source_num_frames, {}, {"nb_frames": 240, "framerate": 30.0}
+ )
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+@pytest.mark.parametrize("target_num_frames", [1, 5])
+def test_short_bucket_selects_only_its_complete_caption(target_num_frames: int) -> None:
+ narrative = {"description": "A red door.", "duration": "8s", "fps": 24.0}
+ transform = TextTransformForVideoTransferChunkedFrames(
+ input_keys=["metas"],
+ args={"caption_config": {"captions": 1.0}, "target_num_frames": target_num_frames, "min_num_frames": 1},
+ )
+ sample = {
+ "metas": {
+ "framerate": 30.0,
+ "nb_frames": 201,
+ "captions": {
+ "caption_structured": json.dumps(
+ {
+ "long": {"start_frame": 0, "end_frame": 201, "caption": json.dumps({"description": "Wrong"})},
+ "short": {
+ "start_frame": 7,
+ "end_frame": 7 + target_num_frames,
+ "caption": json.dumps(narrative),
+ },
+ }
+ )
+ },
+ },
+ }
+ transformed = transform(sample)
+ assert transformed is not None
+ assert transformed["sampled_chunk_key"] == "short"
+ assert (transformed["chunk_start_frame"], transformed["chunk_end_frame"]) == (7, 7 + target_num_frames)
+ assert json.loads(transformed["ai_caption"]) == {"description": narrative["description"]}
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+@pytest.mark.parametrize("control_num_frames", [299, 300])
+def test_target_length_selects_matching_caption_and_rebuilds_metadata(
+ monkeypatch: pytest.MonkeyPatch, control_num_frames: int
+) -> None:
+ narrative = {
+ "description": "The door opens, a person enters, and the door closes.",
+ "duration": "8s",
+ "fps": 24.0,
+ }
+ caption_transform = TextTransformForVideoTransferChunkedFrames(
+ input_keys=["metas"],
+ args={
+ "caption_config": {"captions": 1.0},
+ "target_num_frames": 81,
+ "min_stride": 1,
+ "max_stride": 3,
+ "min_fps": 5.0,
+ "max_fps": 60.0,
+ },
+ )
+ sample = {
+ "metas": {
+ "framerate": 30.0,
+ "nb_frames": 500,
+ "width": 12,
+ "height": 8,
+ "captions": {
+ "caption_structured": json.dumps(
+ {
+ "short": {"start_frame": 0, "end_frame": 9, "caption": json.dumps({"description": "Wrong"})},
+ "complete": {"start_frame": 100, "end_frame": 300, "caption": json.dumps(narrative)},
+ }
+ )
+ },
+ },
+ "video": b"rgb",
+ "persisted_control": b"control",
+ "_persisted_control_meta": {"framerate": 30.0, "nb_frames": 500, "width": 12, "height": 8},
+ "_selected_control_modality": "seg",
+ "__url__": "url",
+ "__key__": "key",
+ }
+ transformed = caption_transform(sample)
+ assert transformed is not None
+ assert transformed["sampled_chunk_key"] == "complete"
+ assert json.loads(transformed["ai_caption"]) == {"description": narrative["description"]}
+
+ parser = _make_target_parser()
+ monkeypatch.setattr(parser, "_validate_and_probe", lambda *_args: True)
+ monkeypatch.setattr(parser, "_probe_video_len", lambda payload: 500 if payload == b"rgb" else control_num_frames)
+ calls: list[list[int]] = []
+
+ def fake_decode(
+ _payload: bytes,
+ indices: list[int],
+ _transforms: object = None,
+ _output_dtype: torch.dtype = torch.uint8,
+ ) -> torch.Tensor:
+ calls.append(list(indices))
+ return torch.zeros(3, len(indices), 8, 12, dtype=torch.uint8) # [C,T,H,W]
+
+ monkeypatch.setattr(parser, "_decode_frames_at", fake_decode)
+ parsed = parser(transformed)
+ if control_num_frames < 300:
+ assert parsed is None
+ assert calls == []
+ return
+
+ assert parsed is not None
+ assert len(calls) == 2
+ assert calls[0] == calls[1]
+ assert (calls[0][0], calls[0][-1]) == (100, 299)
+ assert parsed["video"]["num_frames"] == 81
+ effective_fps = 30.0 * 80 / 199
+ assert parsed["video"]["conditioning_fps"] == pytest.approx(effective_fps)
+ metadata_sample = {
+ "ai_caption": parsed["ai_caption"],
+ "video": parsed["video"]["video"], # [C,T,H,W]
+ "conditioning_fps": parsed["video"]["conditioning_fps"],
+ }
+ metadata_transform = DurationFPSTextTimeStamps(args={"fractional_duration": True, "skip_on_error": False})
+ assert metadata_transform(metadata_sample) is metadata_sample
+ expected_caption = json.dumps({"description": narrative["description"]}) + (
+ f". The video is {81 / effective_fps:.1f} seconds long and is of {effective_fps:.0f} FPS."
+ )
+ assert metadata_sample["ai_caption"] == expected_caption
+ tokenized_captions: list[str] = []
+
+ class CaptionRecorder:
+ def tokenize_text(self, caption: str, system_prompt: str) -> list[int]:
+ tokenized_captions.append(caption)
+ assert "control signals" in system_prompt
+ return [11, 12, 13]
+
+ monkeypatch.setattr(text_tokenizer, "lazy_instantiate", lambda _config: CaptionRecorder())
+ tokenizer = text_tokenizer.TextTokenizerTransformForTransfer(
+ input_keys=["ai_caption"], output_keys=["text_token_ids"], args={"tokenizer_config": {}}
+ )
+ assert tokenizer(metadata_sample) is metadata_sample
+ assert tokenized_captions == [expected_caption]
+ assert metadata_sample["text_token_ids"].tolist() == [11, 12, 13]
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+@pytest.mark.parametrize("fps,min_stride", [(5.0, 1), (30.0, 3)])
+def test_target_caption_rejects_unsupported_effective_fps_and_source_stride(fps: float, min_stride: int) -> None:
+ transform = TextTransformForVideoTransferChunkedFrames(
+ input_keys=["metas"],
+ args={"caption_config": {"captions": 1.0}, "target_num_frames": 81, "min_fps": 5.0, "max_stride": 3},
+ )
+ assert not transform._supports_target_length(
+ 0, 200, {"min_stride": min_stride}, {"nb_frames": 200, "framerate": fps}
+ )
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+@pytest.mark.parametrize("target_num_frames", [None, 81])
+def test_source_duration_fields_are_removed_only_for_exact_target_captions(target_num_frames: int | None) -> None:
+ original_caption = {"description": "The person leaves.", "duration": "8s", "fps": 24.0}
+ transform = TextTransformForVideoTransferChunkedFrames(
+ input_keys=["metas"], args={"caption_config": {"captions": 1.0}, "target_num_frames": target_num_frames}
+ )
+ sample = {
+ "metas": {
+ "framerate": 30.0,
+ "nb_frames": 200,
+ "captions": {
+ "caption_structured": json.dumps(
+ {"complete": {"start_frame": 0, "end_frame": 200, "caption": json.dumps(original_caption)}}
+ )
+ },
+ }
+ }
+ transformed = transform(sample)
+ assert transformed is not None
+ expected_caption = (
+ original_caption if target_num_frames is None else {"description": original_caption["description"]}
+ )
+ assert transformed["ai_caption"] == json.dumps(expected_caption)
diff --git a/cosmos_framework/data/generator/processors/__init__.py b/cosmos_framework/data/generator/processors/__init__.py
index 0b2ac90ff..289ffbdb7 100644
--- a/cosmos_framework/data/generator/processors/__init__.py
+++ b/cosmos_framework/data/generator/processors/__init__.py
@@ -127,10 +127,11 @@ def build_processor(
# (e.g. the top level of nvidia/Cosmos3-Nano, which ships its own
# preprocessor_config.json, tokenizer.json, etc). Avoids the redundant
# upstream Qwen/Qwen3-VL-*-Instruct fetch. Qwen3-VL Nemo Chat variants use
- # their specialized loss-mask wrapper. Renewed Cosmos3-Edge snapshots use
- # the Nemotron bridge; other local artifacts use Qwen3VLProcessor.
+ # their specialized loss-mask wrapper. Nemotron3-Dense-VL and renewed
+ # Cosmos3-Edge snapshots use the Nemotron bridge; other local artifacts use
+ # Qwen3VLProcessor.
if os.path.isdir(tokenizer_type):
- if is_cosmos3_edge_native_snapshot(tokenizer_type):
+ if is_cosmos3_edge_native_snapshot(tokenizer_type) or "NVIDIA-Nemotron-3-Dense-VL" in tokenizer_type:
return Nemotron3DenseVLProcessor(tokenizer_type, cache_dir=cache_dir)
if "Qwen/Qwen3-VL" in tokenizer_type and "Nemo-Chat" in tokenizer_type:
return Qwen3VLNemoChatProcessor(tokenizer_type, cache_dir=cache_dir)
diff --git a/cosmos_framework/data/generator/processors/cosmos3_edge_processing_test.py b/cosmos_framework/data/generator/processors/cosmos3_edge_processing_test.py
index 80cbcca7f..9031c0e42 100644
--- a/cosmos_framework/data/generator/processors/cosmos3_edge_processing_test.py
+++ b/cosmos_framework/data/generator/processors/cosmos3_edge_processing_test.py
@@ -15,6 +15,8 @@
import hashlib
import json
import os
+from pathlib import Path
+from unittest.mock import Mock
import numpy as np
import pytest
@@ -98,6 +100,24 @@ def test_detection_rule_rejects_non_edge_dirs(tmp_path) -> None:
assert not is_cosmos3_edge_native_snapshot(str(tmp_path))
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_build_processor_routes_local_nemotron3_dense_dir(
+ tmp_path: Path,
+ monkeypatch: pytest.MonkeyPatch,
+) -> None:
+ import cosmos_framework.data.generator.processors as processors
+
+ model_dir = tmp_path / "nvidia" / "NVIDIA-Nemotron-3-Dense-VL-2B-BF16-Alignment"
+ model_dir.mkdir(parents=True)
+ expected_processor = object()
+ processor_factory = Mock(return_value=expected_processor)
+ monkeypatch.setattr(processors, "Nemotron3DenseVLProcessor", processor_factory)
+
+ assert processors.build_processor(str(model_dir)) is expected_processor
+ processor_factory.assert_called_once_with(str(model_dir), cache_dir=None)
+
+
@requires_snapshot
def test_text_only_matches_golden(processor) -> None:
inputs = _apply(
diff --git a/cosmos_framework/data/generator/sequence_packing/autoregressive.py b/cosmos_framework/data/generator/sequence_packing/autoregressive.py
index a76928be6..c5cd367d2 100644
--- a/cosmos_framework/data/generator/sequence_packing/autoregressive.py
+++ b/cosmos_framework/data/generator/sequence_packing/autoregressive.py
@@ -3,10 +3,12 @@
"""Autoregressive sequence packing for framewise and chunkwise AR generation."""
-from typing import cast
+from collections.abc import Sequence
+from typing import Any, cast
import torch
+from cosmos_framework.data.generator.augmentors.text_tokenizer import TEXT_SYSTEM_PROMPT_KEY
from cosmos_framework.model.generator.utils.data_and_condition import GenerationDataClean
from cosmos_framework.data.generator.sequence_packing import (
PackedSequence,
@@ -15,6 +17,31 @@
)
+def resolve_text_system_prompt(data_batch: dict[str, Any]) -> Any:
+ """Read tokenizer metadata, falling back to the legacy key only when absent."""
+ return data_batch.get(TEXT_SYSTEM_PROMPT_KEY, data_batch.get("system_prompt"))
+
+
+def caption_system_prompts(sequence_plans: list[SequencePlan], data_batch: dict[str, Any]) -> list[str | None]:
+ """Expand exact per-sample tokenizer prompts over the packed caption slots."""
+ value = resolve_text_system_prompt(data_batch)
+ if value is None or isinstance(value, str):
+ prompts = [value] * len(sequence_plans)
+ elif isinstance(value, (list, tuple)):
+ if len(value) != len(sequence_plans):
+ raise ValueError("Tokenizer system prompts must have one entry per sample.")
+ prompts = list(value)
+ else:
+ raise TypeError("Tokenizer system prompts must be strings or a per-sample sequence.")
+ if any(prompt is not None and not isinstance(prompt, str) for prompt in prompts):
+ raise TypeError("Each tokenizer system prompt must be a string or None.")
+ return [
+ prompt
+ for plan, prompt in zip(sequence_plans, prompts)
+ for _ in range(len(plan.text_view_ids) if plan.text_view_ids is not None else 1)
+ ]
+
+
def pack_input_sequence_autoregressive(
vision_latent: torch.Tensor | None,
action_latent: torch.Tensor | None,
@@ -341,7 +368,7 @@ def pack_input_sequence_autoregressive_batch(
*,
latent_patch_size: int = 1,
condition_frame_indexes_vision: list[int] | None = None,
- frame_idx: int = 0,
+ frame_idx: int | Sequence[int] = 0,
temporal_compression_factor: int = 4,
video_temporal_causal: bool = True,
enable_fps_modulation: bool = True,
@@ -420,6 +447,8 @@ def pack_input_sequence_autoregressive_batch(
input_text_indexes = text_tokens if text_tokens is not None else [[] for _ in range(batch_size)]
input_timesteps = torch.full((batch_size,), timestep, dtype=torch.float32) # [B]
+ if isinstance(frame_idx, Sequence) and len(frame_idx) != batch_size:
+ raise ValueError(f"Expected {batch_size} per-sample frame indices, got {len(frame_idx)}")
initial_offsets: list[int | float] = []
for sample_idx, fps in enumerate(fps_vision):
frame_stride = base_fps / fps
@@ -428,7 +457,8 @@ def pack_input_sequence_autoregressive_batch(
if cached_text_offsets is None
else cached_text_offsets[sample_idx] + unified_3d_mrope_temporal_modality_margin
)
- initial_offsets.append(text_offset + frame_idx * frame_stride)
+ row_frame_idx = frame_idx[sample_idx] if isinstance(frame_idx, Sequence) else frame_idx
+ initial_offsets.append(text_offset + row_frame_idx * frame_stride)
return pack_input_sequence(
sequence_plans=sequence_plans,
diff --git a/cosmos_framework/data/generator/sequence_packing/modality.py b/cosmos_framework/data/generator/sequence_packing/modality.py
index 2143fa3d3..9e9b06fe8 100644
--- a/cosmos_framework/data/generator/sequence_packing/modality.py
+++ b/cosmos_framework/data/generator/sequence_packing/modality.py
@@ -7,6 +7,8 @@
import torch
+from cosmos_framework.data.generator.sequence_packing.runtime import to_device_nonblocking
+
def _empty_long_tensor() -> torch.Tensor:
return torch.empty(0, dtype=torch.long) # [0]
@@ -121,19 +123,23 @@ def __post_init__(self) -> None:
assert isinstance(self.mse_loss_indexes, torch.Tensor), "ModalityData.mse_loss_indexes must be finalized"
def to_cuda(self) -> None:
- """Move all tensor fields to CUDA in-place."""
- self.sequence_indexes = self.sequence_indexes.cuda()
- self.timesteps = self.timesteps.cuda()
- self.mse_loss_indexes = self.mse_loss_indexes.cuda()
- self.tokens = [token.cuda() for token in self.tokens]
- self.condition_mask = [cm.cuda() for cm in self.condition_mask]
- self.noisy_frame_indexes = [ni.cuda() for ni in self.noisy_frame_indexes]
- self.domain_id = [d.cuda() for d in self.domain_id]
+ """Move all tensor fields to CUDA in-place (asynchronous pinned copies, identical values)."""
+ self.sequence_indexes = to_device_nonblocking(self.sequence_indexes, "cuda")
+ self.timesteps = to_device_nonblocking(self.timesteps, "cuda")
+ self.mse_loss_indexes = to_device_nonblocking(self.mse_loss_indexes, "cuda")
+ self.tokens = [to_device_nonblocking(token, "cuda") for token in self.tokens]
+ self.condition_mask = [to_device_nonblocking(cm, "cuda") for cm in self.condition_mask]
+ self.noisy_frame_indexes = [to_device_nonblocking(ni, "cuda") for ni in self.noisy_frame_indexes]
+ self.domain_id = [to_device_nonblocking(d, "cuda") for d in self.domain_id]
# raw_action_dim is optional (e.g., when action-channel masking is disabled).
if self.raw_action_dim is not None:
- self.raw_action_dim = [d.cuda() if d is not None else None for d in self.raw_action_dim]
+ self.raw_action_dim = [
+ to_device_nonblocking(d, "cuda") if d is not None else None for d in self.raw_action_dim
+ ]
if self.action_valid_mask is not None:
- self.action_valid_mask = [m.cuda() if m is not None else None for m in self.action_valid_mask]
+ self.action_valid_mask = [
+ to_device_nonblocking(m, "cuda") if m is not None else None for m in self.action_valid_mask
+ ]
def prepare_attention_mask_per_sample(split_lens, attn_modes, device="cpu"):
diff --git a/cosmos_framework/data/generator/sequence_packing/runtime.py b/cosmos_framework/data/generator/sequence_packing/runtime.py
index a0fd1fe17..6ff892744 100644
--- a/cosmos_framework/data/generator/sequence_packing/runtime.py
+++ b/cosmos_framework/data/generator/sequence_packing/runtime.py
@@ -138,6 +138,23 @@ def _find_non_causal_text_token_idx(
return out
+def to_device_nonblocking(tensor: torch.Tensor, device: torch.device | str) -> torch.Tensor:
+ """Move ``tensor`` to ``device`` without stalling the host.
+
+ A pageable host->device copy (``tensor.cuda()`` / ``torch.tensor(list, device="cuda")``)
+ blocks the CPU until every kernel already queued on the device has finished, which keeps
+ the CPU from running ahead of the GPU in per-frame inference loops. Staging the bytes
+ through pinned memory makes the same copy asynchronous; the values are identical and the
+ caching host allocator keeps the staging block alive until the copy has completed.
+ """
+ target = torch.device(device)
+ if tensor.device == target or tensor.numel() == 0:
+ return tensor.to(target)
+ if target.type == "cuda" and tensor.device.type == "cpu" and not tensor.is_pinned():
+ tensor = tensor.pin_memory()
+ return tensor.to(target, non_blocking=True)
+
+
def _compute_mode_indices_and_offsets(
split_lens: torch.Tensor | List[int], attn_modes: List[str], mode: str, device: torch.device
) -> tuple[torch.Tensor, torch.Tensor]:
@@ -160,8 +177,8 @@ def _compute_mode_indices_and_offsets(
start += split_len
return (
- torch.tensor(indices, dtype=torch.int32, device=device),
- torch.tensor(offsets, dtype=torch.int32, device=device),
+ to_device_nonblocking(torch.tensor(indices, dtype=torch.int32), device),
+ to_device_nonblocking(torch.tensor(offsets, dtype=torch.int32), device),
) # [N_mode_tokens], [N_mode_splits+1]
@@ -318,9 +335,9 @@ def _build_sequence_pack_metadata(
_max_causal_len = max((split_lens[i] for i in range(len(split_lens)) if attn_modes[i] == "causal"), default=0)
_max_full_len = max((split_lens[i] for i in range(len(split_lens)) if attn_modes[i] == "full"), default=0)
- sample_lens_cu = torch.tensor([0] + sample_lens, device=device, dtype=torch.int32) # [N_samples+1]
+ sample_lens_cu = to_device_nonblocking(torch.tensor([0] + sample_lens, dtype=torch.int32), device) # [N_samples+1]
_sample_offsets = torch.cumsum(sample_lens_cu, dim=0, dtype=torch.int32) # [N_samples+1]
- sample_lens_tensor = torch.tensor(sample_lens, device=device, dtype=torch.int64) # [N_samples]
+ sample_lens_tensor = to_device_nonblocking(torch.tensor(sample_lens, dtype=torch.int64), device) # [N_samples]
sample_ids = torch.repeat_interleave(
torch.arange(len(sample_lens), device=device, dtype=torch.int64),
sample_lens_tensor,
diff --git a/cosmos_framework/data/generator/sequence_packing/runtime_test.py b/cosmos_framework/data/generator/sequence_packing/runtime_test.py
new file mode 100644
index 000000000..60b3d4ec8
--- /dev/null
+++ b/cosmos_framework/data/generator/sequence_packing/runtime_test.py
@@ -0,0 +1,51 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+
+import pytest
+import torch
+
+from cosmos_framework.data.generator.sequence_packing.runtime import (
+ prepare_sequence_pack_metadata,
+ to_device_nonblocking,
+)
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_to_device_nonblocking_is_identity_on_the_same_device() -> None:
+ tensor = torch.arange(6, dtype=torch.int32)
+ assert to_device_nonblocking(tensor, "cpu") is tensor
+ assert to_device_nonblocking(tensor, torch.device("cpu")) is tensor
+ empty = torch.empty(0, dtype=torch.int64)
+ assert to_device_nonblocking(empty, "cpu").numel() == 0
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_metadata_from_host_indexes_matches_metadata_from_device_indexes() -> None:
+ """``PackedSequence.to_cuda`` validates the layout with the host copy of ``text_indexes``;
+ the result must not depend on which copy is handed in."""
+ sample_lens = [7, 5]
+ split_lens = [3, 4, 2, 3]
+ attn_modes = ["causal", "full", "causal", "full"]
+ text_indexes = torch.tensor([0, 1, 2, 7, 8], dtype=torch.int64)
+ reference = prepare_sequence_pack_metadata(
+ sample_lens=sample_lens,
+ split_lens=split_lens,
+ attn_modes=attn_modes,
+ packed_und_token_indexes=text_indexes,
+ device=torch.device("cpu"),
+ )
+ other = prepare_sequence_pack_metadata(
+ sample_lens=sample_lens,
+ split_lens=split_lens,
+ attn_modes=attn_modes,
+ packed_und_token_indexes=text_indexes.clone(),
+ device=torch.device("cpu"),
+ )
+ for name in vars(reference):
+ a, b = getattr(reference, name), getattr(other, name)
+ if isinstance(a, torch.Tensor):
+ assert torch.equal(a, b), name
+ else:
+ assert a == b, name
diff --git a/cosmos_framework/data/generator/sequence_packing/sequence.py b/cosmos_framework/data/generator/sequence_packing/sequence.py
index 1dedc46fb..be297c14d 100644
--- a/cosmos_framework/data/generator/sequence_packing/sequence.py
+++ b/cosmos_framework/data/generator/sequence_packing/sequence.py
@@ -19,6 +19,7 @@
from cosmos_framework.data.generator.sequence_packing.runtime import (
SequencePackMetadata,
prepare_sequence_pack_metadata,
+ to_device_nonblocking,
)
if TYPE_CHECKING:
@@ -1256,16 +1257,24 @@ def __post_init__(self) -> None:
)
def to_cuda(self) -> None:
- """Move all tensor fields to CUDA in-place."""
- self.text_ids = self.text_ids.cuda()
- self.text_indexes = self.text_indexes.cuda()
- self.position_ids = self.position_ids.cuda()
+ """Move all tensor fields to CUDA in-place.
+
+ Copies are asynchronous (pinned staging, see ``to_device_nonblocking``) so a per-frame
+ inference loop never waits for the GPU here. The layout metadata depends only on the
+ host-side lengths, so it is built from the host ``text_indexes`` before they move and
+ is kept when the pack already lives on CUDA (AR callers swap device-resident noise and
+ timesteps into the pack and call ``to_cuda`` again every denoising step).
+ """
+ host_text_indexes = self.text_indexes if self.text_indexes.device.type == "cpu" else None
+ self.text_ids = to_device_nonblocking(self.text_ids, "cuda")
+ self.text_indexes = to_device_nonblocking(self.text_indexes, "cuda")
+ self.position_ids = to_device_nonblocking(self.position_ids, "cuda")
if isinstance(self.label_ids, torch.Tensor):
- self.label_ids = self.label_ids.cuda()
+ self.label_ids = to_device_nonblocking(self.label_ids, "cuda")
if isinstance(self.ce_loss_indexes, torch.Tensor):
- self.ce_loss_indexes = self.ce_loss_indexes.cuda()
+ self.ce_loss_indexes = to_device_nonblocking(self.ce_loss_indexes, "cuda")
if isinstance(self.ce_loss_weights, torch.Tensor):
- self.ce_loss_weights = self.ce_loss_weights.cuda()
+ self.ce_loss_weights = to_device_nonblocking(self.ce_loss_weights, "cuda")
if self.vision is not None:
self.vision.to_cuda()
if self.lidar is not None:
@@ -1274,15 +1283,22 @@ def to_cuda(self) -> None:
self.action.to_cuda()
if self.sound is not None:
self.sound.to_cuda()
- self.prepare_sequence_pack_metadata()
+ if host_text_indexes is None and self._sequence_pack_metadata is not None:
+ return # already prepared on CUDA; the layout has not changed
+ self.prepare_sequence_pack_metadata(host_text_indexes=host_text_indexes)
- def prepare_sequence_pack_metadata(self) -> None:
- """Validate and prepare device-specific metadata for this layout."""
+ def prepare_sequence_pack_metadata(self, host_text_indexes: torch.Tensor | None = None) -> None:
+ """Validate and prepare device-specific metadata for this layout.
+
+ ``host_text_indexes`` (a CPU copy of ``text_indexes``) avoids a device->host read of the
+ indexes when the pack has just been moved to CUDA.
+ """
+ indexes = host_text_indexes if host_text_indexes is not None else self.text_indexes
self._sequence_pack_metadata = prepare_sequence_pack_metadata(
sample_lens=self.sample_lens,
split_lens=self.split_lens,
attn_modes=self.attn_modes,
- packed_und_token_indexes=self.text_indexes,
+ packed_und_token_indexes=indexes,
device=self.text_indexes.device,
text_caption_lens=self.text_caption_lens,
)
diff --git a/cosmos_framework/data/generator/sequence_packing/temporal_causal.py b/cosmos_framework/data/generator/sequence_packing/temporal_causal.py
index 2c958204b..f53605a50 100644
--- a/cosmos_framework/data/generator/sequence_packing/temporal_causal.py
+++ b/cosmos_framework/data/generator/sequence_packing/temporal_causal.py
@@ -8,6 +8,7 @@
import torch
from cosmos_framework.data.generator.sequence_packing.mrope import get_3d_mrope_ids_vae_tokens
+from cosmos_framework.data.generator.sequence_packing.runtime import to_device_nonblocking
from cosmos_framework.data.generator.sequence_packing.sequence import PackedSequenceBuilder
@@ -142,15 +143,16 @@ def pack_supertokens_temporal_causal(
# Vision conditioning mask: (T, 1, 1)
condition_set_vision = {idx for idx in condition_frame_indexes_vision if 0 <= idx < latent_t}
- vision_condition_mask = torch.zeros((latent_t, 1, 1), device=device, dtype=dtype) # [T,1,1]
+ # Built on the host and moved asynchronously: writing Python scalars into a CUDA tensor
+ # element by element synchronises the host with the device on every call.
+ vision_condition_mask = torch.zeros((latent_t, 1, 1), dtype=dtype) # [T,1,1]
for fidx in condition_set_vision:
vision_condition_mask[fidx, 0, 0] = 1.0
- vision.condition_mask.append(vision_condition_mask)
+ vision.condition_mask.append(to_device_nonblocking(vision_condition_mask, device))
- vision_noisy_frame_indexes = torch.tensor(
- [idx for idx in range(latent_t) if idx not in condition_set_vision],
- device=device,
- dtype=torch.long,
+ vision_noisy_frame_indexes = to_device_nonblocking(
+ torch.tensor([idx for idx in range(latent_t) if idx not in condition_set_vision], dtype=torch.long),
+ device,
) # [N_noisy_frames]
vision.noisy_frame_indexes.append(vision_noisy_frame_indexes)
diff --git a/cosmos_framework/inference/common/distillation_export.py b/cosmos_framework/inference/common/distillation_export.py
index 43d73371b..82acccc8d 100644
--- a/cosmos_framework/inference/common/distillation_export.py
+++ b/cosmos_framework/inference/common/distillation_export.py
@@ -8,6 +8,18 @@
from pathlib import Path, PurePath
from typing import Any
+_PLACEHOLDER_BUCKET = "bucket"
+
+
+def _placeholder_object_store_uri(value: str) -> str:
+ """Replace the bucket in a ``scheme://bucket/key`` URI, keeping the key."""
+ scheme, _, remainder = value.partition("://")
+ _, _, key = remainder.partition("/")
+ return f"{scheme}://{_PLACEHOLDER_BUCKET}/{key}" if key else f"{scheme}://{_PLACEHOLDER_BUCKET}"
+
+
+# Where the LiDAR VAE lives inside a published artifact. One per artifact, so no filename map.
+_PUBLIC_LIDAR_VAE_PATH = "pretrained/tokenizers/lidar/diffusion_pytorch_model.safetensors"
_PUBLIC_WAN_VAE_PATHS = {
"Wan2.2_VAE.pth": "pretrained/tokenizers/video/wan2pt2/Wan2.2_VAE.pth",
}
@@ -176,20 +188,32 @@ def sanitize_student_public_model_config(
if not isinstance(config, dict):
raise TypeError("Expected model config to be a dictionary.")
- for tokenizer_key in ("tokenizer", "sound_tokenizer"):
+ for tokenizer_key in ("tokenizer", "sound_tokenizer", "lidar_tokenizer"):
tokenizer_config = config.get(tokenizer_key)
if not isinstance(tokenizer_config, dict):
continue
if "bucket_name" in tokenizer_config:
tokenizer_config["bucket_name"] = "bucket"
if "object_store_credential_path_pretrained" in tokenizer_config:
- tokenizer_config["object_store_credential_path_pretrained"] = ""
+ # None, not "": VideoTokenizerInterface treats None as "no credentials" and anything
+ # else as a path, so an empty string is a path that does not exist and it raises.
+ tokenizer_config["object_store_credential_path_pretrained"] = None
if tokenizer_key == "tokenizer" and "vae_path" in tokenizer_config:
tokenizer_config["vae_path"] = _normalize_public_dependency_path(
tokenizer_config["vae_path"],
field_name="tokenizer.vae_path",
public_paths=_PUBLIC_WAN_VAE_PATHS,
)
+ if tokenizer_key == "lidar_tokenizer":
+ # Swept rather than listed: V0 adds latent_stats_path on the same bucket, and the
+ # next such field would be missed again.
+ for field, value in list(tokenizer_config.items()):
+ if field != "vae_path" and isinstance(value, str) and "://" in value:
+ tokenizer_config[field] = _placeholder_object_store_uri(value)
+ if tokenizer_key == "lidar_tokenizer" and "vae_path" in tokenizer_config:
+ # The internal LiDAR path is a full URI, so the bucket rides inside the value and
+ # a filename map would not remove it. Rewrite to the published location.
+ tokenizer_config["vae_path"] = _PUBLIC_LIDAR_VAE_PATH
vlm_config = config.get("vlm_config")
if not isinstance(vlm_config, dict):
diff --git a/cosmos_framework/inference/common/distillation_export_test.py b/cosmos_framework/inference/common/distillation_export_test.py
index 716cfb11a..d1d874f4e 100644
--- a/cosmos_framework/inference/common/distillation_export_test.py
+++ b/cosmos_framework/inference/common/distillation_export_test.py
@@ -264,12 +264,12 @@ def test_sanitize_student_public_model_config_removes_internal_loaders() -> None
"config": {
"tokenizer": {
"bucket_name": "bucket",
- "object_store_credential_path_pretrained": "",
+ "object_store_credential_path_pretrained": None,
"vae_path": "pretrained/tokenizers/video/wan2pt2/Wan2.2_VAE.pth",
},
"sound_tokenizer": {
"bucket_name": "bucket",
- "object_store_credential_path_pretrained": "",
+ "object_store_credential_path_pretrained": None,
"avae_path": "pretrained/tokenizers/audio/avae/avae.ckpt",
},
"vlm_config": {
@@ -477,3 +477,58 @@ class _DefaultBaseModel:
model_cls, _ = distillation_export.resolve_student_base_model({}, default_base_model=_DefaultBaseModel)
assert model_cls is _DefaultBaseModel
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_public_export_removes_every_lidar_internal_path() -> None:
+ """vae_path goes to the published location; any other object-store path loses its bucket.
+
+ V0 adds latent_stats_path, pointing at the same internal bucket, so rewriting vae_path alone
+ published it through a field nobody had listed.
+ """
+ from cosmos_framework.inference.common.distillation_export import sanitize_student_public_model_config
+
+ model = {
+ "config": {
+ "lidar_tokenizer": {
+ "bucket_name": "an-internal-bucket",
+ "object_store_credential_path_pretrained": "credentials/secret.secret",
+ "vae_path": "s3://an-internal-bucket/lidar/tokenizer/v0/checkpoints/iter_1.pt",
+ "latent_stats_path": "s3://an-internal-bucket/lidar/tokenizer/latent_stats.pt",
+ }
+ }
+ }
+ sanitize_student_public_model_config(model)
+ lidar = model["config"]["lidar_tokenizer"]
+ assert lidar["bucket_name"] == "bucket"
+ # None, not "": VideoTokenizerInterface reads None as "no credentials" and anything else as
+ # a path, so an empty string is a path that does not exist and it raises at load time. The
+ # published artifact only surfaced this end to end -- this assertion had the wrong value too.
+ assert lidar["object_store_credential_path_pretrained"] is None
+ assert lidar["vae_path"] == "pretrained/tokenizers/lidar/diffusion_pytorch_model.safetensors"
+ assert lidar["latent_stats_path"] == "s3://bucket/lidar/tokenizer/latent_stats.pt"
+ assert "an-internal-bucket" not in str(model)
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_public_export_leaves_relative_lidar_fields_alone() -> None:
+ """A value with no scheme carries no bucket. vae_path is rewritten regardless, because the
+ internal name carries a training iteration and the published artifact is safetensors."""
+ from cosmos_framework.inference.common.distillation_export import sanitize_student_public_model_config
+
+ model = {
+ "config": {
+ "lidar_tokenizer": {
+ "vae_path": "tokenizer/v1p2/iter_000030000.pt",
+ "latent_stats_path": "tokenizer/latent_stats.pt",
+ "dtype": "float32",
+ }
+ }
+ }
+ sanitize_student_public_model_config(model)
+ lidar = model["config"]["lidar_tokenizer"]
+ assert lidar["vae_path"] == "pretrained/tokenizers/lidar/diffusion_pytorch_model.safetensors"
+ assert lidar["latent_stats_path"] == "tokenizer/latent_stats.pt"
+ assert lidar["dtype"] == "float32"
diff --git a/cosmos_framework/model/attention/frontend.py b/cosmos_framework/model/attention/frontend.py
index 6d3e8e917..cd9a7a417 100644
--- a/cosmos_framework/model/attention/frontend.py
+++ b/cosmos_framework/model/attention/frontend.py
@@ -562,7 +562,7 @@ def multi_dimensional_attention_varlen(
value (Tensor): 4-D value tensor with sequence-packed layout
(`[1, seqlen_total, heads_kv, head_dim_v]`)
- metadata (dict): Pre-computed varlen metadata from `cosmos_framework.varlen.generate_multi_dim_varlen_parameters`.
+ metadata (dict): Pre-computed varlen metadata from `cosmos_framework.model.attention.varlen.generate_multi_dim_varlen_parameters`.
scale (float | None): Attention scale. Defaults to head_dim ** -0.5.
diff --git a/cosmos_framework/model/generator/mot/activation_marks.py b/cosmos_framework/model/generator/mot/activation_marks.py
new file mode 100644
index 000000000..2ac2487be
--- /dev/null
+++ b/cosmos_framework/model/generator/mot/activation_marks.py
@@ -0,0 +1,143 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+
+"""Marking individual call sites for selective activation checkpointing.
+
+``save_ops_regex`` selects by dispatched op name, which cannot separate calls
+that run the same kernel. The decomposed multiview attention runs four FMHA
+calls per layer -- three sensor folds plus the causal pass -- and they are worth
+very different amounts to keep: the same-view fold is ~96% of forward attention
+time and ~94% of backward, while the others are cheap to recompute. A regex
+either keeps all four or none.
+
+A call site marks itself instead::
+
+ k = mark_next_activation(k)
+ out, lse = attention(q, k, v, ...)
+
+The policy sees ``cosmos3::keep_next_activation`` before the op it marks and keeps
+that one. What counts as markable is ``save_ops_regex``: a mark stays pending
+until an op matching it arrives, so the gathers, reshapes and the clone the
+attention frontend makes in between cannot take the mark -- four such ops sit
+between the marker and the kernel at the real call site.
+
+Marking only decides anything when ``save_only_marked_ops`` is set. Without it the
+regex keeps every op it matches, marked or not, which is what existing configs do
+-- and then the marker is pure cost, so it is not emitted at all: ``enable_marking``
+is called from ``_apply_selective_ac`` when that config field is on, and until it is,
+``mark_next_activation`` hands its argument straight back. The switch is read while
+Dynamo traces, so it is baked into the compiled graph rather than tested per step.
+
+The switch only ever turns on, and production has no way to turn it off -- see
+``enable_marking`` for why turning it off from a second model's config is a silent
+whole-run regression rather than a saving.
+
+The marker must be functional and its result must be consumed, which means it
+copies what it marks. Mark the smallest tensor the call takes. On the AV shape
+that is K or V: Cosmos3 16B runs 32 query heads against 8 KV heads, and under
+CP16 that is 2 against 1 per rank, so the copy is a quarter to a half of the
+output the fold would otherwise stash. The arithmetic still comes out ahead: on a
+three-call stand-in, marking one rather than saving all three moved per-step peak
+from 2574 MiB to 1928 MiB, clone included.
+
+The copy is transient. The marker is ``MUST_RECOMPUTE``, so inside a checkpointed
+region it is rebuilt during recompute rather than stored, and nothing carries
+across steps -- measured at 0 MiB resident after each of four consecutive steps.
+"""
+
+import torch
+
+_MARK_OP_NAME = "keep_next_activation"
+
+MARK_OP_QUALNAME = f"cosmos3::{_MARK_OP_NAME}"
+
+
+@torch.library.custom_op(MARK_OP_QUALNAME, mutates_args=())
+def _keep_next_activation(tensor: torch.Tensor) -> torch.Tensor:
+ # A copy, not a view: inductor rejects custom ops whose output aliases an
+ # input, and a functional op's result has to be consumed for the marker to
+ # keep its place in the trace.
+ return tensor.clone()
+
+
+@_keep_next_activation.register_fake
+def _(tensor: torch.Tensor) -> torch.Tensor:
+ return torch.empty_like(tensor)
+
+
+def _backward(ctx, grad: torch.Tensor) -> torch.Tensor:
+ return grad
+
+
+_keep_next_activation.register_autograd(_backward)
+
+
+_MARKING_ENABLED = False
+
+
+def enable_marking() -> None:
+ """Turn call-site marking on for this process. There is deliberately no way back.
+
+ Called from ``parallelize_unified_mot._apply_selective_ac`` when the model's
+ ``activation_checkpointing.save_only_marked_ops`` is set -- at setup, before the
+ first forward, so the answer is fixed by the time anything traces.
+
+ Process-level because the call sites that mark are deep inside attention and have
+ no view of the checkpointing config, and one-way because the two mistakes are not
+ the same size. Dynamo installs an equality guard on the switch, so turning it off
+ for a second model invalidates the first model's compiled code on its very next
+ forward -- no shape change needed -- and the recompiled graph has no marker in it.
+ A policy built with ``save_only_marked_ops`` then keeps *nothing*: the model silently
+ reverts to recomputing all four folds, ~50 ms/step, with no error. Measured, not
+ reasoned about: the same block goes from saving ``[False, True, False]`` to
+ ``[False, False, False]`` across that transition.
+
+ The cost of the opposite mistake is a clone a later model does not need. It is
+ bounded and visible: only the maskless folds mark at all, the marker is
+ ``MUST_RECOMPUTE`` so it is 0 bytes resident under selective or full AC, and it is
+ a quarter to a half of one fold's output in transient bandwidth. So the switch
+ latches, and a model that would rather not pay it scopes marking to itself.
+ """
+ global _MARKING_ENABLED
+ _MARKING_ENABLED = True
+
+
+def reset_marking_for_tests() -> None:
+ """Turn marking back off. Tests only -- see ``enable_marking`` for why.
+
+ Named for its one legitimate caller so that a config-driven ``enable_marking(False)``
+ cannot be written by accident; that call is the regression described above.
+ """
+ global _MARKING_ENABLED
+ _MARKING_ENABLED = False
+
+
+def marking_enabled() -> bool:
+ """Whether ``mark_next_activation`` currently emits anything."""
+ return _MARKING_ENABLED
+
+
+def mark_next_activation(tensor: torch.Tensor) -> torch.Tensor:
+ """Ask selective AC to keep the output of the next op that reads ``tensor``.
+
+ Pass the smallest tensor the marked call takes, and use its return value::
+
+ k = mark_next_activation(k)
+ out, lse = attention(q, k, v, ...)
+
+ Binds to the next op matching ``save_ops_regex``, in trace order, rather than
+ to an index: a call site behind a branch cannot shift the mark onto its
+ neighbour, and the reshapes between the mark and the kernel cannot absorb it.
+
+ A pass-through unless ``enable_marking`` has been called, so a model that does not
+ use ``save_only_marked_ops`` pays nothing for the call sites that mark. When it is
+ on the marker is a copy, which the checkpoint recomputes and discards.
+ """
+ if not _MARKING_ENABLED:
+ return tensor
+ return _keep_next_activation(tensor)
+
+
+def is_mark_op(op_name: str) -> bool:
+ """Whether a dispatched op name is the marker."""
+ return _MARK_OP_NAME in op_name
diff --git a/cosmos_framework/model/generator/mot/activation_marks_test.py b/cosmos_framework/model/generator/mot/activation_marks_test.py
new file mode 100644
index 000000000..efbb6bedb
--- /dev/null
+++ b/cosmos_framework/model/generator/mot/activation_marks_test.py
@@ -0,0 +1,830 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+
+"""Call-site marking for selective activation checkpointing.
+
+Runs against ``cosmos_framework.model.attention`` with varlen arguments, which is what the
+decomposed multiview attention calls and what selects NATTEN's varlen FMHA on
+Blackwell -- a ``torch.library`` custom op inductor cannot decompose, rather than
+an aten op it can. ``F.scaled_dot_product_attention`` would exercise neither.
+
+The compiled cases carry the weight: three cheaper marking mechanisms work eager
+and fail under ``torch.compile``, which is how a decoder layer runs, so an
+eager-only suite would pass on all of them.
+"""
+
+import gc
+import re
+
+import pytest
+import torch
+from torch import nn
+from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import checkpoint_wrapper as ptd_checkpoint_wrapper
+from torch.utils.checkpoint import CheckpointPolicy, create_selective_checkpoint_contexts
+
+from cosmos_framework.model.attention import attention
+from cosmos_framework.utils.helper_test import RunIf
+from cosmos_framework.model.generator.mot.activation_marks import (
+ MARK_OP_QUALNAME,
+ enable_marking,
+ is_mark_op,
+ mark_next_activation,
+ marking_enabled,
+ reset_marking_for_tests,
+)
+from cosmos_framework.model.generator.mot.parallelize_unified_mot import make_selective_ac_policy
+
+HEADS = 8
+KV_HEADS = 2
+HEAD_DIM = 128
+HIDDEN = HEADS * HEAD_DIM
+TOKENS = 4096
+SEGMENTS = 2
+SEGMENT = TOKENS // SEGMENTS
+
+
+def _policy(
+ decisions: list[tuple[str, bool]], save_ops: tuple[str, ...] = ("fmha",), save_only_marked_ops: bool = True
+):
+ """The shipped policy, wrapped to record what it decided.
+
+ Deliberately not a reimplementation. An earlier version of this file had one,
+ it consulted the mark only on attention ops while the real policy consulted it
+ on every op, and it hid a defect where the mark landed on the gather that
+ precedes the kernel rather than on the kernel.
+ """
+ real = make_selective_ac_policy([re.compile(pattern) for pattern in save_ops], save_only_marked_ops)
+
+ def policy(ctx, func, *args, **kwargs) -> CheckpointPolicy:
+ verdict = real(ctx, func, *args, **kwargs)
+ name = getattr(func, "__name__", str(func))
+ decisions.append((name, verdict == CheckpointPolicy.MUST_SAVE))
+ return verdict
+
+ return policy
+
+
+@pytest.fixture(autouse=True)
+def _marking_off_by_default():
+ """Every test states its own precondition; the switch is process-level.
+
+ Without this a test that enables marking would leave it on for the rest of the
+ session, and the cases that check what happens *without* marks would stop
+ checking anything.
+ """
+ reset_marking_for_tests()
+ yield
+ reset_marking_for_tests()
+
+
+def _natten_can_run_this_block() -> bool:
+ """Whether NATTEN's varlen FMHA is usable here for the shapes ``_Block`` runs.
+
+ The blocks pin ``backend="natten"`` rather than letting the frontend choose, because
+ the policy matches ops by name and only NATTEN's kernels are named "fmha". NATTEN is
+ the frontend's pick for varlen on Blackwell but not on Hopper, where flash3 outranks
+ it -- pinning is what keeps this coverage on both. Pinning an *incompatible* backend
+ raises, though, so compatibility is checked here and the suite skips rather than
+ errors on a device NATTEN cannot serve.
+ """
+ if not torch.cuda.is_available():
+ return False
+ try:
+ from cosmos_framework.model.attention.backends import is_backend_compatible
+
+ return is_backend_compatible(
+ backend="natten",
+ query_shape=torch.Size((1, TOKENS, HEADS, HEAD_DIM)),
+ key_shape=torch.Size((1, TOKENS, KV_HEADS, HEAD_DIM)),
+ value_shape=torch.Size((1, TOKENS, KV_HEADS, HEAD_DIM)),
+ dtype=torch.bfloat16,
+ device=torch.device("cuda"),
+ requires_grad=True,
+ is_causal=False,
+ causal_type=None,
+ is_varlen=True,
+ raise_error=False,
+ )
+ except Exception: # noqa: BLE001 - a gate that cannot answer should skip, not error.
+ return False
+
+
+_NATTEN_RUNS_HERE = _natten_can_run_this_block()
+
+
+def _natten_is_the_frontend_choice() -> bool:
+ """Whether the frontend would *pick* NATTEN here, not merely whether it can run.
+
+ A different question from ``_natten_can_run_this_block``, and the peak-memory
+ comparison below is the one place the distinction matters: what the blocks pin is
+ NATTEN's kernel, but what the surrounding cost of recomputing looks like is the
+ arch's, and those diverge.
+ """
+ if not torch.cuda.is_available():
+ return False
+ try:
+ from cosmos_framework.model.attention.backends import choose_backend
+
+ return (
+ choose_backend(
+ query_shape=torch.Size((1, TOKENS, HEADS, HEAD_DIM)),
+ key_shape=torch.Size((1, TOKENS, KV_HEADS, HEAD_DIM)),
+ value_shape=torch.Size((1, TOKENS, KV_HEADS, HEAD_DIM)),
+ dtype=torch.bfloat16,
+ device=torch.device("cuda"),
+ requires_grad=True,
+ is_causal=False,
+ causal_type=None,
+ is_varlen=True,
+ raise_error=False,
+ )
+ == "natten"
+ )
+ except Exception: # noqa: BLE001 - a gate that cannot answer should skip, not error.
+ return False
+
+
+_NATTEN_IS_THE_FRONTEND_CHOICE = _natten_is_the_frontend_choice()
+
+
+def _saved_attention(decisions: list[tuple[str, bool]]) -> list[bool]:
+ return [saved for name, saved in decisions if "fmha" in name or "attention" in name]
+
+
+def _saved_names(decisions: list[tuple[str, bool]]) -> list[str]:
+ return [name for name, saved in decisions if saved]
+
+
+class _Block(nn.Module):
+ """Three varlen attention calls; ``mark_at`` says which one marks itself.
+
+ Three because the decomposed layer runs the same kernel several times over
+ different folds, which is the case a name-matching policy cannot separate.
+ """
+
+ def __init__(self, mark_at: int | tuple[int, ...] | None) -> None:
+ super().__init__()
+ self.marks = () if mark_at is None else (mark_at,) if isinstance(mark_at, int) else tuple(mark_at)
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ # Derived from the input rather than fixed at construction, so a second token
+ # count retraces the region instead of failing the view.
+ tokens = x.shape[0]
+ segment = tokens // SEGMENTS
+ offsets = torch.arange(SEGMENTS + 1, dtype=torch.int32, device=x.device) * segment
+ for index in range(3):
+ q = x.view(1, tokens, HEADS, HEAD_DIM)
+ kv = q[:, :, :KV_HEADS]
+ if index in self.marks:
+ kv = mark_next_activation(kv)
+ out = attention(
+ q,
+ kv,
+ kv,
+ cumulative_seqlen_Q=offsets,
+ cumulative_seqlen_KV=offsets,
+ max_seqlen_Q=segment,
+ max_seqlen_KV=segment,
+ backend="natten",
+ )
+ x = out.reshape(tokens, HIDDEN)
+ return x
+
+
+class _InterveningBlock(nn.Module):
+ """Marks call 1, then evaluates its other operands -- as the real call site does."""
+
+ def __init__(self) -> None:
+ super().__init__()
+ self.register_buffer("offsets", torch.arange(SEGMENTS + 1, dtype=torch.int32) * SEGMENT, persistent=False)
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ for index in range(3):
+ base = x.view(1, TOKENS, HEADS, HEAD_DIM)
+ kv = base[:, :, :KV_HEADS]
+ if index == 1:
+ kv = mark_next_activation(kv)
+ # Gathers and slices dispatched *after* the mark, before the kernel.
+ gather = torch.arange(TOKENS, device=x.device)
+ q = base[:, gather]
+ v = base[:, gather][:, :, :KV_HEADS]
+ out = attention(
+ q,
+ kv,
+ v,
+ cumulative_seqlen_Q=self.offsets,
+ cumulative_seqlen_KV=self.offsets,
+ max_seqlen_Q=SEGMENT,
+ max_seqlen_KV=SEGMENT,
+ backend="natten",
+ )
+ x = out.reshape(TOKENS, HIDDEN)
+ return x
+
+
+def _run(
+ mark_at: int | tuple[int, ...] | None,
+ compiled: bool,
+ save_ops: tuple[str, ...] = ("fmha",),
+ save_only_marked_ops: bool = True,
+ block: nn.Module | None = None,
+) -> tuple[list[tuple[str, bool]], torch.Tensor]:
+ torch._dynamo.reset()
+ decisions: list[tuple[str, bool]] = []
+ # Where ``_apply_selective_ac`` does it: before anything is traced. Inside
+ # ``context_fn`` would be too late under compile, since the region is traced before
+ # the partitioner calls it -- the marker would already have been a pass-through, and
+ # the two compiled cases that check marking fail exactly that way.
+ enable_marking()
+ block = ptd_checkpoint_wrapper(
+ (block if block is not None else _Block(mark_at)).cuda(),
+ context_fn=lambda: create_selective_checkpoint_contexts(_policy(decisions, save_ops, save_only_marked_ops)),
+ )
+ if compiled:
+ block = torch.compile(block, fullgraph=True)
+ torch.manual_seed(0)
+ x = torch.randn(TOKENS, HIDDEN, device="cuda", dtype=torch.bfloat16, requires_grad=True)
+ block(x).float().sum().backward()
+ torch.cuda.synchronize()
+ return decisions, x.grad.clone()
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_the_attention_op_regex_covers_every_backend_that_can_be_selected() -> None:
+ """The marked policy's eligibility list, checked against the ops actually registered.
+
+ A marked call site is only kept if some op matches ``save_ops_regex``, so a backend
+ missing from that list makes marking silently inert -- the config says it is on, no
+ op is eligible, the mark is never consumed and the layer recomputes everything. Not
+ hypothetical: ``["fmha"]`` covers NATTEN on every arch but nothing else, and on sm90
+ flash3 is ranked ahead of NATTEN and takes these calls, which is how the mechanism
+ came to do nothing on Hopper while reading as enabled.
+
+ Enumerated from the dispatcher rather than hard-coded, so a renamed or newly added
+ kernel fails here instead of in a training run's memory profile. cuDNN is out of
+ scope by design: it rejects varlen and every marked call site is varlen.
+ """
+ from cosmos_framework.model.attention.backends import get_backend_list
+ from cosmos_framework.configs.base.defaults.activation_checkpointing import (
+ ATTENTION_FORWARD_OPS_REGEX,
+ )
+
+ patterns = [re.compile(pattern) for pattern in ATTENTION_FORWARD_OPS_REGEX]
+
+ def covered(op_name: str) -> bool:
+ # What the policy matches on: ``func.__name__``, which drops the namespace.
+ return any(pattern.search(f"{op_name.split('::')[-1]}.default") for pattern in patterns)
+
+ # Backend name -> the namespace its ops register under. ``None`` means no marked call
+ # site can reach it: cuDNN rejects varlen and every mark sits on a varlen fold.
+ namespaces = {
+ "natten": "natten::",
+ "flash2": "flash_attn::",
+ "flash3": "flash_attn_3::",
+ "flash4": "flash_attn_4::",
+ "cudnn": None,
+ }
+ # Every backend the frontend can return, on any arch it supports. Driving the check
+ # from here rather than a hardcoded list is the point: ``flash4`` is already installed
+ # in the GB200 image but commented out of ``get_backend_list``, and the day it is
+ # enabled it ranks ahead of NATTEN on sm100 -- at which point marking would go inert
+ # on the primary training arch exactly as it did on Hopper, and this should say so.
+ selectable = {backend for arch in (75, 80, 86, 90, 100, 103, 110, 120, 121) for backend in get_backend_list(arch)}
+ assert selectable, "no backend is selectable anywhere, so this would pass vacuously"
+ unknown = selectable - set(namespaces)
+ assert not unknown, f"a backend was added without saying where its ops live: {sorted(unknown)}"
+
+ registered = torch._C._dispatch_get_all_op_names()
+ prefixes = tuple(namespaces[backend] for backend in selectable if namespaces[backend])
+ # NATTEN also registers neighbourhood kernels (na1d/na2d/na3d), which serve
+ # ``multi_dimensional_attention`` rather than the decomposed folds. Nothing marks into
+ # those, so they are out of scope: the folds dispatch full attention.
+ neighbourhood = ("na1d", "na2d", "na3d")
+
+ def in_scope(name: str) -> bool:
+ return name.startswith(prefixes) and not any(kernel in name for kernel in neighbourhood)
+
+ forwards = [name for name in registered if in_scope(name) and name.endswith("_forward")]
+ # flash3 ships as ``flash_attn_3_nv`` and is absent from aarch64 images, where it is
+ # not a candidate anyway. Its op name is pinned by docker/Dockerfile.base and read from
+ # that tag's source, so cover it whether or not this image has the package.
+ if "flash3" in selectable:
+ forwards.append("flash_attn_3::_flash_attn_forward")
+
+ assert forwards, "no attention backend registered an op, so this would pass vacuously"
+ missing = [name for name in forwards if not covered(name)]
+ assert not missing, f"selectable backends the marked policy cannot keep: {missing}"
+
+ # Backward ops must not match: the policy runs over the forward, and a pattern loose
+ # enough to catch them is matching on something other than what it means.
+ caught = [name for name in registered if in_scope(name) and name.endswith("_backward") and covered(name)]
+ assert not caught, f"the regex reaches backward ops: {caught}"
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_the_schema_default_is_the_regex_the_test_above_validates() -> None:
+ """Every selective config inherits the list checked against the dispatcher, not a subset.
+
+ The default used to be ``["fmha"]``, which is NATTEN's op name rather than attention's:
+ it matched on sm100, where cuDNN and flash2 both reject varlen and NATTEN wins the call,
+ and matched nothing on sm90, where flash3 is ranked first and takes it. So a config
+ asking for selective AC got it on one arch and full recompute on the other, silently
+ and with the same ``mode="selective"`` in both. The constant is what makes that
+ arch-independent, and it only helps the configs that get it by default -- a config
+ naming its own list is a config that can be wrong about the backend again.
+
+ Copied per instance rather than shared: ``attrs`` hands the factory's value to the
+ config, so a returned module constant would let one config's ``append`` rewrite the
+ policy of every other config built in the same process.
+ """
+ from cosmos_framework.configs.base.defaults.activation_checkpointing import (
+ ATTENTION_FORWARD_OPS_REGEX,
+ ActivationCheckpointingConfig,
+ )
+
+ assert ActivationCheckpointingConfig().save_ops_regex == ATTENTION_FORWARD_OPS_REGEX
+
+ mutated = ActivationCheckpointingConfig()
+ mutated.save_ops_regex.append("randn")
+ assert ActivationCheckpointingConfig().save_ops_regex == ATTENTION_FORWARD_OPS_REGEX
+ assert "randn" not in ATTENTION_FORWARD_OPS_REGEX
+
+
+@RunIf(min_gpus=1)
+@pytest.mark.L0
+@pytest.mark.GPU
+@pytest.mark.skipif(
+ not _NATTEN_RUNS_HERE, reason="These blocks pin NATTEN's varlen FMHA, which this device cannot run."
+)
+class TestActivationMarks:
+ """Marking one call among several that dispatch the same kernel."""
+
+ @pytest.mark.parametrize("compiled", [False, True])
+ def test_only_the_marked_call_is_saved(self, compiled: bool) -> None:
+ """The point of the mechanism, and the compiled case is the one that matters.
+
+ The decomposed attention runs four identical FMHA ops per layer, so a
+ name-matching policy keeps all of them or none.
+ """
+ decisions, _ = _run(mark_at=1, compiled=compiled)
+ assert _saved_attention(decisions)[:3] == [False, True, False], _saved_names(decisions)
+
+ @pytest.mark.parametrize("compiled", [False, True])
+ def test_an_unmarked_block_saves_nothing(self, compiled: bool) -> None:
+ """Under ``save_only_marked_ops`` an eligible op still needs a mark."""
+ decisions, _ = _run(mark_at=None, compiled=compiled)
+ assert _saved_names(decisions) == [], _saved_names(decisions)
+
+ @pytest.mark.parametrize("compiled", [False, True])
+ def test_without_save_only_marked_ops_the_regex_decides_alone(self, compiled: bool) -> None:
+ """The default, and what every existing config relies on.
+
+ Marks are inert here: the regex keeps every op it matches, so a marked
+ block and an unmarked one save the same set.
+ """
+ marked, _ = _run(mark_at=1, compiled=compiled, save_only_marked_ops=False)
+ unmarked, _ = _run(mark_at=None, compiled=compiled, save_only_marked_ops=False)
+ assert _saved_attention(marked)[:3] == [True, True, True], _saved_names(marked)
+ assert _saved_attention(unmarked)[:3] == [True, True, True], _saved_names(unmarked)
+
+ @pytest.mark.parametrize("compiled", [False, True])
+ def test_a_mark_is_not_taken_by_ops_between_it_and_the_kernel(self, compiled: bool) -> None:
+ """A pending mark waits for an *eligible* op, not the next op of any kind.
+
+ The real call site evaluates the gathers for Q and V after the mark line,
+ and the attention frontend clones internally, so four ops sit between the
+ marker and the kernel -- index, index, slice, clone. Consuming on the next
+ op of any kind put the mark on a gather and saved that instead.
+ """
+ decisions, _ = _run(mark_at=1, compiled=compiled, block=_InterveningBlock())
+ assert _saved_attention(decisions)[:3] == [False, True, False], _saved_names(decisions)
+ assert all("fmha" in name or "attention" in name for name in _saved_names(decisions)), _saved_names(decisions)
+
+ @pytest.mark.parametrize("compiled", [False, True])
+ def test_marking_does_not_change_gradients(self, compiled: bool) -> None:
+ """Marking decides what is kept, never what is computed.
+
+ Compared against a floor the test measures rather than a fixed tolerance:
+ SDPA's backward is not deterministic, so the same configuration run twice
+ already differs, and bitwise equality would be unreachable for reasons
+ that have nothing to do with marking.
+ """
+ _, plain = _run(mark_at=None, compiled=compiled)
+ _, plain_again = _run(mark_at=None, compiled=compiled)
+ _, marked = _run(mark_at=1, compiled=compiled)
+
+ assert (plain != 0).any(), "a zero gradient would compare equal and prove nothing"
+ scale = plain.abs().max()
+ # Two floors, because either alone is misleading. The run-to-run spread
+ # catches a non-deterministic backward, but two runs can land bit-identical
+ # and report zero; bf16 round-off at this scale is the floor underneath
+ # that. Both were measured at 6.25e-02 on the NATTEN path, which is ~1.3e-3
+ # relative -- bf16's mantissa, not a difference marking caused.
+ observed = (plain - plain_again).abs().max()
+ floor = torch.maximum(observed * 2, scale * 2**-8)
+ difference = (marked - plain).abs().max()
+ assert difference <= floor, f"marked differs by {difference}; run-to-run {observed}, bf16 floor {scale * 2**-8}"
+
+ @pytest.mark.parametrize("compiled", [False, True])
+ def test_each_mark_is_consumed_by_exactly_one_call(self, compiled: bool) -> None:
+ """Two marked call sites keep two calls, and the unmarked one between them is not kept.
+
+ A mark that armed the policy without disarming it would keep every eligible op
+ after the first mark, which looks identical to working when only one call is
+ marked. Marking calls 0 and 2 is what separates the two.
+ """
+ decisions, _ = _run(mark_at=(0, 2), compiled=compiled)
+ assert _saved_attention(decisions)[:3] == [True, False, True], _saved_names(decisions)
+
+ def test_marks_do_not_leak_between_regions_under_compile(self) -> None:
+ """The eager leak check, compiled -- where the policy runs at trace time.
+
+ Two blocks share nothing but the process. If policy state outlived a region, the
+ second block's calls would inherit the first block's marking, and under compile
+ that mistake is baked into a graph rather than made once per step.
+
+ Each block records into its own list, because the compiled region is traced more
+ than once and a single list interleaves the passes.
+ """
+ torch._dynamo.reset()
+ enable_marking()
+ marked_block: list[tuple[str, bool]] = []
+ plain_block: list[tuple[str, bool]] = []
+ blocks = [
+ torch.compile(
+ ptd_checkpoint_wrapper(
+ _Block(marks).cuda(),
+ context_fn=lambda decisions=decisions: create_selective_checkpoint_contexts(_policy(decisions)),
+ ),
+ fullgraph=True,
+ )
+ for marks, decisions in ((2, marked_block), (None, plain_block))
+ ]
+ x = torch.randn(TOKENS, HIDDEN, device="cuda", dtype=torch.bfloat16, requires_grad=True)
+ blocks[1](blocks[0](x)).float().sum().backward()
+ torch.cuda.synchronize()
+
+ folds = _saved_attention(marked_block)
+ assert folds and len(folds) % 3 == 0, f"expected whole passes over three calls, saw {len(folds)}"
+ assert folds == [False, False, True] * (len(folds) // 3), _saved_names(marked_block)
+ assert _saved_names(plain_block) == [], _saved_names(plain_block)
+
+ def test_marking_holds_across_repeated_steps(self) -> None:
+ """What is kept must not drift once training is under way.
+
+ The policy carries state -- a pending-mark flag -- and ``context_fn`` is called
+ once per checkpointed forward. A flag left armed at the end of one step would
+ keep the *first* call of the next one, silently and only from step two, which a
+ single-step test cannot see.
+
+ Eager, because that is where there is something per-step to watch: under compile
+ the policy is consulted while tracing and not again, so later steps record no
+ decisions at all. The compiled equivalent of this claim is that the graph does
+ not drift, which ``test_a_compiled_marked_block_is_stable_across_steps`` checks
+ through its gradients and its memory instead.
+ """
+ torch._dynamo.reset()
+ enable_marking()
+ decisions: list[tuple[str, bool]] = []
+ block = ptd_checkpoint_wrapper(
+ _Block(1).cuda(),
+ context_fn=lambda: create_selective_checkpoint_contexts(_policy(decisions)),
+ )
+ x = torch.randn(TOKENS, HIDDEN, device="cuda", dtype=torch.bfloat16, requires_grad=True)
+ for step in range(4):
+ decisions.clear()
+ block(x).float().sum().backward()
+ x.grad = None
+ torch.cuda.synchronize()
+ assert _saved_attention(decisions)[:3] == [False, True, False], f"step {step}: {_saved_names(decisions)}"
+
+ def test_a_compiled_marked_block_is_stable_across_steps(self) -> None:
+ """A compiled marked block repeats itself exactly, and holds nothing between steps.
+
+ Under compile the policy is consulted once while tracing, so what a later step
+ does can only be read off its results. Gradients that stay put say the graph is
+ not being re-partitioned differently, and flat memory says the kept activation
+ and the marker's copy are released with the graph rather than accumulating.
+
+ Against a measured floor rather than bitwise, for the reason
+ ``test_marking_does_not_change_gradients`` gives: NATTEN's backward is not
+ deterministic, so two runs of one configuration already differ slightly.
+ """
+ torch._dynamo.reset()
+ enable_marking()
+ block = torch.compile(
+ ptd_checkpoint_wrapper(
+ _Block(1).cuda(),
+ context_fn=lambda: create_selective_checkpoint_contexts(_policy([])),
+ ),
+ fullgraph=True,
+ )
+ torch.manual_seed(0)
+ x = torch.randn(TOKENS, HIDDEN, device="cuda", dtype=torch.bfloat16, requires_grad=True)
+
+ block(x).float().sum().backward()
+ first = x.grad.clone()
+ x.grad = None
+ block(x).float().sum().backward()
+ # The same step twice: whatever these differ by is the backward's own noise, and
+ # anything marking did wrong later has to clear it to be visible.
+ floor = torch.maximum((x.grad - first).abs().max() * 2, first.abs().max() * 2**-8)
+ x.grad = None
+ torch.cuda.synchronize()
+ gc.collect()
+ torch.cuda.empty_cache()
+ baseline = torch.cuda.memory_allocated()
+
+ resident = []
+ for step in range(3):
+ block(x).float().sum().backward()
+ drift = (x.grad - first).abs().max()
+ assert drift <= floor, f"step {step} drifted by {drift}, floor {floor}"
+ x.grad = None
+ torch.cuda.synchronize()
+ gc.collect()
+ resident.append(torch.cuda.memory_allocated() - baseline)
+ assert len(set(resident)) == 1, f"memory accumulated across steps: {resident}"
+
+ def test_marking_survives_a_recompilation(self) -> None:
+ """A second shape retraces the region, and the new graph must mark it too.
+
+ The switch is read while Dynamo traces, so it is a guard on the compiled graph
+ rather than a test run per step. Packing gives a different token count from step
+ to step, so retracing is routine here rather than exotic, and a retrace that
+ dropped the mark would quietly go back to recomputing the expensive fold.
+
+ The graph count is asserted as well, so the test cannot pass by never retracing
+ -- a dynamic-shape graph reused for both sizes would check nothing.
+ """
+ torch._dynamo.reset()
+ enable_marking()
+ decisions: list[tuple[str, bool]] = []
+ block = torch.compile(
+ ptd_checkpoint_wrapper(
+ _Block(1).cuda(),
+ context_fn=lambda: create_selective_checkpoint_contexts(_policy(decisions)),
+ ),
+ fullgraph=True,
+ )
+ graphs = []
+ for tokens in (TOKENS, TOKENS * 2):
+ decisions.clear()
+ x = torch.randn(tokens, HIDDEN, device="cuda", dtype=torch.bfloat16, requires_grad=True)
+ block(x).float().sum().backward()
+ torch.cuda.synchronize()
+ graphs.append(torch._dynamo.utils.counters["stats"]["unique_graphs"])
+ assert _saved_attention(decisions)[:3] == [False, True, False], (
+ f"{tokens} tokens: {_saved_names(decisions)}"
+ )
+ assert graphs[1] > graphs[0], f"the second shape reused the first graph, so nothing retraced: {graphs}"
+
+ def test_the_marker_passes_its_gradient_through_unchanged(self) -> None:
+ """The marker is bookkeeping, so its backward is the identity.
+
+ Registered by hand, and a marker that scaled or dropped a gradient would move
+ the model's without any test of what is *kept* noticing.
+ """
+ enable_marking()
+ source = torch.randn(64, 32, device="cuda", dtype=torch.float32, requires_grad=True)
+ seed = torch.randn_like(source)
+ mark_next_activation(source).backward(seed)
+ assert source.grad is not None
+ assert torch.equal(source.grad, seed)
+
+ def test_a_forward_without_a_backward_keeps_nothing_resident(self) -> None:
+ """Inference steps interleave with training ones, and must not pin activations.
+
+ Under ``no_grad`` there is no checkpoint to save into, so the marked fold and
+ the marker's own copy should both be gone once the forward returns.
+ """
+ torch._dynamo.reset()
+ enable_marking()
+ block = ptd_checkpoint_wrapper(
+ _Block(1).cuda(),
+ context_fn=lambda: create_selective_checkpoint_contexts(_policy([])),
+ )
+ x = torch.randn(TOKENS, HIDDEN, device="cuda", dtype=torch.bfloat16)
+ with torch.no_grad():
+ block(x)
+ torch.cuda.synchronize()
+ gc.collect()
+ torch.cuda.empty_cache()
+ baseline = torch.cuda.memory_allocated()
+
+ resident = []
+ for _ in range(3):
+ with torch.no_grad():
+ block(x)
+ torch.cuda.synchronize()
+ gc.collect()
+ resident.append(torch.cuda.memory_allocated() - baseline)
+ assert set(resident) == {0}, f"an inference step pinned memory: {resident}"
+
+ @pytest.mark.skipif(
+ not _NATTEN_IS_THE_FRONTEND_CHOICE,
+ reason="Peak-memory ordering at this block size is an arch property; measured on sm100.",
+ )
+ def test_marking_costs_less_memory_than_saving_every_eligible_op(self) -> None:
+ """The arithmetic the mechanism exists for, counting the clone it pays for.
+
+ Keeping the one marked call has to cost less than keeping all three, or the
+ mechanism buys nothing -- and the marker's copy is inside the measurement, so
+ this is the net figure rather than the saving before its cost.
+
+ Only against saving everything, and only where the ordering has been measured.
+ Recomputing is not reliably the cheapest of the three: its transient forward
+ buffers are live alongside the backward's, which at this block's size puts its
+ peak above both others while on the production shape it sits below them.
+
+ The same effect inverts this test's own claim on sm90, where it was measured at
+ 168.5 MiB marked against 160.6 MiB saving all three -- marking retains ~14 MiB
+ less (one output plus the K clone, against three outputs) but pays more than that
+ back in workspace for the two forwards it recomputes. So the bytes comparison is
+ a property of arch and shape rather than of marking, and is asserted only where
+ it was measured. What marking does to the *number* of kept activations is
+ arch-independent and covered by ``test_only_the_marked_call_is_saved``.
+ """
+ peaks = {}
+ for label, mark_at, save_only in (("marked", 1, True), ("save_all", None, False)):
+ torch._dynamo.reset()
+ gc.collect()
+ torch.cuda.empty_cache()
+ torch.cuda.reset_peak_memory_stats()
+ _run(mark_at=mark_at, compiled=True, save_only_marked_ops=save_only)
+ torch.cuda.synchronize()
+ peaks[label] = torch.cuda.max_memory_allocated()
+
+ assert peaks["marked"] < peaks["save_all"], peaks
+
+ def test_a_mark_binds_to_the_next_op_not_an_index(self) -> None:
+ """A call site behind a branch must not shift the mark onto its neighbour.
+
+ ``cross_view`` is conditional and ``gen_to_und`` has two forms, so the
+ number of attention calls varies per step. An ordinal policy would drift;
+ binding to the next op in trace order does not.
+ """
+ for mark_at in (0, 1, 2):
+ decisions, _ = _run(mark_at=mark_at, compiled=False)
+ expected = [index == mark_at for index in range(3)]
+ assert _saved_attention(decisions)[:3] == expected, f"mark at {mark_at}: {_saved_names(decisions)}"
+
+ def test_the_regex_still_gates_what_a_mark_can_keep(self) -> None:
+ """A mark cannot keep an op the regex does not make eligible."""
+ decisions, _ = _run(mark_at=1, compiled=False, save_ops=("nothing_matches_this",))
+ assert _saved_names(decisions) == [], _saved_names(decisions)
+
+ def test_without_save_only_marked_ops_every_eligible_op_is_kept(self) -> None:
+ """The pre-existing behaviour, unchanged."""
+ decisions, _ = _run(mark_at=None, compiled=False, save_only_marked_ops=False)
+ assert _saved_attention(decisions)[:3] == [True, True, True], _saved_names(decisions)
+
+ def test_the_marker_does_not_leak_between_regions(self) -> None:
+ """Policy state is per region, so a mark cannot arm the next block's first op."""
+ torch._dynamo.reset()
+ enable_marking()
+ decisions: list[tuple[str, bool]] = []
+ blocks = [
+ ptd_checkpoint_wrapper(
+ _Block(2 if index == 0 else None).cuda(),
+ context_fn=lambda: create_selective_checkpoint_contexts(_policy(decisions)),
+ )
+ for index in range(2)
+ ]
+ x = torch.randn(TOKENS, HIDDEN, device="cuda", dtype=torch.bfloat16, requires_grad=True)
+ nn.Sequential(*blocks)(x).float().sum().backward()
+ torch.cuda.synchronize()
+ # First block marks its last call; the second block marks nothing.
+ assert _saved_attention(decisions)[:6] == [False, False, True, False, False, False], _saved_names(decisions)
+
+ def test_the_clone_is_transient(self) -> None:
+ """The copy the marker makes must not survive the step.
+
+ It is MUST_RECOMPUTE, so a checkpointed region rebuilds it during recompute
+ rather than storing it.
+ """
+ torch._dynamo.reset()
+ enable_marking()
+ decisions: list[tuple[str, bool]] = []
+ block = ptd_checkpoint_wrapper(
+ _Block(0).cuda(),
+ context_fn=lambda: create_selective_checkpoint_contexts(_policy(decisions)),
+ )
+ x = torch.randn(TOKENS, HIDDEN, device="cuda", dtype=torch.bfloat16, requires_grad=True)
+ block(x).float().sum().backward()
+ x.grad = None
+ torch.cuda.synchronize()
+ gc.collect()
+ torch.cuda.empty_cache()
+ baseline = torch.cuda.memory_allocated()
+
+ resident = []
+ for _ in range(3):
+ block(x).float().sum().backward()
+ x.grad = None
+ torch.cuda.synchronize()
+ gc.collect()
+ resident.append(torch.cuda.memory_allocated() - baseline)
+ assert len(set(resident)) == 1, f"memory accumulated across steps: {resident}"
+
+ def test_the_op_is_registered_under_its_advertised_name(self) -> None:
+ namespace, name = MARK_OP_QUALNAME.split("::")
+ assert hasattr(getattr(torch.ops, namespace), name)
+ assert is_mark_op(f"{name}.default")
+ assert not is_mark_op("blackwell_fmha_forward.default")
+
+ def test_marking_returns_an_equal_tensor(self) -> None:
+ """A copy, because inductor rejects a custom op whose output aliases its input."""
+ enable_marking()
+ source = torch.randn(8, 16, device="cuda", dtype=torch.bfloat16)
+ marked = mark_next_activation(source)
+ assert torch.equal(marked, source)
+ assert marked.data_ptr() != source.data_ptr()
+
+ def test_marking_is_free_until_a_policy_asks_for_it(self) -> None:
+ """The copy is the whole cost of a marked call site, so an unmarked model pays none.
+
+ Every model reaching ``multiview_maskless_attention`` runs its mark line, but
+ only the configs setting ``save_only_marked_ops`` can do anything with a mark.
+ Off, the marker hands its argument straight back -- the same tensor, not a copy
+ of it -- so the line costs nothing rather than a clone per layer per step.
+ """
+ assert marking_enabled() is False, "the fixture should leave this off"
+ source = torch.randn(8, 16, device="cuda", dtype=torch.bfloat16)
+ assert mark_next_activation(source) is source
+
+ enable_marking()
+ assert marking_enabled() is True
+ assert mark_next_activation(source) is not source
+
+ @pytest.mark.parametrize(
+ "mode,save_only_marked_ops,expected",
+ [
+ pytest.param("selective", True, True, id="selective_marked"),
+ pytest.param("selective", False, False, id="selective_plain"),
+ pytest.param("full", True, False, id="full"),
+ ],
+ )
+ def test_wrapping_a_module_turns_marking_on_from_its_config(
+ self, mode: str, save_only_marked_ops: bool, expected: bool
+ ) -> None:
+ """The switch and the policy come from one config field, so they cannot disagree.
+
+ A policy that consults marks while the call sites emit none keeps *nothing* --
+ a silent memory regression rather than an error -- so the two are set together
+ where the module is wrapped rather than left to a caller to keep in step.
+
+ ``mode="full"`` recomputes the whole block and has no policy to consult, so the
+ flag buys nothing there and the clone is not worth paying for. Driven through
+ ``apply_ac_to_module``, the entry point the model actually calls, so the
+ dispatch is covered along with the switch.
+
+ Each case starts from off, which is what the fixture guarantees. What happens
+ when it starts from *on* is a separate question with a separate answer, in
+ ``test_a_later_unmarked_model_cannot_turn_marking_off``.
+ """
+ from cosmos_framework.configs.base.defaults.activation_checkpointing import (
+ ActivationCheckpointingConfig,
+ )
+ from cosmos_framework.model.generator.mot.parallelize_unified_mot import apply_ac_to_module
+
+ config = ActivationCheckpointingConfig(mode=mode, save_only_marked_ops=save_only_marked_ops)
+ apply_ac_to_module(nn.Identity(), config)
+ assert marking_enabled() is expected
+
+ def test_a_later_unmarked_model_cannot_turn_marking_off(self) -> None:
+ """The switch latches, and that is the decision rather than an oversight.
+
+ One process can build several networks -- an EMA copy, a distillation teacher --
+ and a second one wrapped without marks must not disarm the first. Dynamo guards
+ on the switch, so turning it off invalidates the first model's compiled code on
+ its very next forward, and the recompiled graph has no marker in it; a policy
+ built with ``save_only_marked_ops`` then keeps nothing and the model silently
+ reverts to recomputing every fold.
+
+ The cost of latching is the opposite mistake: a later maskless model pays a clone
+ it cannot use. That one is bounded -- the marker is ``MUST_RECOMPUTE``, so it is
+ transient rather than resident -- and it is the cheaper of the two by more than
+ an order of magnitude.
+ """
+ from cosmos_framework.configs.base.defaults.activation_checkpointing import (
+ ActivationCheckpointingConfig,
+ )
+ from cosmos_framework.model.generator.mot.parallelize_unified_mot import apply_ac_to_module
+
+ apply_ac_to_module(nn.Identity(), ActivationCheckpointingConfig(mode="selective", save_only_marked_ops=True))
+ assert marking_enabled() is True
+
+ for config in (
+ ActivationCheckpointingConfig(mode="selective"),
+ ActivationCheckpointingConfig(mode="full"),
+ ):
+ apply_ac_to_module(nn.Identity(), config)
+ assert marking_enabled() is True, f"{config.mode} disarmed a model that needs marks"
+
+ # And no production caller can reach the off path: it takes no argument.
+ import inspect
+
+ assert not inspect.signature(enable_marking).parameters, "enable_marking must not take a value"
diff --git a/cosmos_framework/model/generator/mot/attention_test.py b/cosmos_framework/model/generator/mot/attention_test.py
index dc33f9b51..9a23bb8bd 100644
--- a/cosmos_framework/model/generator/mot/attention_test.py
+++ b/cosmos_framework/model/generator/mot/attention_test.py
@@ -1910,6 +1910,44 @@ class _MultiviewMasklessBatch:
und_v: torch.Tensor # [N_und,kv_heads,head_dim]
+def _natten_is_the_varlen_backend() -> bool:
+ """Whether a varlen GQA call on this device dispatches NATTEN's FMHA.
+
+ The selective-AC policy matches ops by name, and ``["fmha"]`` only matches where
+ NATTEN is the *selected* backend: on Blackwell cuDNN and flash2 both refuse varlen
+ and NATTEN wins, while on Hopper flash3 outranks it and nothing is named "fmha".
+ ``NATTEN_SUPPORTED`` says only that it is importable, which is true on both, so it
+ does not gate these -- the H200 CI runner is exactly where that distinction bit.
+ """
+ if not torch.cuda.is_available():
+ return False
+ try:
+ from cosmos_framework.model.attention.backends import choose_backend
+
+ return (
+ choose_backend(
+ query_shape=torch.Size((1, 128, 8, 64)),
+ key_shape=torch.Size((1, 128, 4, 64)),
+ value_shape=torch.Size((1, 128, 4, 64)),
+ dtype=torch.bfloat16,
+ device=torch.device("cuda"),
+ requires_grad=True,
+ is_causal=False,
+ causal_type=None,
+ is_varlen=True,
+ return_lse=True,
+ raise_error=False,
+ )
+ == "natten"
+ )
+ except Exception: # noqa: BLE001 - a gate that cannot answer should skip, not error.
+ return False
+
+
+_NATTEN_IS_VARLEN_BACKEND = _natten_is_the_varlen_backend()
+_NOT_NATTEN_REASON = "The op-name policy only matches where NATTEN is the selected varlen backend."
+
+
def _multiview_maskless_batch(
*,
und_len: int,
@@ -2631,6 +2669,456 @@ def forward(self) -> torch.Tensor:
)
+@pytest.mark.L0
+@pytest.mark.GPU
+@pytest.mark.skipif(not torch.cuda.is_available(), reason="The attention kernels require a GPU.")
+@pytest.mark.skipif(not NATTEN_SUPPORTED, reason="merge_attentions requires NATTEN.")
+@pytest.mark.parametrize("save_only_marked_ops", [False, True])
+def test_marking_keeps_one_maskless_fold_where_the_regex_keeps_all_four(save_only_marked_ops: bool) -> None:
+ """On the real folds, the mark separates calls ``["fmha"]`` cannot tell apart.
+
+ The three sensor passes and the causal one dispatch the same op, so the regex keeps
+ all four or none, while the same-view fold is the only one worth keeping -- ~96% of
+ forward attention time and ~94% of backward, the others cheap to recompute. That is
+ why ``multiview_maskless_attention`` marks that fold at its call site, and why the
+ cam+LiDAR arm configuring ``save_only_marked_ops`` gets selective AC back instead of
+ the ``mode="full"`` its siblings fell back to.
+
+ Run against the real folds rather than a stand-in because the mark has to survive
+ what this call site puts between it and the kernel: the gathers for Q and V are
+ evaluated after the mark line and the attention frontend reshapes internally, so
+ several ops dispatch in between. An earlier version that gave the mark to the next
+ op of any kind put it on a gather, and the fold it was meant for was recomputed
+ anyway.
+
+ Counted as a ratio because the compiled region is traced more than once and the
+ policy object is shared across those traces, so the absolute count is a multiple of
+ the four folds rather than four.
+ """
+ from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import (
+ checkpoint_wrapper as ptd_checkpoint_wrapper,
+ )
+ from torch.utils.checkpoint import CheckpointPolicy, create_selective_checkpoint_contexts
+
+ from cosmos_framework.configs.base.defaults.activation_checkpointing import (
+ ATTENTION_FORWARD_OPS_REGEX,
+ )
+ from cosmos_framework.model.generator.mot.activation_marks import enable_marking, reset_marking_for_tests
+ from cosmos_framework.model.generator.mot.parallelize_unified_mot import make_selective_ac_policy
+
+ device = torch.device("cuda")
+ num_views, frames_per_view, patch_h, patch_w = 3, 2, 2, 3
+ token_shape = (num_views * frames_per_view, patch_h, patch_w)
+
+ # What ``_apply_selective_ac`` does at setup, and it has to happen before anything
+ # traces: ``context_fn`` runs after the region is traced, so flipping it there would
+ # leave the marker a pass-through in the compiled graph.
+ if save_only_marked_ops:
+ enable_marking()
+
+ # The shipped policy, wrapped only to record its verdicts. A reimplementation here
+ # would be free to agree with a policy production no longer runs.
+ eligible = [re.compile(pattern) for pattern in ATTENTION_FORWARD_OPS_REGEX]
+ real_policy = make_selective_ac_policy(eligible, save_only_marked_ops=save_only_marked_ops)
+
+ def _is_attention(op_name: str) -> bool:
+ # The same test the policy applies. Filtering on "fmha" here instead would make
+ # the assertions NATTEN-only while the policy is not, so on an arch where flash3
+ # serves these folds every one of them would be invisible and the test would fail
+ # for having found nothing rather than for anything about marking.
+ return any(pattern.search(op_name) for pattern in eligible)
+
+ verdicts: list[tuple[str, bool]] = []
+
+ def _policy(ctx, func, *args, **kwargs):
+ verdict = real_policy(ctx, func, *args, **kwargs)
+ verdicts.append((getattr(func, "__name__", str(func)), verdict == CheckpointPolicy.MUST_SAVE))
+ return verdict
+
+ batch = _multiview_maskless_batch(
+ und_len=5,
+ token_shape=token_shape,
+ num_views=num_views,
+ num_q_heads=8,
+ num_kv_heads=4,
+ head_dim=64,
+ device=device,
+ seed=0,
+ )
+ for pack in batch.packs:
+ for key in ("causal_seq", "full_only_seq"):
+ pack[key].requires_grad_(True)
+ plan = _plan(num_views, token_shape, device, _padded_gen_tokens(batch.packs[0]))
+ num_gen_tokens = batch.gen_q.shape[0]
+
+ class _Layer(torch.nn.Module):
+ def forward(self) -> torch.Tensor:
+ return get_gen_seq(multiview_attention(*batch.packs, maskless_plan=plan))[:num_gen_tokens]
+
+ layer = ptd_checkpoint_wrapper(
+ _Layer(),
+ context_fn=lambda: create_selective_checkpoint_contexts(_policy),
+ preserve_rng_state=False,
+ )
+ torch.manual_seed(1)
+ seed_grad = torch.randn(num_gen_tokens, 8 * 64, device=device, dtype=torch.bfloat16)
+ try:
+ # Compiled for the neighbouring test's reason: eager + non-reentrant checkpointing
+ # raises CheckpointError on this path, since NATTEN's merge backward unpacks thrice.
+ torch.compile(layer)().backward(seed_grad)
+ finally:
+ reset_marking_for_tests()
+
+ folds = [saved for name, saved in verdicts if _is_attention(name)]
+ assert folds and len(folds) % 4 == 0, f"expected whole passes over the four folds, saw {len(folds)}"
+ # One mark per pass. A second marked call site would double what each layer keeps,
+ # and every other assertion here would still hold.
+ marks = [name for name, _ in verdicts if "keep_next_activation" in name]
+ expected_marks = len(folds) // 4 if save_only_marked_ops else 0
+ assert len(marks) == expected_marks, f"expected {expected_marks} marks, saw {len(marks)}"
+ assert not any(saved for name, saved in verdicts if not _is_attention(name)), (
+ f"a non-attention op was kept: {sorted({name for name, saved in verdicts if saved and not _is_attention(name)})}"
+ )
+ saved_per_pass = 4 * sum(folds) / len(folds)
+ if save_only_marked_ops:
+ assert saved_per_pass == 1, f"expected the marked fold alone, kept {saved_per_pass} of 4"
+ else:
+ # The behaviour every existing config relies on: the regex decides alone, and it
+ # matches all four folds.
+ assert saved_per_pass == 4, f"the regex should keep all four folds, kept {saved_per_pass}"
+
+
+@pytest.mark.L0
+@pytest.mark.GPU
+@pytest.mark.skipif(not torch.cuda.is_available(), reason="The attention kernels require a GPU.")
+@pytest.mark.skipif(not NATTEN_SUPPORTED, reason="merge_attentions requires NATTEN.")
+def test_the_marked_maskless_fold_is_the_same_view_pass(monkeypatch: pytest.MonkeyPatch) -> None:
+ """Which of the four folds is marked, not merely that one of them is.
+
+ Keeping one fold only pays if it is the expensive one. The same-view pass is ~96%
+ of forward attention time and ~94% of backward; the other three are cheap to
+ recompute, so marking one of *those* would spend a clone and a stashed activation
+ to save almost nothing. Every other test here would still pass.
+
+ Checked by identity rather than by shape or position: three of the four folds run
+ on the same token count here, so nothing about the tensor says which pass it
+ belongs to.
+ """
+ from cosmos_framework.model.generator.mot import multiview_maskless_attention as maskless_module
+ from cosmos_framework.model.generator.mot.activation_marks import enable_marking, reset_marking_for_tests
+
+ device = torch.device("cuda")
+ num_views, frames_per_view, patch_h, patch_w = 3, 2, 2, 3
+ token_shape = (num_views * frames_per_view, patch_h, patch_w)
+
+ marked: list[torch.Tensor] = []
+ keys: list[torch.Tensor] = []
+ real_mark, real_attention = maskless_module.mark_next_activation, maskless_module.attention
+
+ def _spy_mark(tensor: torch.Tensor) -> torch.Tensor:
+ result = real_mark(tensor)
+ marked.append(result)
+ return result
+
+ def _spy_attention(query, key, value, **kwargs):
+ keys.append(key)
+ return real_attention(query, key, value, **kwargs)
+
+ monkeypatch.setattr(maskless_module, "mark_next_activation", _spy_mark)
+ monkeypatch.setattr(maskless_module, "attention", _spy_attention)
+
+ batch = _multiview_maskless_batch(
+ und_len=5,
+ token_shape=token_shape,
+ num_views=num_views,
+ num_q_heads=8,
+ num_kv_heads=4,
+ head_dim=64,
+ device=device,
+ seed=0,
+ )
+ plan = _plan(num_views, token_shape, device, _padded_gen_tokens(batch.packs[0]))
+ try:
+ enable_marking()
+ multiview_attention(*batch.packs, maskless_plan=plan)
+ finally:
+ reset_marking_for_tests()
+
+ assert len(marked) == 1, f"expected one marked call site, found {len(marked)}"
+ assert len(keys) > 1, "the decomposition should run several attention calls"
+ # Pass 1 of ``multiview_maskless_gen_attention`` is the same-view fold.
+ took_the_mark = [index for index, key in enumerate(keys) if key is marked[0]]
+ assert took_the_mark == [0], f"the mark reached folds {took_the_mark}, not the same-view pass"
+
+
+@pytest.mark.L0
+@pytest.mark.GPU
+@pytest.mark.skipif(not torch.cuda.is_available(), reason="The attention kernels require a GPU.")
+@pytest.mark.skipif(not NATTEN_SUPPORTED, reason="merge_attentions requires NATTEN.")
+def test_a_marked_maskless_layer_is_stable_across_steps() -> None:
+ """Training runs thousands of steps, and this one has to repeat itself.
+
+ Under compile the policy is consulted while tracing and not again, so a later step
+ can only be read off its results. Gradients that stay put say the partitioner keeps
+ making the same choice; flat memory says the kept fold and the marker's copy go with
+ the graph instead of accumulating -- the failure that would show up on a cluster as
+ a slow climb into OOM rather than as a wrong number.
+
+ Against a measured floor rather than bitwise: NATTEN's backward is not
+ deterministic, so two runs of one configuration already differ slightly.
+ """
+ import gc
+
+ from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import (
+ checkpoint_wrapper as ptd_checkpoint_wrapper,
+ )
+ from torch.utils.checkpoint import create_selective_checkpoint_contexts
+
+ from cosmos_framework.configs.base.defaults.activation_checkpointing import (
+ ATTENTION_FORWARD_OPS_REGEX,
+ )
+ from cosmos_framework.model.generator.mot.activation_marks import enable_marking, reset_marking_for_tests
+ from cosmos_framework.model.generator.mot.parallelize_unified_mot import make_selective_ac_policy
+
+ device = torch.device("cuda")
+ num_views, frames_per_view, patch_h, patch_w = 3, 2, 2, 3
+ token_shape = (num_views * frames_per_view, patch_h, patch_w)
+
+ batch = _multiview_maskless_batch(
+ und_len=5,
+ token_shape=token_shape,
+ num_views=num_views,
+ num_q_heads=8,
+ num_kv_heads=4,
+ head_dim=64,
+ device=device,
+ seed=0,
+ )
+ leaf = batch.packs[0]["full_only_seq"]
+ leaf.requires_grad_(True)
+ plan = _plan(num_views, token_shape, device, _padded_gen_tokens(batch.packs[0]))
+ num_gen_tokens = batch.gen_q.shape[0]
+
+ class _Layer(torch.nn.Module):
+ def forward(self) -> torch.Tensor:
+ return get_gen_seq(multiview_attention(*batch.packs, maskless_plan=plan))[:num_gen_tokens]
+
+ try:
+ enable_marking()
+ layer = torch.compile(
+ ptd_checkpoint_wrapper(
+ _Layer(),
+ context_fn=lambda: create_selective_checkpoint_contexts(
+ make_selective_ac_policy(
+ [re.compile(pattern) for pattern in ATTENTION_FORWARD_OPS_REGEX], save_only_marked_ops=True
+ )
+ ),
+ preserve_rng_state=False,
+ )
+ )
+ torch.manual_seed(1)
+ seed_grad = torch.randn(num_gen_tokens, 8 * 64, device=device, dtype=torch.bfloat16)
+
+ layer().backward(seed_grad)
+ first = leaf.grad.clone()
+ leaf.grad = None
+ layer().backward(seed_grad)
+ # The same step twice: what these differ by is the backward's own noise, and any
+ # later drift has to clear it to mean anything.
+ floor = torch.maximum((leaf.grad - first).abs().max() * 2, first.abs().max() * 2**-8)
+ leaf.grad = None
+ torch.cuda.synchronize()
+ gc.collect()
+ torch.cuda.empty_cache()
+ baseline = torch.cuda.memory_allocated()
+
+ resident = []
+ for step in range(3):
+ layer().backward(seed_grad)
+ drift = (leaf.grad - first).abs().max()
+ assert drift <= floor, f"step {step} drifted by {drift}, floor {floor}"
+ leaf.grad = None
+ torch.cuda.synchronize()
+ gc.collect()
+ resident.append(torch.cuda.memory_allocated() - baseline)
+ finally:
+ reset_marking_for_tests()
+
+ assert len(set(resident)) == 1, f"memory accumulated across steps: {resident}"
+
+
+@pytest.mark.skipif(not _NATTEN_IS_VARLEN_BACKEND, reason=_NOT_NATTEN_REASON)
+@pytest.mark.L0
+@pytest.mark.GPU
+@pytest.mark.skipif(not torch.cuda.is_available(), reason="The attention kernels require a GPU.")
+@pytest.mark.skipif(not NATTEN_SUPPORTED, reason="merge_attentions requires NATTEN.")
+def test_marking_costs_less_memory_than_keeping_every_maskless_fold() -> None:
+ """The reason the AV recipes can turn selective AC back on, measured on the real folds.
+
+ These recipes fell back to ``mode="full"`` because ``["fmha"]`` kept all four folds.
+ Keeping the marked one has to cost less than that, with the marker's copy inside the
+ measurement, or the config change makes things worse rather than better.
+ """
+ import gc
+
+ from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import (
+ checkpoint_wrapper as ptd_checkpoint_wrapper,
+ )
+ from torch.utils.checkpoint import create_selective_checkpoint_contexts
+
+ from cosmos_framework.configs.base.defaults.activation_checkpointing import (
+ ATTENTION_FORWARD_OPS_REGEX,
+ )
+ from cosmos_framework.model.generator.mot.activation_marks import enable_marking, reset_marking_for_tests
+ from cosmos_framework.model.generator.mot.parallelize_unified_mot import make_selective_ac_policy
+
+ device = torch.device("cuda")
+ num_views, frames_per_view, patch_h, patch_w = 3, 2, 2, 3
+ token_shape = (num_views * frames_per_view, patch_h, patch_w)
+
+ def _peak(save_only_marked_ops: bool) -> int:
+ torch._dynamo.reset()
+ if save_only_marked_ops:
+ enable_marking()
+ else:
+ reset_marking_for_tests()
+ batch = _multiview_maskless_batch(
+ und_len=5,
+ token_shape=token_shape,
+ num_views=num_views,
+ num_q_heads=8,
+ num_kv_heads=4,
+ head_dim=64,
+ device=device,
+ seed=0,
+ )
+ for pack in batch.packs:
+ for key in ("causal_seq", "full_only_seq"):
+ pack[key].requires_grad_(True)
+ plan = _plan(num_views, token_shape, device, _padded_gen_tokens(batch.packs[0]))
+ num_gen_tokens = batch.gen_q.shape[0]
+
+ class _Layer(torch.nn.Module):
+ def forward(self) -> torch.Tensor:
+ return get_gen_seq(multiview_attention(*batch.packs, maskless_plan=plan))[:num_gen_tokens]
+
+ layer = torch.compile(
+ ptd_checkpoint_wrapper(
+ _Layer(),
+ context_fn=lambda: create_selective_checkpoint_contexts(
+ make_selective_ac_policy(
+ [re.compile(pattern) for pattern in ATTENTION_FORWARD_OPS_REGEX],
+ save_only_marked_ops=save_only_marked_ops,
+ )
+ ),
+ preserve_rng_state=False,
+ )
+ )
+ torch.manual_seed(1)
+ seed_grad = torch.randn(num_gen_tokens, 8 * 64, device=device, dtype=torch.bfloat16)
+ gc.collect()
+ torch.cuda.empty_cache()
+ torch.cuda.reset_peak_memory_stats()
+ layer().backward(seed_grad)
+ torch.cuda.synchronize()
+ return torch.cuda.max_memory_allocated()
+
+ try:
+ marked = _peak(save_only_marked_ops=True)
+ keep_all = _peak(save_only_marked_ops=False)
+ finally:
+ reset_marking_for_tests()
+
+ assert marked < keep_all, f"marked peaked at {marked}, keeping every fold at {keep_all}"
+
+
+@pytest.mark.L0
+@pytest.mark.GPU
+@pytest.mark.skipif(not torch.cuda.is_available(), reason="The attention kernels require a GPU.")
+@pytest.mark.skipif(not NATTEN_SUPPORTED, reason="merge_attentions requires NATTEN.")
+def test_marking_a_maskless_fold_does_not_change_its_gradients() -> None:
+ """Marking decides what is kept, never what is computed.
+
+ Not implied by the unmarked checkpointing test above: the mark clones K before the
+ kernel, so the kernel saves the clone, and ``merge_attentions`` repairs a branch by
+ finding the storage its kernel saved *by data pointer*. An extra tensor between the
+ call site and the kernel is exactly the shape of thing that has broken that link
+ before, silently and by ~70%.
+
+ Exact against an uncheckpointed reference, for that test's reason: recompute reruns
+ the same kernels on the same inputs, so a difference is a missed patch rather than
+ arithmetic.
+ """
+ from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import (
+ checkpoint_wrapper as ptd_checkpoint_wrapper,
+ )
+ from torch.utils.checkpoint import create_selective_checkpoint_contexts
+
+ from cosmos_framework.configs.base.defaults.activation_checkpointing import (
+ ATTENTION_FORWARD_OPS_REGEX,
+ )
+ from cosmos_framework.model.generator.mot.activation_marks import enable_marking, reset_marking_for_tests
+ from cosmos_framework.model.generator.mot.parallelize_unified_mot import make_selective_ac_policy
+
+ device = torch.device("cuda")
+ num_views, frames_per_view, patch_h, patch_w = 3, 2, 2, 3
+ token_shape = (num_views * frames_per_view, patch_h, patch_w)
+
+ def _grads(checkpointed: bool) -> dict[str, torch.Tensor]:
+ batch = _multiview_maskless_batch(
+ und_len=5,
+ token_shape=token_shape,
+ num_views=num_views,
+ num_q_heads=8,
+ num_kv_heads=4,
+ head_dim=64,
+ device=device,
+ seed=0,
+ )
+ leaves: dict[str, torch.Tensor] = {}
+ for name, pack in zip("qkv", batch.packs):
+ for key in ("causal_seq", "full_only_seq"):
+ pack[key].requires_grad_(True)
+ leaves[f"{name}.{key}"] = pack[key]
+ plan = _plan(num_views, token_shape, device, _padded_gen_tokens(batch.packs[0]))
+ num_gen_tokens = batch.gen_q.shape[0]
+
+ class _Layer(torch.nn.Module):
+ def forward(self) -> torch.Tensor:
+ return get_gen_seq(multiview_attention(*batch.packs, maskless_plan=plan))[:num_gen_tokens]
+
+ layer: torch.nn.Module = _Layer()
+ # Set before the region is traced, as ``_apply_selective_ac`` does. The reference
+ # leaves it off, so it also confirms the marked graph and the plain one agree.
+ if checkpointed:
+ enable_marking()
+ policy = make_selective_ac_policy(
+ [re.compile(pattern) for pattern in ATTENTION_FORWARD_OPS_REGEX], save_only_marked_ops=True
+ )
+ layer = ptd_checkpoint_wrapper(
+ layer,
+ context_fn=lambda: create_selective_checkpoint_contexts(policy),
+ preserve_rng_state=False,
+ )
+ torch.manual_seed(1)
+ seed_grad = torch.randn(num_gen_tokens, 8 * 64, device=device, dtype=torch.bfloat16)
+ torch.compile(layer)().backward(seed_grad)
+ return {name: leaf.grad for name, leaf in leaves.items() if leaf.grad is not None}
+
+ try:
+ expected = _grads(checkpointed=False)
+ actual = _grads(checkpointed=True)
+ finally:
+ reset_marking_for_tests()
+
+ assert expected, "no leaf took a gradient, so this would pass vacuously"
+ assert set(expected) == set(actual), "marked checkpointing changed which leaves take a gradient"
+ for name, want in expected.items():
+ torch.testing.assert_close(
+ actual[name].double(), want.double(), atol=0, rtol=0, msg=lambda m, n=name: f"{n}: {m}"
+ )
+
+
@pytest.mark.L0
@pytest.mark.CPU
@torch.no_grad()
diff --git a/cosmos_framework/model/generator/mot/causal_attention.py b/cosmos_framework/model/generator/mot/causal_attention.py
index 3cffd2edc..b8d06a737 100644
--- a/cosmos_framework/model/generator/mot/causal_attention.py
+++ b/cosmos_framework/model/generator/mot/causal_attention.py
@@ -1842,9 +1842,10 @@ def attention_AR_gen_only(
und K/V is padded to ``S_und_max`` while the rolling gen history is
already fixed-size. The real und prefix is compacted with gen K/V so
varlen attention ignores only the padded text suffix.
- - **Static-shape** (``gen_k_buf_full`` set, ``gen_k_hist`` ``None``):
- cat ``[und || curr || gen_buf]`` (real positions contiguous from
- offset 0; padding in the gen-buf tail) and call ``attention()``
+ - **Static-shape** (``kv_k_static`` set, ``gen_k_hist`` ``None``):
+ write the current frame in place into the pooled ``[und | curr | hist |
+ pad]`` buffer (real positions contiguous from offset 0; padding in the
+ tail) and call ``attention()`` on that buffer
with the varlen kwargs ``cumulative_seqlen_Q`` /
``cumulative_seqlen_KV`` pre-built outside the captured region in
``ARMemoryState.init`` (see ``ARMemoryValue.cu_seqlens_q_t`` /
@@ -1875,9 +1876,9 @@ def attention_AR_gen_only(
k_gen = get_gen_seq(packed_key_states) # [S_curr, H_kv, D]
v_gen = get_gen_seq(packed_value_states) # [S_curr, H_kv, D]
- if memory_value.batch_size > 1:
- if memory_value.for_cuda_graphs or memory_value.post_saturation_static_compile:
- raise ValueError("Batched AR attention supports only the eager dynamic-shape path")
+ if memory_value.batch_size > 1 and not memory_value.for_cuda_graphs:
+ if memory_value.post_saturation_static_compile:
+ raise ValueError("Batched AR attention does not support post-saturation static compile")
if len(memory_value.gen_lens) != memory_value.batch_size:
raise ValueError(f"Expected {memory_value.batch_size} generation lengths, got {memory_value.gen_lens}")
if len(memory_value.und_lens) != memory_value.batch_size:
@@ -1968,24 +1969,75 @@ def attention_AR_gen_only(
v_curr = v_gen_real.unsqueeze(0) # [1, S_gen_real, H_kv, D]
if memory_value.for_cuda_graphs:
- # Static-shape branch. Real positions live in [0, S_und + gen_len +
- # real_gen_cache_len); the gen-buffer tail is zero-padding to a
- # fixed max size. Putting the current frame *before* the gen
- # buffer keeps real positions contiguous from offset 0, so a
- # single ``cumulative_seqlen_KV = [0, real_total_kv_len]``
- # restricts the kernel to the real prefix without any padding
- # hole. RoPE was applied to each K vector at projection time,
- # so order within the seq dim is irrelevant for correctness.
- assert memory_value.und_k_cached is not None and memory_value.und_v_cached is not None, (
- "static-shape branch requires the und cache to be populated"
- )
- assert memory_value.gen_k_buf_full is not None
- assert memory_value.gen_v_buf_full is not None
+ # Static-shape branch. ``kv_k_static`` / ``kv_v_static`` are one pooled
+ # ``[und | curr | hist | pad]`` buffer per K and V: real positions live in
+ # [0, S_und + gen_len + real_gen_cache_len) and the tail is padding to a
+ # fixed max size, so a single ``cumulative_seqlen_KV = [0,
+ # real_total_kv_len]`` restricts the kernel to the real prefix without any
+ # padding hole. The und region was primed once per generation and the
+ # history region rebuilt outside the compiled region; only the current
+ # frame is written here, in place, instead of materialising the whole
+ # sequence with ``cat`` on every forward. RoPE was applied to each K
+ # vector at projection time, so order within the seq dim is irrelevant
+ # for correctness.
assert memory_value.cu_seqlens_q_t is not None
assert memory_value.cu_seqlens_kv_t is not None
+ if memory_value.kv_k_static is None:
+ # Legacy static layout: separate ``und_k_cached`` / ``gen_k_buf_full`` buffers
+ # (real history prefix + padded tail) materialised as ``[und | curr |
+ # gen_buf_full]`` per forward. Kept for callers that build ``ARMemoryValue``
+ # by hand; the production state hands the block the composite buffer below.
+ assert memory_value.und_k_cached is not None and memory_value.und_v_cached is not None, (
+ "static-shape branch requires the und cache to be populated"
+ )
+ assert memory_value.gen_k_buf_full is not None and memory_value.gen_v_buf_full is not None, (
+ "static-shape branch requires either kv_k_static or gen_k_buf_full"
+ )
+ k_legacy = torch.cat([memory_value.und_k_cached, k_curr, memory_value.gen_k_buf_full], dim=1)
+ v_legacy = torch.cat([memory_value.und_v_cached, v_curr, memory_value.gen_v_buf_full], dim=1)
+ attn_result = attention(
+ query=q_gen.unsqueeze(0), # [1, S_curr, H, D]
+ key=k_legacy, # [1, KV_LEN_MAX, H_kv, D]
+ value=v_legacy,
+ cumulative_seqlen_Q=memory_value.cu_seqlens_q_t,
+ cumulative_seqlen_KV=memory_value.cu_seqlens_kv_t,
+ max_seqlen_Q=gen_len,
+ max_seqlen_KV=memory_value.max_seqlen_KV,
+ is_causal=False,
+ return_lse=False,
+ backend="natten",
+ )
+ assert isinstance(attn_result, torch.Tensor)
+ gen_out = attn_result.squeeze(0).flatten(-2, -1) # [S_curr, H*D]
+ output = from_und_gen_splits(gen_out.new_empty(0, gen_out.shape[-1]), gen_out, packed_query_states)
+ return output, None
+ assert memory_value.kv_v_static is not None, "static-shape branch requires both composite K and V buffers"
+ batch_rows = memory_value.batch_size
+ assert memory_value.kv_k_static.shape[1] == batch_rows * memory_value.max_seqlen_KV, (
+ f"static K/V buffer holds {memory_value.kv_k_static.shape[1]} tokens, "
+ f"expected {batch_rows} x max_seqlen_KV={memory_value.max_seqlen_KV}"
+ )
- k_full = torch.cat([memory_value.und_k_cached, k_curr, memory_value.gen_k_buf_full], dim=1)
- v_full = torch.cat([memory_value.und_v_cached, v_curr, memory_value.gen_v_buf_full], dim=1)
+ k_full = memory_value.kv_k_static # [1, B*R, H_kv, D]
+ v_full = memory_value.kv_v_static # [1, B*R, H_kv, D]
+ curr_start = memory_value.static_curr_offset
+ if batch_rows == 1:
+ # Literally the single-row write; kept as-is so the B=1 compiled graph is unchanged.
+ k_full[:, curr_start : curr_start + gen_len].copy_(k_curr)
+ v_full[:, curr_start : curr_start + gen_len].copy_(v_curr)
+ else:
+ # Row r's current frame sits at r*R + curr_start; one strided copy for all rows. The
+ # packed gen stream is sample-major with ``gen_len`` tokens per row.
+ row_stride = memory_value.static_row_stride
+ num_kv_heads, head_dim = k_gen.shape[-2], k_gen.shape[-1]
+ k_rows = k_gen[: batch_rows * gen_len].view(batch_rows, gen_len, num_kv_heads, head_dim) # [B,g,H_kv,D]
+ v_rows = v_gen[: batch_rows * gen_len].view(batch_rows, gen_len, num_kv_heads, head_dim) # [B,g,H_kv,D]
+ k_full.view(batch_rows, row_stride, num_kv_heads, head_dim)[:, curr_start : curr_start + gen_len].copy_(
+ k_rows
+ )
+ v_full.view(batch_rows, row_stride, num_kv_heads, head_dim)[:, curr_start : curr_start + gen_len].copy_(
+ v_rows
+ )
# ``cu_seqlens_q_t`` and ``cu_seqlens_kv_t`` are pre-built outside
# the captured region (in ``ARMemoryState.init``) as ``[2]`` int32
diff --git a/cosmos_framework/model/generator/mot/causal_context_parallel_test.py b/cosmos_framework/model/generator/mot/causal_context_parallel_test.py
index add2dabd5..f839b1bec 100644
--- a/cosmos_framework/model/generator/mot/causal_context_parallel_test.py
+++ b/cosmos_framework/model/generator/mot/causal_context_parallel_test.py
@@ -1056,16 +1056,25 @@ def _cp_attention_ar_gen_static(
cu_seqlens_kv_t = torch.tensor([0, real_total_kv_len], device=device, dtype=torch.int32)
real_gen_cache_len_t = torch.tensor([real_len], device=device, dtype=torch.int32)
max_seqlen_KV = s_und + gen_len + max_gen_cache_tokens
+ # Composite ``[und | curr | hist | pad]`` buffer; the attention writes ``curr`` in place.
+ curr_room_k = gen_k_buf.new_zeros((1, gen_len, gen_k_buf.shape[2], gen_k_buf.shape[3]))
+ curr_room_v = gen_v_buf.new_zeros((1, gen_len, gen_v_buf.shape[2], gen_v_buf.shape[3]))
+ kv_k_static = torch.cat([und_k_cached, curr_room_k, gen_k_buf], dim=1).contiguous()
+ kv_v_static = torch.cat([und_v_cached, curr_room_v, gen_v_buf], dim=1).contiguous()
mv = ARMemoryValue(
- und_k_cached=und_k_cached,
- und_v_cached=und_v_cached,
+ und_k_cached=None,
+ und_v_cached=None,
gen_k_hist=None,
gen_v_hist=None,
frame_idx=frame_idx,
gen_len=gen_len,
- gen_k_buf_full=gen_k_buf,
- gen_v_buf_full=gen_v_buf,
+ gen_k_buf_full=None,
+ gen_v_buf_full=None,
+ kv_k_static=kv_k_static,
+ kv_v_static=kv_v_static,
+ static_curr_offset=s_und,
+ static_hist_offset=s_und + gen_len,
real_gen_cache_len_t=real_gen_cache_len_t,
cu_seqlens_q_t=cu_seqlens_q_t,
cu_seqlens_kv_t=cu_seqlens_kv_t,
diff --git a/cosmos_framework/model/generator/mot/multiview_maskless_attention.py b/cosmos_framework/model/generator/mot/multiview_maskless_attention.py
index 974d6a7f0..a118661a1 100644
--- a/cosmos_framework/model/generator/mot/multiview_maskless_attention.py
+++ b/cosmos_framework/model/generator/mot/multiview_maskless_attention.py
@@ -28,6 +28,7 @@
MultiviewAttentionConfig,
resolve_caption_scope,
)
+from cosmos_framework.model.generator.mot.activation_marks import mark_next_activation
from cosmos_framework.model.generator.mot.merge_bridge import BridgeFn, MergeAttentionsBridge
from cosmos_framework.data.generator.sequence_packing.runtime import (
SequencePack,
@@ -813,9 +814,16 @@ def multiview_maskless_gen_attention(
# and the kernel's own output is that order -- no gather, no bridge. A control item puts a
# view's tokens in two runs instead, which costs the gather and the bridge back.
view_gather = plan.same_view_gather
+ # Keep this fold's output under selective AC rather than recomputing it: it is
+ # ~96% of forward attention time and ~94% of backward, against three other
+ # calls running the same kernel that a name-matching policy cannot tell apart.
+ # The mark goes on K because it is the smallest operand the call takes -- 32 query
+ # heads against 8 KV heads, 2 against 1 per rank under CP16 -- and marking copies
+ # what it marks.
+ same_view_k = mark_next_activation((k if view_gather is None else k[view_gather]).unsqueeze(0))
same_view_out, same_view_lse = attention(
(q if view_gather is None else q[view_gather]).unsqueeze(0), # [1,N_gen,heads,head_dim]
- (k if view_gather is None else k[view_gather]).unsqueeze(0), # [1,N_gen,kv_heads,head_dim]
+ same_view_k, # [1,N_gen,kv_heads,head_dim]
(v if view_gather is None else v[view_gather]).unsqueeze(0), # [1,N_gen,kv_heads,head_dim]
cumulative_seqlen_Q=plan.same_view_offsets,
cumulative_seqlen_KV=plan.same_view_offsets,
diff --git a/cosmos_framework/model/generator/mot/parallelize_unified_mot.py b/cosmos_framework/model/generator/mot/parallelize_unified_mot.py
index 95257c14c..d8f813b91 100644
--- a/cosmos_framework/model/generator/mot/parallelize_unified_mot.py
+++ b/cosmos_framework/model/generator/mot/parallelize_unified_mot.py
@@ -34,6 +34,7 @@
from cosmos_framework.configs.base.defaults.activation_checkpointing import ActivationCheckpointingConfig
from cosmos_framework.configs.base.defaults.compile import CompileConfig
+from cosmos_framework.model.generator.mot.activation_marks import enable_marking, is_mark_op
from cosmos_framework.model.generator.mot.attention import SplitInfo, dispatch_attention
from cosmos_framework.model.generator.mot.context_parallel_utils import context_parallel_attention
from cosmos_framework.model.generator.utils.memory import KVToStore, MemoryValue
@@ -127,25 +128,60 @@ def forward(
)
+def make_selective_ac_policy(save_ops_regex: "list[re.Pattern]", save_only_marked_ops: bool = False):
+ """The policy selective AC consults for every op inside a checkpointed block.
+
+ ``save_ops_regex`` says which ops are eligible to be kept. With
+ ``save_only_marked_ops``, an eligible op is kept only where the model asked for it
+ with :func:`~cosmos_framework.model.generator.mot.activation_marks.mark_next_activation`
+ -- which is how four calls running the same kernel can be told apart, since a
+ name cannot tell them apart. Without it every eligible op is kept, which is
+ what every existing config does.
+
+ A pending mark is consumed by the next *eligible* op rather than the next op.
+ At the real call site, the gathers for Q and V and the clone made by the attention
+ frontend all dispatch between the marker and the kernel; consuming the mark on
+ the next op of any kind puts the mark on a gather, not the kernel.
+ """
+ # Per region: ``context_fn`` is called once per checkpointed forward, so a
+ # mark cannot leak from one block into the next.
+ marked = {"armed": False}
+
+ def wrapped_policy(ctx, func, *args, **kwargs) -> CheckpointPolicy:
+ op_name = getattr(func, "__name__", str(func))
+ if is_mark_op(op_name):
+ marked["armed"] = True
+ # The marker is a copy, not something worth keeping.
+ return CheckpointPolicy.MUST_RECOMPUTE
+ if not any(pattern.search(op_name) for pattern in save_ops_regex):
+ return CheckpointPolicy.MUST_RECOMPUTE
+ if not save_only_marked_ops:
+ return CheckpointPolicy.MUST_SAVE
+ if marked["armed"]:
+ marked["armed"] = False
+ return CheckpointPolicy.MUST_SAVE
+ return CheckpointPolicy.MUST_RECOMPUTE
+
+ return wrapped_policy
+
+
def _apply_selective_ac(
module: nn.Module,
ac: ActivationCheckpointingConfig,
) -> nn.Module:
"""Apply per-op selective activation checkpointing to ``module``."""
save_ops_regex = [re.compile(pattern) for pattern in ac.save_ops_regex]
-
- def _get_custom_policy():
- def wrapped_policy(ctx, func, *args, **kwargs) -> CheckpointPolicy:
- op_name = getattr(func, "__name__", str(func))
- if any(pattern.search(op_name) for pattern in save_ops_regex):
- return CheckpointPolicy.MUST_SAVE
- return CheckpointPolicy.MUST_RECOMPUTE
-
- return wrapped_policy
+ if ac.save_only_marked_ops:
+ # Here rather than at the call sites: this is the one place that knows the
+ # policy wants marks, and it runs before the first forward, so the marker is
+ # a pass-through for every model that does not ask for it.
+ enable_marking()
return ptd_checkpoint_wrapper(
module,
- context_fn=lambda: create_selective_checkpoint_contexts(_get_custom_policy()),
+ context_fn=lambda: create_selective_checkpoint_contexts(
+ make_selective_ac_policy(save_ops_regex, save_only_marked_ops=ac.save_only_marked_ops)
+ ),
preserve_rng_state=ac.preserve_rng_state,
determinism_check=ac.determinism_check,
)
@@ -358,12 +394,21 @@ def apply_compile(model: nn.Module, config: CompileConfig) -> None:
if config.coordinate_descent_tuning:
compile_options["coordinate_descent_tuning"] = True
+ # A whole-forward CUDA-graph capture replays the compiled block's wrapper as recorded GPU work
+ # only. With symbolic shapes inductor materialises Python floats (e.g. RMSNorm epsilons) and
+ # sizes through CPU code plus a pinned host->device copy inside that wrapper; on replay the
+ # CPU side never runs and the copy reads freed host memory. Static shapes bake those values
+ # into the kernels, so forward scope forces ``dynamic=False`` (as the post-saturation static
+ # compile already does).
+ forward_scope_capture = config.use_cuda_graphs and config.cuda_graph_scope == "forward"
for layer_id, block in model.model.layers.named_children():
block = torch.compile(
block,
fullgraph=True,
- dynamic=config.compile_dynamic,
- mode="reduce-overhead" if config.use_cuda_graphs else None,
+ dynamic=False if forward_scope_capture else config.compile_dynamic,
+ # Forward-scope capture wraps the whole AR forward in one explicit graph; inductor's
+ # per-block CUDA-graph trees cannot nest inside it, so they are only used for "block".
+ mode="reduce-overhead" if (config.use_cuda_graphs and config.cuda_graph_scope == "block") else None,
options=compile_options or None,
)
# Instance-attribute override, not a subclass/wrapper module: OptimizedModule already
diff --git a/cosmos_framework/model/generator/mot/parallelize_unified_mot_test.py b/cosmos_framework/model/generator/mot/parallelize_unified_mot_test.py
index 482ab0fb5..69a66d674 100644
--- a/cosmos_framework/model/generator/mot/parallelize_unified_mot_test.py
+++ b/cosmos_framework/model/generator/mot/parallelize_unified_mot_test.py
@@ -389,3 +389,62 @@ def test_compile_runs_before_fsdp_wrapping(self) -> None:
parallelize_unified_mot(model, parallel_dims, compile_config, SimpleNamespace())
assert calls == ["ac", "compile", "fsdp"]
+
+
+class TestCudaGraphScope:
+ """``cuda_graph_scope`` decides whether inductor's CUDA-graph trees wrap each block.
+
+ With scope "forward" the AR loop captures one explicit graph around the whole forward, and a
+ capture cannot nest the per-block graph trees, so the blocks must compile with ``mode=None``.
+ """
+
+ @staticmethod
+ def _model_with_one_layer() -> torch.nn.Module:
+ block = torch.nn.Identity()
+ layers = torch.nn.Module()
+ layers.register_module("0", block)
+ inner = torch.nn.Module()
+ inner.layers = layers
+ model = torch.nn.Module()
+ model.model = inner
+ return model
+
+ @pytest.mark.parametrize(
+ ("use_cuda_graphs", "scope", "expected_mode"),
+ [(True, "block", "reduce-overhead"), (True, "forward", None), (False, "block", None), (False, "forward", None)],
+ )
+ def test_block_scope_alone_uses_reduce_overhead(self, use_cuda_graphs: bool, scope: str, expected_mode) -> None:
+ model = self._model_with_one_layer()
+ seen_modes: list[object] = []
+
+ def fake_compile(block, **kwargs):
+ seen_modes.append(kwargs["mode"])
+ return block
+
+ with patch("cosmos_framework.model.generator.mot.parallelize_unified_mot.torch.compile", fake_compile):
+ apply_compile(model, CompileConfig(enabled=True, use_cuda_graphs=use_cuda_graphs, cuda_graph_scope=scope))
+
+ assert seen_modes == [expected_mode]
+
+ @pytest.mark.parametrize("scope", ["block", "forward"])
+ def test_forward_scope_forces_static_shapes(self, scope: str) -> None:
+ model = self._model_with_one_layer()
+ seen_dynamic: list[object] = []
+
+ def fake_compile(block, **kwargs):
+ seen_dynamic.append(kwargs["dynamic"])
+ return block
+
+ with patch("cosmos_framework.model.generator.mot.parallelize_unified_mot.torch.compile", fake_compile):
+ apply_compile(
+ model, CompileConfig(enabled=True, use_cuda_graphs=True, cuda_graph_scope=scope, compile_dynamic=True)
+ )
+
+ # Symbolic shapes stage Python floats through pinned host memory inside the compiled wrapper,
+ # which a whole-forward capture would replay from freed memory; forward scope pins static shapes.
+ assert seen_dynamic == [scope == "block"]
+
+ def test_scope_is_validated_and_defaults_to_block(self) -> None:
+ assert CompileConfig().cuda_graph_scope == "block"
+ with pytest.raises(ValueError):
+ CompileConfig(cuda_graph_scope="whole") # type: ignore[arg-type]
diff --git a/cosmos_framework/model/generator/mot/parallelize_vfm_network.py b/cosmos_framework/model/generator/mot/parallelize_vfm_network.py
index 06ee39df0..1df56dbbc 100644
--- a/cosmos_framework/model/generator/mot/parallelize_vfm_network.py
+++ b/cosmos_framework/model/generator/mot/parallelize_vfm_network.py
@@ -30,10 +30,16 @@ def apply_compile(model: torch.nn.Module, config: CompileConfig):
if config.coordinate_descent_tuning:
inductor_options["coordinate_descent_tuning"] = True
+ # Whole-forward CUDA graphs (``cuda_graph_scope="forward"``) capture these heads inside the
+ # recorded AR forward: per-region CUDA-graph trees cannot replay inside that capture, and a
+ # symbolic-shape wrapper would stage host values through pinned memory that a replay later
+ # reads after it is freed. That scope therefore compiles the heads like the decoder blocks
+ # (inductor only, static shapes); every other configuration keeps the dynamic single graph.
+ forward_scope_capture = config.use_cuda_graphs and config.cuda_graph_scope == "forward"
compile_options = {
"fullgraph": True,
- "dynamic": True,
- "mode": "reduce-overhead" if config.use_cuda_graphs else None,
+ "dynamic": not forward_scope_capture,
+ "mode": "reduce-overhead" if (config.use_cuda_graphs and not forward_scope_capture) else None,
"options": inductor_options or None,
}
diff --git a/cosmos_framework/model/generator/mot/parallelize_vfm_network_test.py b/cosmos_framework/model/generator/mot/parallelize_vfm_network_test.py
index 3bc97aff6..55b9f0c50 100644
--- a/cosmos_framework/model/generator/mot/parallelize_vfm_network_test.py
+++ b/cosmos_framework/model/generator/mot/parallelize_vfm_network_test.py
@@ -13,12 +13,17 @@
runs in) alongside the memory it is there to buy.
"""
+from types import SimpleNamespace
+from unittest.mock import patch
+
import pytest
import torch
from torch import nn
from cosmos_framework.utils.helper_test import RunIf
from cosmos_framework.configs.base.defaults.activation_checkpointing import ActivationCheckpointingConfig
+from cosmos_framework.configs.base.defaults.compile import CompileConfig
+from cosmos_framework.model.generator.mot import parallelize_vfm_network as module
from cosmos_framework.model.generator.mot.modeling_utils import TimestepEmbedder
from cosmos_framework.model.generator.mot.parallelize_vfm_network import apply_ac
@@ -188,3 +193,36 @@ def retained(mode: str) -> int:
# The float32 [N, HIDDEN] intermediates inside the MLP (the first Linear's
# output and the SiLU's) are what stop being pinned; assert at least both.
assert saved >= 2 * one_activation, f"only freed {saved / 2**20:.0f} MiB"
+
+
+@pytest.mark.CPU
+@pytest.mark.parametrize(
+ "use_cuda_graphs,scope,expected_mode,expected_dynamic",
+ [
+ (False, "block", None, True),
+ (True, "block", "reduce-overhead", True),
+ (True, "forward", None, False),
+ ],
+)
+def test_vfm_heads_follow_the_cuda_graph_scope(
+ use_cuda_graphs: bool, scope: str, expected_mode: str | None, expected_dynamic: bool
+) -> None:
+ """Forward-scope capture records the heads inside the AR forward: no nested CUDA-graph trees,
+ static shapes (no host staging in the wrapper); other configurations keep the dynamic graph."""
+ heads = ["_encode_text", "_encode_vision", "_encode_action", "_decode_vision", "_decode_action"]
+ model = SimpleNamespace(**{name: object() for name in heads})
+ seen: list[dict] = []
+
+ def fake_compile(fn, **kwargs):
+ seen.append(kwargs)
+ return fn
+
+ with patch.object(module.torch, "compile", fake_compile):
+ module.apply_compile(
+ model, CompileConfig(enabled=True, use_cuda_graphs=use_cuda_graphs, cuda_graph_scope=scope)
+ )
+
+ assert len(seen) == len(heads)
+ assert all(kwargs["mode"] == expected_mode for kwargs in seen)
+ assert all(kwargs["dynamic"] is expected_dynamic for kwargs in seen)
+ assert all(kwargs["fullgraph"] is True for kwargs in seen)
diff --git a/cosmos_framework/model/generator/mot/post_saturation/cuda_graph.py b/cosmos_framework/model/generator/mot/post_saturation/cuda_graph.py
index 4a86907e1..e4893aedb 100644
--- a/cosmos_framework/model/generator/mot/post_saturation/cuda_graph.py
+++ b/cosmos_framework/model/generator/mot/post_saturation/cuda_graph.py
@@ -189,6 +189,18 @@ def __init__(self) -> None:
self._capture_stream: torch.cuda.Stream | None = None
self._graph_pool: Any | None = None
+ @staticmethod
+ def _base_kind(kind: str) -> str:
+ """Validate ``kind`` = ``"denoise" | "refresh"`` with an optional ``":"`` suffix.
+
+ The variant (e.g. the Transfer history limit of a seed) keys separate graphs whose
+ captured ``memory_info`` differs while the forward structure is the same.
+ """
+ base_kind = kind.split(":", 1)[0]
+ if base_kind not in {"denoise", "refresh"}:
+ raise ValueError(f"Unsupported post-saturation CUDA Graph kind={kind!r}")
+ return base_kind
+
def reset_for_new_generation(self) -> None:
"""Discard graphs bound to the previous generation's cache storage."""
had_captures = bool(self._runners)
@@ -208,8 +220,7 @@ def run(
memory_info: dict[str, Any],
) -> dict[str, Any]:
"""Capture on first use, then replay a branch-specific coarse graph."""
- if kind not in {"denoise", "refresh"}:
- raise ValueError(f"Unsupported post-saturation CUDA Graph kind={kind!r}")
+ self._base_kind(kind)
key = (kind, branch)
runner = self._runners.get(key)
if runner is None:
diff --git a/cosmos_framework/model/generator/mot/unified_mot.py b/cosmos_framework/model/generator/mot/unified_mot.py
index f17cc9a05..2c95692e0 100644
--- a/cosmos_framework/model/generator/mot/unified_mot.py
+++ b/cosmos_framework/model/generator/mot/unified_mot.py
@@ -97,49 +97,37 @@
def _pad_packed_tokens_by_sample(
tokens: torch.Tensor, # [N_padded,H,D]
- sample_ids: torch.Tensor, # [N_padded]
+ lengths: tuple[int, ...],
num_real_tokens: int,
- batch_size: int,
-) -> tuple[torch.Tensor, tuple[int, ...]]: # ([B,S_max,H,D], tuple[B])
+) -> torch.Tensor: # [B,S_max,H,D]
"""Restore the batch dimension needed to store packed K/V in the AR cache.
- Attention projects one flat, sample-major token stream, but the batched
- cache stores independent ``[B,S,H,D]`` rows. A leading ``unsqueeze(0)``
- would merge all samples into one history; a reshape cannot handle unequal
- prompt lengths. This helper copies each sample's tokens into its own row
- and zero-pads shorter rows to the longest sequence in this pathway.
-
- ``sample_ids`` are SequencePack's batch-local row ordinals, not persistent
- dataset or request IDs. The real prefix must already be grouped in
- increasing sample order, with IDs in ``[0,batch_size)``; this function
- counts boundaries but does not sort or gather interleaved samples. The
- caller must keep the same batch-row order while reusing the AR cache.
-
- Args:
- tokens: Packed K or V, ``[N_padded,H,D]``. Only the leading
- ``num_real_tokens`` entries are copied; alignment padding is ignored.
- sample_ids: Integer row ordinal for each token, ``[N_padded]``.
- num_real_tokens: Length of the real token prefix shared by both inputs.
- batch_size: Number of cache rows, including samples empty in this pathway.
-
- Returns:
- Zero-padded ``[B,S_max,H,D]`` tokens and the ``B`` real sequence lengths.
- Padding is storage only: attention must use real lengths, not attend
- to the zero tails. ``ARMemoryState`` obtains the same lengths from the
- pack and carries them alongside the cache.
+ Attention projects one flat, sample-major token stream, but the cache stores independent
+ ``[B,S,H,D]`` rows. ``lengths`` are the per-sample token counts of this pathway, taken from
+ the pack's host-side layout (``per_sample_pathway_lengths``), so no device->host sync happens
+ inside the compiled block. Equal lengths (every AR generation split, and ``B=1``) are a pure
+ view of the real prefix; unequal prompt lengths are copied into a zero-padded tensor whose
+ padding is storage only -- readers must use the real lengths.
"""
- real_sample_ids = sample_ids[:num_real_tokens] # [N]
- lengths_tensor = torch.bincount(real_sample_ids, minlength=batch_size) # [B]
- lengths = tuple(int(length) for length in lengths_tensor.tolist())
- max_length = max(lengths, default=0)
+ if sum(lengths) != num_real_tokens:
+ raise AssertionError(f"Packed token lengths {lengths} sum to {sum(lengths)}, expected {num_real_tokens}")
+ batch_size = len(lengths)
+ if batch_size == 1:
+ # The legacy single-sample path, kept literally: no shape arithmetic on lengths, which are
+ # symbolic under dynamic-shape compilation and may disagree with the (possibly empty)
+ # traced stream (e.g. the und stream of a gen-only frame).
+ return tokens[:num_real_tokens].unsqueeze(0) # [1,S,H,D]
+ # Traced under torch.compile with symbolic lengths (dynamic shapes): keep to builtins Dynamo
+ # handles on SymInts (no ``max(..., default=)``).
+ max_length = max(lengths) if lengths else 0
+ if all(length == max_length for length in lengths):
+ return tokens[:num_real_tokens].view(batch_size, max_length, *tokens.shape[1:]) # [B,S,H,D]
padded = tokens.new_zeros((batch_size, max_length, *tokens.shape[1:])) # [B,S_max,H,D]
offset = 0
for sample_idx, length in enumerate(lengths):
padded[sample_idx, :length] = tokens[offset : offset + length] # [S_i,H,D]
offset += length
- if offset != num_real_tokens:
- raise AssertionError(f"Packed token lengths sum to {offset}, expected {num_real_tokens}")
- return padded, lengths
+ return padded
# -----------------------------------------------------------------------------
@@ -848,34 +836,28 @@ def forward(
# of raw k_und_. Without the norm, k_und_for_gen_ is not defined, so
# fall back to k_und_.
k_und_to_store = k_und_for_gen_ if self.k_norm_und_for_gen is not None else k_und_
+ # Undo sample packing before cache writes: batch row i must keep sample i's history
+ # across AR steps. One construction for every batch size and no device->host sync
+ # inside the compiled block: AR packs give every sample the same generation split, so
+ # the gen rows are a view of the real prefix (B=1 included); prompt lengths come from
+ # the memory state, which read them on the host at frame 0. Unequal prompts are
+ # zero-padded and ARMemoryState excludes the padding on reads.
memory_batch_size = int(getattr(memory_value, "batch_size", 1))
- if memory_batch_size > 1:
- # Undo sample packing before cache writes: batch row i must keep
- # sample i's history across AR steps. UND prompt lengths may
- # differ, so pad each pathway independently; ARMemoryState uses
- # the pack's real per-row lengths to exclude padding on reads.
- gen_k_batched, gen_lengths = _pad_packed_tokens_by_sample(
- k_gen_, pack["_full_only_sample_ids"], gen_len, memory_batch_size
- ) # [B,S_gen,H,D], tuple[B]
- gen_v_batched, gen_v_lengths = _pad_packed_tokens_by_sample(
- v_gen, pack["_full_only_sample_ids"], gen_len, memory_batch_size
- ) # [B,S_gen,H,D], tuple[B]
- und_k_batched, und_lengths = _pad_packed_tokens_by_sample(
- k_und_to_store, pack["_causal_sample_ids"], und_len, memory_batch_size
- ) # [B,S_und_max,H,D], tuple[B]
- und_v_batched, und_v_lengths = _pad_packed_tokens_by_sample(
- v_und, pack["_causal_sample_ids"], und_len, memory_batch_size
- ) # [B,S_und_max,H,D], tuple[B]
- if gen_lengths != gen_v_lengths or und_lengths != und_v_lengths:
- raise AssertionError("Packed K/V sample lengths differ")
- kv_to_store = (gen_k_batched, gen_v_batched, und_k_batched, und_v_batched)
+ if gen_len % memory_batch_size:
+ raise AssertionError(f"{gen_len} generation tokens do not split evenly over {memory_batch_size} rows")
+ gen_lengths = (gen_len // memory_batch_size,) * memory_batch_size
+ if und_len == 0:
+ und_lengths = (0,) * memory_batch_size
+ elif memory_batch_size == 1:
+ und_lengths = (und_len,)
else:
- kv_to_store = (
- k_gen_[:gen_len].unsqueeze(0),
- v_gen[:gen_len].unsqueeze(0),
- k_und_to_store[:und_len].unsqueeze(0),
- v_und[:und_len].unsqueeze(0),
- )
+ und_lengths = tuple(memory_value.und_lens)
+ kv_to_store = (
+ _pad_packed_tokens_by_sample(k_gen_, gen_lengths, gen_len), # [B,S_gen,H,D]
+ _pad_packed_tokens_by_sample(v_gen, gen_lengths, gen_len), # [B,S_gen,H,D]
+ _pad_packed_tokens_by_sample(k_und_to_store, und_lengths, und_len), # [B,S_und_max,H,D]
+ _pad_packed_tokens_by_sample(v_und, und_lengths, und_len), # [B,S_und_max,H,D]
+ )
# Attention compute is local-head under both sequence-sharded and
# replicated attention I/O layouts. The difference here is the output
diff --git a/cosmos_framework/model/generator/mot/unified_mot_test.py b/cosmos_framework/model/generator/mot/unified_mot_test.py
index d90beac02..954c08d77 100644
--- a/cosmos_framework/model/generator/mot/unified_mot_test.py
+++ b/cosmos_framework/model/generator/mot/unified_mot_test.py
@@ -207,16 +207,9 @@ def forward(
def test_pad_packed_tokens_by_sample_supports_unequal_prompt_lengths() -> None:
"""Sample-major packed prompts become independent padded cache rows."""
tokens = torch.arange(8 * 2 * 3, dtype=torch.float32).reshape(8, 2, 3) # [S_padded,H,D]
- sample_ids = torch.tensor([0, 0, 1, 1, 1, 1, 0, 1]) # [S_padded]
- padded, lengths = _pad_packed_tokens_by_sample(
- tokens,
- sample_ids,
- num_real_tokens=6,
- batch_size=2,
- ) # [B,S_max,H,D], tuple[B]
+ padded = _pad_packed_tokens_by_sample(tokens, (2, 4), num_real_tokens=6) # [B,S_max,H,D]
- assert lengths == (2, 4)
assert padded.shape == (2, 4, 2, 3)
torch.testing.assert_close(padded[0, :2], tokens[:2])
assert torch.count_nonzero(padded[0, 2:]) == 0
@@ -228,18 +221,64 @@ def test_pad_packed_tokens_by_sample_supports_unequal_prompt_lengths() -> None:
def test_pad_packed_tokens_by_sample_preserves_single_sample_layout() -> None:
"""B=1 produces the same leading-token tensor as the legacy unsqueeze path."""
tokens = torch.arange(7 * 2 * 3, dtype=torch.float32).reshape(7, 2, 3) # [S_padded,H,D]
- sample_ids = torch.zeros(7, dtype=torch.long) # [S_padded]
-
- padded, lengths = _pad_packed_tokens_by_sample(
- tokens,
- sample_ids,
- num_real_tokens=5,
- batch_size=1,
- ) # [1,S_real,H,D], tuple[1]
+
+ padded = _pad_packed_tokens_by_sample(tokens, (5,), num_real_tokens=5) # [1,S_real,H,D]
legacy = tokens[:5].unsqueeze(0) # [1,S_real,H,D]
- assert lengths == (5,)
torch.testing.assert_close(padded, legacy)
+ assert padded.data_ptr() == tokens.data_ptr() # a view, like the legacy unsqueeze
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_pad_packed_tokens_by_sample_single_sample_is_the_legacy_unsqueeze() -> None:
+ """B=1 must not reshape by length: under dynamic shapes the length is symbolic and the traced
+ und stream of a gen-only frame is empty, where ``view(1, n, ...)`` fails while ``unsqueeze`` is exact."""
+ import torch._dynamo
+
+ empty = torch.empty(0, 2, 3) # [0,H,D] und stream of a gen-only frame
+ out = _pad_packed_tokens_by_sample(empty, (0,), num_real_tokens=0)
+ assert out.shape == (1, 0, 2, 3)
+
+ def split(tokens: torch.Tensor, und_len: torch.Tensor) -> torch.Tensor:
+ n = int(und_len.shape[0]) # symbolic under dynamic=True, like pack["_num_causal_tokens"]
+ return _pad_packed_tokens_by_sample(tokens[:n], (n,), n)
+
+ compiled = torch.compile(split, backend="eager", dynamic=True, fullgraph=True)
+ tokens = torch.arange(5 * 2 * 3, dtype=torch.float32).reshape(5, 2, 3)
+ for n in (5, 3, 0):
+ torch.testing.assert_close(compiled(tokens, torch.zeros(n)), tokens[:n].unsqueeze(0))
+ torch._dynamo.reset()
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_pad_packed_tokens_by_sample_traces_with_symbolic_lengths() -> None:
+ """Per-block CUDA graphs compile the block with dynamic shapes: the row split must trace with SymInt lengths."""
+ import torch._dynamo
+
+ def split_rows(tokens: torch.Tensor, rows: int) -> torch.Tensor:
+ gen_len = tokens.shape[0] - 3 # symbolic under dynamic=True (3 padding rows)
+ return _pad_packed_tokens_by_sample(tokens, (gen_len // rows,) * rows, gen_len)
+
+ compiled = torch.compile(split_rows, backend="eager", dynamic=True, fullgraph=True)
+ for n_real, rows in ((6, 1), (8, 2), (12, 3)):
+ tokens = torch.arange((n_real + 3) * 2 * 3, dtype=torch.float32).reshape(n_real + 3, 2, 3) # [S_padded,H,D]
+ torch.testing.assert_close(compiled(tokens, rows), tokens[:n_real].view(rows, n_real // rows, 2, 3))
+ torch._dynamo.reset()
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_pad_packed_tokens_by_sample_equal_lengths_are_a_view() -> None:
+ """Equal per-sample lengths (every AR generation split) need no copy."""
+ tokens = torch.arange(9 * 2 * 3, dtype=torch.float32).reshape(9, 2, 3) # [S_padded,H,D]
+ batched = _pad_packed_tokens_by_sample(tokens, (3, 3), num_real_tokens=6) # [2,3,H,D]
+ assert batched.shape == (2, 3, 2, 3)
+ assert batched.data_ptr() == tokens.data_ptr()
+ torch.testing.assert_close(batched[1], tokens[3:6])
+ with pytest.raises(AssertionError):
+ _pad_packed_tokens_by_sample(tokens, (3, 2), num_real_tokens=6)
def _kv(batch: int, seqlen: int, num_kv_heads: int = 2, head_dim: int = 4) -> torch.Tensor:
diff --git a/cosmos_framework/model/generator/omni_mot_causal_model.py b/cosmos_framework/model/generator/omni_mot_causal_model.py
index 6f4b9c01a..9b1f81501 100644
--- a/cosmos_framework/model/generator/omni_mot_causal_model.py
+++ b/cosmos_framework/model/generator/omni_mot_causal_model.py
@@ -15,6 +15,7 @@
import contextlib
import itertools
from collections.abc import Callable, Generator, Iterable, Sequence
+from dataclasses import dataclass
from typing import Any, Literal, cast
from unittest.mock import patch
@@ -27,11 +28,13 @@
import cosmos_framework.model.generator.omni_mot_model as omni_mot_model_module
from cosmos_framework.configs.base.defaults.model_config import OmniMoTModelConfig
+from cosmos_framework.data.generator.augmentors.text_tokenizer import TEXT_SYSTEM_PROMPT_KEY
from cosmos_framework.model.generator.omni_mot_model import OmniMoTModel, _broadcast_seed, _per_view_caption_groups
from cosmos_framework.model.generator.utils.data_and_condition import GenerationDataClean
from cosmos_framework.model.generator.utils.memory import MemoryState
from cosmos_framework.data.generator.sequence_packing import PackedSequence, build_sequence_plans_from_data_batch
from cosmos_framework.data.generator.sequence_packing.modality import compute_text_split_length
+from cosmos_framework.data.generator.sequence_packing.runtime import to_device_nonblocking
from cosmos_framework.configs.base.defaults.causal_flex_attention import CausalFlexAttentionConfig
from cosmos_framework.configs.base.defaults.replay_attention import (
TeacherForcingKVImplementation,
@@ -47,6 +50,7 @@
from cosmos_framework.model.generator.mot.causal_flex_attention import build_teacher_forcing_clean_target_token_indexes
from cosmos_framework.model.generator.mot.post_saturation.installer import install_ar_post_saturation_mode
from cosmos_framework.model.generator.mot.post_saturation.runtime import (
+ ARPostSaturationRuntime,
is_ar_post_saturation_cuda_graph_frame,
is_ar_post_saturation_static_compile_frame,
reset_ar_post_saturation_runtime_for_generation,
@@ -64,6 +68,7 @@
ARMemoryState,
DualKVCache,
FlexARMemoryState,
+ KVBufferPool,
TeacherForcingMemoryState,
)
from cosmos_framework.model.generator.utils.kv_storage_backend import validate_kv_cache_dtype
@@ -71,6 +76,7 @@
from cosmos_framework.data.generator.sequence_packing.autoregressive import (
pack_input_sequence_autoregressive,
pack_input_sequence_autoregressive_batch,
+ resolve_text_system_prompt,
)
from cosmos_framework.utils.generator.data_batch import condition_frame_indexes_vision_from_batch
@@ -414,6 +420,64 @@ def _validate_kv_cache_dtype_supports_cuda_graphs(kv_cache_dtype: str | None, cu
)
+def uses_ar_forward_cuda_graph(model: Any) -> bool:
+ """Whether AR frames on the static-shape path replay one explicit CUDA graph per whole forward.
+
+ ``compile.cuda_graph_scope == "forward"`` compiles the decoder blocks without inductor's
+ per-block CUDA-graph trees (see ``apply_compile``) and instead captures ``denoise`` for
+ cache index >= 1 through :class:`ARPostSaturationRuntime`; the machinery is the one the
+ post-saturation ``"cuda-graph"`` mode uses, so the two are mutually exclusive.
+ """
+ compile_config = model.config.compile
+ return bool(
+ compile_config.enabled
+ and compile_config.use_cuda_graphs
+ and compile_config.cuda_graph_scope == "forward"
+ and compile_config.ar_post_saturation_mode == "default"
+ )
+
+
+def get_ar_forward_cuda_graph_runtime(model: Any) -> ARPostSaturationRuntime:
+ """Return the model-owned coarse-graph runtime for ``cuda_graph_scope="forward"`` (created lazily)."""
+ runtime = getattr(model, "_ar_forward_cuda_graph_runtime", None)
+ if not isinstance(runtime, ARPostSaturationRuntime):
+ runtime = ARPostSaturationRuntime()
+ model._ar_forward_cuda_graph_runtime = runtime
+ return runtime
+
+
+def run_ar_forward_cuda_graph(
+ model: Any,
+ *,
+ kind: str,
+ branch: str,
+ packed_seq: PackedSequence,
+ memory_info: dict[str, Any],
+) -> dict[str, Any]:
+ """Capture on first use, then replay, one whole-forward CUDA graph keyed by ``(kind, branch)``."""
+ return get_ar_forward_cuda_graph_runtime(model).run_cuda_graph(
+ model=model,
+ kind=kind,
+ branch=branch,
+ packed_seq=packed_seq,
+ memory_info=memory_info,
+ )
+
+
+@dataclass(frozen=True)
+class StreamingTransferStep:
+ """One control chunk of the batched streaming Transfer loop, with optional per-row episode restarts.
+
+ ``reset_rows`` restart those rows at this step: their K/V history and mRoPE frame clock start
+ over and, when ``seeds`` is given, so does their noise stream. Prompts are kept. Only tensor
+ values change, so compiled kernels and captured CUDA graphs keep running.
+ """
+
+ control: torch.Tensor # [B,C,T,H,W]
+ reset_rows: tuple[int, ...] = ()
+ seeds: tuple[int, ...] | None = None # per reset row, same order as ``reset_rows``
+
+
def _validate_attention_sink_config(kv_cache_inference_size: int | None, attention_sink_size: int) -> None:
"""Validate AR-inference attention-sink cache settings."""
if attention_sink_size < 0:
@@ -652,7 +716,7 @@ def maybe_install_ar_post_saturation_mode(self) -> None:
def memory_init_training(
self,
gen_data_clean: GenerationDataClean,
- data_batch: dict[str, torch.Tensor],
+ data_batch: dict[str, Any],
input_text_indexes: list[list[int]],
) -> tuple[GenerationDataClean, dict]:
"""Prepare per-step memory info for causal training.
@@ -673,7 +737,13 @@ def memory_init_training(
gen_data_clean = self._truncate_for_chunkwise_tf(gen_data_clean)
self._assert_chunkwise_tf_shape(gen_data_clean)
+ # Keep the exact tokenizer task with its CP owner's cached training payload.
+ # Later CP window slots ignore the rank-local raw batch entirely.
+ text_system_prompt = resolve_text_system_prompt(data_batch)
+ if isinstance(text_system_prompt, (list, tuple)):
+ text_system_prompt = list(text_system_prompt)
return gen_data_clean, {
+ TEXT_SYSTEM_PROMPT_KEY: text_system_prompt,
"skip_text": False,
"initial_temporal_offset": 0,
"dual_kv_cache": None,
@@ -1265,8 +1335,6 @@ def build_memory_state(
if dual_kv_cache is not None:
vision_token_shapes = packed_seq.vision.token_shapes if packed_seq.vision else None
if use_ar_rolling:
- if batched_ar:
- raise ValueError("Batched AR does not support the compiled rolling-cache path")
# Static-shape AR inference at frame >= 1 (compile + CG).
# ``for_cuda_graphs=True`` makes ``read_for_layer`` return
# the full preallocated gen buffer + a real-length scalar
@@ -1290,8 +1358,11 @@ def build_memory_state(
write_gen_cache=write_gen_cache,
kv_head_shard_rank=kv_head_shard_rank,
kv_head_shard_size=kv_head_shard_size,
+ coarse_cuda_graph=coarse_cuda_graph,
+ stage_gen_cache_writes=stage_gen_cache_writes,
transfer_history_sink_tokens=transfer_history_sink_tokens,
transfer_history_max_tokens=transfer_history_max_tokens,
+ batched=batched_ar,
)
return ARMemoryState(
dual_kv_cache=dual_kv_cache,
@@ -1371,7 +1442,7 @@ def iter_samples_from_batch_autoregressive_streaming_transfer(
self,
*,
data_batch: dict[str, Any],
- control_latent_chunks: Iterable[torch.Tensor], # items: [B,C,T,H,W]
+ control_latent_chunks: Iterable[torch.Tensor | StreamingTransferStep], # items: [B,C,T,H,W]
num_frames: int,
seeds: list[int],
guidance: float = 1.0,
@@ -1411,8 +1482,14 @@ def iter_samples_from_batch_autoregressive_streaming_transfer(
"Batched streaming Transfer requires teacher_forcing_replay_policy.control_visibility='causal', "
"controls_read_strict_past_clean_rgb=True, and clean_pass_causality='frame'"
)
- if self.config.compile.enabled:
- raise ValueError("Batched streaming Transfer requires eager attention")
+ # Compile + CUDA graphs run the same static-shape path as single-view Transfer, with one
+ # buffer row per sample; compile without CUDA graphs runs the dynamic-shape path.
+ use_ar_rolling_path = (
+ self.config.compile.enabled
+ and self.config.compile.use_cuda_graphs
+ and self.config.compile.ar_post_saturation_mode == "default"
+ )
+ _validate_kv_cache_dtype_supports_cuda_graphs(self.config.kv_cache_dtype, use_ar_rolling_path)
if self.parallel_dims is not None and (
getattr(self.parallel_dims, "cp_enabled", False) or getattr(self.parallel_dims, "cfgp_enabled", False)
):
@@ -1429,6 +1506,7 @@ def iter_samples_from_batch_autoregressive_streaming_transfer(
raise ValueError("Finite-window batched streaming Transfer currently requires chunk_size=1")
reset_ar_post_saturation_runtime_for_generation(self)
+ self._reset_ar_forward_cuda_graph_runtime_for_generation()
has_negative_prompt = has_negative_prompt or f"neg_{self.input_caption_key}" in data_batch
cond_text_tokens, uncond_text_tokens = self._get_inference_text_tokens(data_batch, has_negative_prompt)
if len(cond_text_tokens) != batch_size or len(uncond_text_tokens) != batch_size:
@@ -1451,14 +1529,19 @@ def iter_samples_from_batch_autoregressive_streaming_transfer(
else:
gen_cache_size = 2 * num_frames + 1
physical_attention_sink_size = 0
+ kv_buffer_pool = self._ar_kv_buffer_pool_for_generation()
+ preallocate_ring = chunk_size == 1 # rings need uniform entry sizes; the cache is always finite here
dual_kv_cache = [
DualKVCache(
gen_cache_size=gen_cache_size,
kv_cache_dtype=self.config.kv_cache_dtype,
kv_cache_kernel_impl=self.config.kv_cache_kernel_impl,
attention_sink_size=physical_attention_sink_size,
+ buffer_pool=kv_buffer_pool,
+ pool_slot=layer_idx,
+ preallocate_ring=preallocate_ring,
)
- for _ in range(num_layers)
+ for layer_idx in range(num_layers)
]
dual_kv_cache_uncond = (
[
@@ -1467,8 +1550,11 @@ def iter_samples_from_batch_autoregressive_streaming_transfer(
kv_cache_dtype=self.config.kv_cache_dtype,
kv_cache_kernel_impl=self.config.kv_cache_kernel_impl,
attention_sink_size=physical_attention_sink_size,
+ buffer_pool=kv_buffer_pool,
+ pool_slot=num_layers + layer_idx,
+ preallocate_ring=preallocate_ring,
)
- for _ in range(num_layers)
+ for layer_idx in range(num_layers)
]
if cfg_active
else None
@@ -1495,6 +1581,8 @@ def iter_samples_from_batch_autoregressive_streaming_transfer(
enable_fps_modulation = bool(self.config.diffusion_expert_config.enable_fps_modulation)
base_fps = float(self.config.diffusion_expert_config.base_fps)
control_chunks = iter(control_latent_chunks)
+ seeds = list(seeds)
+ row_position_base = [0] * batch_size # per-row mRoPE frame clock origin (episode start)
transfer_history_cache_idx = 0
transfer_history_sink_tokens = 0
transfer_history_control_max_tokens: int | None = None
@@ -1503,9 +1591,32 @@ def iter_samples_from_batch_autoregressive_streaming_transfer(
for chunk_start, chunk_end in _iter_ar_chunk_ranges(0, num_frames, chunk_size):
chunk_len = chunk_end - chunk_start
try:
- control_latent = next(control_chunks).to(**self.tensor_kwargs) # [B,C,chunk_len,H,W]
+ step = next(control_chunks)
except StopIteration as error:
raise ValueError(f"Missing streamed control chunk for frames [{chunk_start}, {chunk_end})") from error
+ reset_rows: tuple[int, ...] = ()
+ if isinstance(step, StreamingTransferStep):
+ reset_rows = tuple(step.reset_rows)
+ control_latent = step.control.to(**self.tensor_kwargs) # [B,C,chunk_len,H,W]
+ if reset_rows:
+ if not use_ar_rolling_path:
+ raise ValueError("Per-row episode restarts require the compiled CUDA-graph AR path")
+ if chunk_start == 0:
+ raise ValueError("Rows cannot restart at the first step; they start fresh there already")
+ if step.seeds is not None and len(step.seeds) != len(reset_rows):
+ raise ValueError(f"Expected {len(reset_rows)} seeds for the reset rows, got {len(step.seeds)}")
+ for position, row in enumerate(reset_rows):
+ if row < 0 or row >= batch_size:
+ raise IndexError(f"reset row {row} outside [0, {batch_size})")
+ row_position_base[row] = chunk_start
+ if step.seeds is not None:
+ seeds[row] = int(step.seeds[position])
+ for cache_list in (dual_kv_cache, dual_kv_cache_uncond or []):
+ for cache in cache_list:
+ cache.gen_cache.reset_rows(reset_rows, at_frame_idx=transfer_history_cache_idx)
+ else:
+ control_latent = step.to(**self.tensor_kwargs) # [B,C,chunk_len,H,W]
+ row_positions = [chunk_start - base for base in row_position_base] # per-row episode-local frame
if control_latent.ndim != 5 or control_latent.shape[0] != batch_size:
raise ValueError(
f"Expected control latent [B,C,T,H,W] with B={batch_size}, got {tuple(control_latent.shape)}"
@@ -1529,7 +1640,7 @@ def iter_samples_from_batch_autoregressive_streaming_transfer(
self._seed_frame_into_kv_cache(
frame_latent=control_latent,
frame_idx=transfer_history_cache_idx,
- position_frame_idx=chunk_start,
+ position_frame_idx=row_positions,
dual_kv_cache=dual_kv_cache,
dual_kv_cache_uncond=dual_kv_cache_uncond,
cond_text_tokens=cond_text_tokens if include_text else None,
@@ -1555,12 +1666,13 @@ def iter_samples_from_batch_autoregressive_streaming_transfer(
transfer_history_sink_tokens=transfer_history_sink_tokens,
transfer_history_max_tokens=transfer_history_control_max_tokens,
batched_ar=True,
+ use_ar_rolling_path=use_ar_rolling_path,
)
transfer_history_cache_idx += 1
noise_rows = []
for sample_idx, seed in enumerate(seeds):
- generator = torch.Generator(device=control_latent.device).manual_seed(seed + chunk_start)
+ generator = torch.Generator(device=control_latent.device).manual_seed(seed + row_positions[sample_idx])
noise_row = torch.empty_like(control_latent[sample_idx]).normal_(generator=generator) # [C,T,H,W]
noise_rows.append(noise_row)
curr_vision_latent = torch.stack(noise_rows, dim=0) # [B,C,T,H,W]
@@ -1572,7 +1684,7 @@ def iter_samples_from_batch_autoregressive_streaming_transfer(
special_tokens=self.llm_special_tokens,
latent_patch_size=patch_size,
condition_frame_indexes_vision=[],
- frame_idx=chunk_start,
+ frame_idx=row_positions,
temporal_compression_factor=tcf,
video_temporal_causal=video_temporal_causal,
enable_fps_modulation=enable_fps_modulation,
@@ -1589,7 +1701,7 @@ def iter_samples_from_batch_autoregressive_streaming_transfer(
special_tokens=self.llm_special_tokens,
latent_patch_size=patch_size,
condition_frame_indexes_vision=[],
- frame_idx=chunk_start,
+ frame_idx=row_positions,
temporal_compression_factor=tcf,
video_temporal_causal=video_temporal_causal,
enable_fps_modulation=enable_fps_modulation,
@@ -1610,7 +1722,7 @@ def iter_samples_from_batch_autoregressive_streaming_transfer(
gen_data_clean=gen_data_clean,
dual_kv_cache=dual_kv_cache,
dual_kv_cache_uncond=dual_kv_cache_uncond,
- frame_idx=chunk_start,
+ frame_idx=row_positions,
cache_frame_idx=transfer_history_cache_idx,
num_frames=num_frames,
guidance=guidance,
@@ -1622,7 +1734,7 @@ def iter_samples_from_batch_autoregressive_streaming_transfer(
distilled_num_steps=distilled_num_steps,
fps_vision_list=fps_vision_list,
fps_action_list=[24.0] * batch_size,
- use_ar_rolling_path=False,
+ use_ar_rolling_path=use_ar_rolling_path,
transfer_history_sink_tokens=transfer_history_sink_tokens,
transfer_history_max_tokens=transfer_history_target_max_tokens,
batched_ar=True,
@@ -1639,7 +1751,7 @@ def iter_samples_from_batch_autoregressive_streaming_transfer(
self._seed_frame_into_kv_cache(
frame_latent=frame_latent,
frame_idx=transfer_history_cache_idx,
- position_frame_idx=frame_idx,
+ position_frame_idx=[frame_idx - base for base in row_position_base],
dual_kv_cache=dual_kv_cache,
dual_kv_cache_uncond=dual_kv_cache_uncond,
cond_text_tokens=None,
@@ -1664,6 +1776,7 @@ def iter_samples_from_batch_autoregressive_streaming_transfer(
transfer_history_sink_tokens=transfer_history_sink_tokens,
transfer_history_max_tokens=transfer_history_target_max_tokens,
batched_ar=True,
+ use_ar_rolling_path=use_ar_rolling_path,
)
transfer_history_cache_idx += 1
yield {"vision": denoised_chunk[:, :, local_idx : local_idx + 1]} # [B,C,1,H,W]
@@ -1674,6 +1787,29 @@ def iter_samples_from_batch_autoregressive_streaming_transfer(
return
raise ValueError(f"Received an extra streamed control chunk with shape {tuple(extra_chunk.shape)}")
+ def _reset_ar_forward_cuda_graph_runtime_for_generation(self) -> None:
+ """Validate the whole-forward CUDA-graph scope and drop last generation's captures."""
+ if self.config.compile.cuda_graph_scope == "forward" and self.config.compile.enabled:
+ if not self.config.compile.use_cuda_graphs or self.config.compile.ar_post_saturation_mode != "default":
+ raise ValueError(
+ "compile.cuda_graph_scope='forward' requires compile.use_cuda_graphs=True and "
+ "ar_post_saturation_mode='default' (it captures the static-shape AR forward itself)."
+ )
+ if self.parallel_dims is not None and self.parallel_dims.cfgp_enabled:
+ # CFG parallelism bypasses the whole-forward graphs, and the static-shape blocks
+ # (compiled with dynamic=False for capture) would recompile per frame.
+ raise ValueError("compile.cuda_graph_scope='forward' requires cfgp_size=1.")
+ # Graphs are bound to this generation's KV caches (und cache addresses); recapture per generation.
+ get_ar_forward_cuda_graph_runtime(self).reset_for_new_generation()
+
+ def _ar_kv_buffer_pool_for_generation(self) -> KVBufferPool:
+ """Model-owned pool for the per-layer K/V workspaces (stable addresses across generations)."""
+ kv_buffer_pool = getattr(self, "_ar_kv_buffer_pool", None)
+ if not isinstance(kv_buffer_pool, KVBufferPool):
+ kv_buffer_pool = KVBufferPool()
+ self._ar_kv_buffer_pool = kv_buffer_pool
+ return kv_buffer_pool
+
@torch.no_grad()
def iter_samples_from_batch_autoregressive(
self,
@@ -1781,6 +1917,7 @@ def iter_samples_from_batch_autoregressive(
)
if sampler_mode not in ("rf", "distilled"):
raise ValueError(f"sampler_mode must be 'rf' or 'distilled', got {sampler_mode!r}")
+ 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():
@@ -1941,14 +2078,16 @@ def iter_samples_from_batch_autoregressive(
# ``use_ar_rolling_path`` enables the static-shape AR path
# (``ARMemoryState(for_cuda_graphs=True)``) at frame >= 1; gated
# by the velocity_fn on ``frame_idx > 0`` so frame 0 stays on
- # the dynamic-shape branch.
+ # the dynamic-shape branch. Single-view video_transfer uses the
+ # same path from cache idx 1 (first target denoise) onward; its
+ # sink/recent history limits are applied per cache entry inside
+ # ``ARMemoryState.init`` so the padded buffer's real prefix equals
+ # the dynamic branch's sliced history.
use_ar_rolling_path = (
self.config.compile.enabled
and self.config.compile.use_cuda_graphs
and self.config.compile.ar_post_saturation_mode == "default"
)
- if is_transfer and self.config.compile.enabled:
- raise ValueError("video_transfer AR inference requires eager attention; run with --no-use-torch-compile.")
_validate_kv_cache_dtype_supports_cuda_graphs(self.config.kv_cache_dtype, use_ar_rolling_path)
# Chunkwise AR: generate ``chunk_size`` latent frames jointly per AR step
@@ -2001,14 +2140,23 @@ def iter_samples_from_batch_autoregressive(
physical_attention_sink_size = self.config.attention_sink_size
if gen_cache_size is None and use_ar_rolling_path:
gen_cache_size = num_frames
+ # Large per-layer workspaces (in-place K/V rings and the CUDA-graph
+ # static history buffers) come from a model-owned pool so they keep
+ # their addresses across generations; see ``KVBufferPool``. Rings need
+ # a finite cache; unbounded caches keep the lazily grown clone storage.
+ kv_buffer_pool = self._ar_kv_buffer_pool_for_generation()
+ preallocate_ring = gen_cache_size is not None and chunk_size == 1 # rings need uniform entry sizes
dual_kv_cache = [
DualKVCache(
gen_cache_size=gen_cache_size,
kv_cache_dtype=self.config.kv_cache_dtype,
kv_cache_kernel_impl=self.config.kv_cache_kernel_impl,
attention_sink_size=physical_attention_sink_size,
+ buffer_pool=kv_buffer_pool,
+ pool_slot=layer_idx,
+ preallocate_ring=preallocate_ring,
)
- for _ in range(num_layers)
+ for layer_idx in range(num_layers)
]
dual_kv_cache_uncond = (
[
@@ -2017,8 +2165,11 @@ def iter_samples_from_batch_autoregressive(
kv_cache_dtype=self.config.kv_cache_dtype,
kv_cache_kernel_impl=self.config.kv_cache_kernel_impl,
attention_sink_size=physical_attention_sink_size,
+ buffer_pool=kv_buffer_pool,
+ pool_slot=num_layers + layer_idx,
+ preallocate_ring=preallocate_ring,
)
- for _ in range(num_layers)
+ for layer_idx in range(num_layers)
]
if cfg_active and not cfgp_enabled
else None
@@ -2218,7 +2369,9 @@ def iter_samples_from_batch_autoregressive(
enable_fps_mod=_enable_fps_mod,
base_fps=_base_fps,
modality_margin=_margin,
- use_ar_rolling_path=False,
+ # Cache idx 0 (text + first control) stays dynamic inside the
+ # seed; every later control seed shares the static-shape path.
+ use_ar_rolling_path=use_ar_rolling_path,
condition_frame_indexes_vision=list(range(chunk_len)),
transfer_history_sink_tokens=transfer_history_sink_tokens,
transfer_history_max_tokens=transfer_history_control_max_tokens,
@@ -2871,7 +3024,7 @@ def _seed_frame_into_kv_cache(
base_fps: float,
modality_margin: int,
use_ar_rolling_path: bool = False,
- position_frame_idx: int | None = None,
+ position_frame_idx: int | list[int] | None = None,
condition_frame_indexes_vision: list[int] | None = None,
transfer_history_sink_tokens: int = 0,
transfer_history_max_tokens: int | None = None,
@@ -2921,6 +3074,8 @@ def _seed_frame_into_kv_cache(
assert isinstance(seed, int)
sigma = self.config.sigma_diffusion_forcing
# Deterministic ε shared across cond/uncond packs.
+ if isinstance(position_frame_idx, list):
+ raise ValueError("diffusion_forcing seeding does not support per-row frame indices")
g = torch.Generator(device=frame_latent.device).manual_seed(seed + position_frame_idx)
eps = torch.empty_like(frame_latent).normal_(generator=g) # [1,C,T,H,W]
frame_in = sigma * eps + (1.0 - sigma) * frame_latent # [1,C,T,H,W]
@@ -3003,6 +3158,10 @@ def _build_pack(
use_ar_rolling = use_ar_rolling_path and frame_idx > 0
post_saturation_static_compile = is_ar_post_saturation_static_compile_frame(self, frame_idx)
post_saturation_cuda_graph = is_ar_post_saturation_cuda_graph_frame(self, frame_idx)
+ # Whole-forward graph per (kind, branch); the Transfer history limit joins the key
+ # because control seeds and RGB refreshes see different history lengths.
+ forward_cuda_graph = use_ar_rolling and not cfgp_enabled and uses_ar_forward_cuda_graph(self)
+ forward_graph_kind = f"refresh:{transfer_history_max_tokens}"
memory_info = {
"dual_kv_cache": dual_kv_cache,
"use_rolling_kv_cache": False,
@@ -3036,6 +3195,14 @@ def _build_pack(
packed_seq=cond_pack,
memory_info=memory_info,
)
+ elif forward_cuda_graph:
+ run_ar_forward_cuda_graph(
+ self,
+ kind=forward_graph_kind,
+ branch="conditional",
+ packed_seq=cond_pack,
+ memory_info=memory_info,
+ )
else:
torch.compiler.cudagraph_mark_step_begin()
memory = self.build_memory_state(cond_pack, memory_info)
@@ -3065,6 +3232,14 @@ def _build_pack(
packed_seq=uncond_pack,
memory_info=memory_info_uncond,
)
+ elif forward_cuda_graph:
+ run_ar_forward_cuda_graph(
+ self,
+ kind=forward_graph_kind,
+ branch="unconditional",
+ packed_seq=uncond_pack,
+ memory_info=memory_info_uncond,
+ )
else:
torch.compiler.cudagraph_mark_step_begin()
memory_uncond = self.build_memory_state(uncond_pack, memory_info_uncond)
@@ -3102,7 +3277,7 @@ def _run_distilled_ar_sampler(
initial_noise: torch.Tensor, # [B,N_tokens_flat]
*,
seed: int | list[int],
- frame_idx: int,
+ frame_idx: int | list[int],
num_frames: int | None = None,
distilled_num_steps: int | None = None,
) -> torch.Tensor: # [B,N_tokens_flat]
@@ -3115,6 +3290,14 @@ def _run_distilled_ar_sampler(
sample_type = self.config.fixed_step_sampler_config.sample_type
max_timestep = float(self.config.rectified_flow_inference_config.num_train_timesteps)
x = initial_noise.float() # [B,N_tokens_flat]
+ # One asynchronous copy of the schedule instead of a synchronising host->device scalar per
+ # step (same float32 rounding of the same Python values; the deltas are formed in double
+ # precision exactly as the previous per-step ``torch.as_tensor(sigma_next - sigma_cur)``).
+ sigmas = to_device_nonblocking(torch.tensor(full_t_list, dtype=torch.float32), x.device) # [N_steps+1]
+ sigma_deltas = to_device_nonblocking(
+ torch.tensor([nxt - cur for cur, nxt in zip(full_t_list[:-1], full_t_list[1:])], dtype=torch.float32),
+ x.device,
+ ) # [N_steps]
for step_idx, (sigma_cur, sigma_next) in enumerate(zip(full_t_list[:-1], full_t_list[1:])):
timestep = torch.full(
(x.shape[0], 1),
@@ -3123,14 +3306,14 @@ def _run_distilled_ar_sampler(
device=x.device,
) # [B,1]
velocity = velocity_fn(x, timestep).float() # [B,N_tokens_flat]
- sigma_cur_tensor = torch.as_tensor(sigma_cur, dtype=torch.float32, device=x.device) # []
+ sigma_cur_tensor = sigmas[step_idx] # []
x0_pred = x - sigma_cur_tensor * velocity # [B,N_tokens_flat]
if sigma_next == 0.0:
x = x0_pred # [B,N_tokens_flat]
continue
- sigma_next_tensor = torch.as_tensor(sigma_next, dtype=torch.float32, device=x.device) # []
+ sigma_next_tensor = sigmas[step_idx + 1] # []
if sample_type == "ode":
- delta_sigma = torch.as_tensor(sigma_next - sigma_cur, dtype=torch.float32, device=x.device) # []
+ delta_sigma = sigma_deltas[step_idx] # []
x = x + delta_sigma * velocity # [B,N_tokens_flat]
elif sample_type == "sde":
# Use a mixed seed so different frame/step pairs cannot collide
@@ -3142,12 +3325,15 @@ def _run_distilled_ar_sampler(
raise ValueError(f"Expected {x.shape[0]} seeds, got {len(seed)}")
noise_rows = []
for sample_idx, sample_seed in enumerate(seed):
- step_seed = int(sample_seed) + int(frame_idx) * 1_000_003 + (int(step_idx) + 1) * 9_176
+ row_frame = frame_idx[sample_idx] if isinstance(frame_idx, (list, tuple)) else frame_idx
+ step_seed = int(sample_seed) + int(row_frame) * 1_000_003 + (int(step_idx) + 1) * 9_176
generator = torch.Generator(device=x.device).manual_seed(step_seed)
noise_row = torch.empty_like(x[sample_idx]).normal_(generator=generator) # [N_tokens_flat]
noise_rows.append(noise_row)
noise = torch.stack(noise_rows, dim=0) # [B,N_tokens_flat]
else:
+ if isinstance(frame_idx, (list, tuple)):
+ raise ValueError("Per-row frame indices require a per-row seed list")
step_seed = int(seed) + int(frame_idx) * 1_000_003 + (int(step_idx) + 1) * 9_176
generator = torch.Generator(device=x.device).manual_seed(step_seed)
noise = torch.empty_like(x).normal_(generator=generator) # [B,N_tokens_flat]
@@ -3243,7 +3429,7 @@ def _run_ar_sampler(
num_steps: int,
shift: float,
seed: int | list[int],
- sample_idx: int,
+ sample_idx: int | list[int],
num_frames: int | None,
distilled_num_steps: int | None,
) -> torch.Tensor: # [B,N_tokens_flat]
@@ -3303,7 +3489,7 @@ def generate_next_frame(
seed: int | list[int],
fps_vision_list: list[float],
fps_action_list: list[float],
- frame_idx: int | None = None,
+ frame_idx: int | list[int] | None = None,
cache_frame_idx: int | None = None,
num_frames: int | None = None,
normalize_cfg: bool = False,
@@ -3357,6 +3543,8 @@ def generate_next_frame(
if frame_idx is None:
frame_idx = 0
if cache_frame_idx is None:
+ if isinstance(frame_idx, list):
+ raise ValueError("Per-row frame indices require an explicit cache_frame_idx")
cache_frame_idx = frame_idx
# packed_seq now passed as parameter, containing all frames up to current
@@ -3438,6 +3626,7 @@ def velocity_fn(
use_ar_rolling = use_ar_rolling_path and cache_frame_idx > 0
post_saturation_static_compile = is_ar_post_saturation_static_compile_frame(self, cache_frame_idx)
post_saturation_cuda_graph = is_ar_post_saturation_cuda_graph_frame(self, cache_frame_idx)
+ forward_cuda_graph = use_ar_rolling and not cfgp_enabled and uses_ar_forward_cuda_graph(self)
def run_branch(
pack: PackedSequence,
@@ -3475,6 +3664,14 @@ def run_branch(
packed_seq=pack,
memory_info=memory_info,
)
+ elif forward_cuda_graph:
+ output = run_ar_forward_cuda_graph(
+ self,
+ kind=f"denoise:{transfer_history_max_tokens}",
+ branch=branch,
+ packed_seq=pack,
+ memory_info=memory_info,
+ )
else:
torch.compiler.cudagraph_mark_step_begin()
memory = self.build_memory_state(pack, memory_info)
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 e15644825..8f8828f25 100644
--- a/cosmos_framework/model/generator/omni_mot_causal_model_test.py
+++ b/cosmos_framework/model/generator/omni_mot_causal_model_test.py
@@ -18,7 +18,7 @@
uninitialized.
"""
-from types import SimpleNamespace
+from types import MethodType, SimpleNamespace
from typing import Literal, cast
from unittest.mock import MagicMock, patch
@@ -2413,6 +2413,139 @@ def test_video_transfer_finite_window_rejects_non_history_replay(self) -> None:
with pytest.raises(ValueError, match="controls_read_strict_past_clean_rgb=True"):
self._run(model, "video_transfer")
+ @pytest.mark.L0
+ @pytest.mark.CPU
+ def test_video_transfer_with_cuda_graphs_routes_cached_forwards_to_the_rolling_path(self) -> None:
+ """compile + CUDA graphs: control seeds, RGB seeds and denoise all get use_ar_rolling_path=True."""
+ model = self._make_model_mock()
+ model.config.compile.enabled = True
+ model.config.compile.use_cuda_graphs = True
+ model.config.compile.ar_post_saturation_mode = "default"
+ model.config.kv_cache_inference_size = 3
+ model.config.attention_sink_size = 1
+
+ result, _ = self._run(model, "video_transfer")
+
+ assert result["vision"].shape == (1, self.C, self.T, self.H, self.W)
+ seed_calls = model._seed_frame_into_kv_cache.call_args_list
+ assert len(seed_calls) == 2 * self.T - 1 # T control seeds + (T-1) RGB refreshes
+ assert all(call.kwargs["use_ar_rolling_path"] is True for call in seed_calls)
+ target_calls = model.generate_next_frame.call_args_list
+ assert len(target_calls) == self.T
+ assert all(call.kwargs["use_ar_rolling_path"] is True for call in target_calls)
+
+ @pytest.mark.L0
+ @pytest.mark.CPU
+ def test_video_transfer_with_compile_but_no_cuda_graphs_stays_on_the_dynamic_path(self) -> None:
+ """compile without CUDA graphs is allowed and keeps every forward dynamic-shape."""
+ model = self._make_model_mock()
+ model.config.compile.enabled = True
+ model.config.compile.use_cuda_graphs = False
+ model.config.compile.ar_post_saturation_mode = "default"
+
+ self._run(model, "video_transfer")
+
+ seed_calls = model._seed_frame_into_kv_cache.call_args_list
+ assert seed_calls and all(call.kwargs["use_ar_rolling_path"] is False for call in seed_calls)
+ target_calls = model.generate_next_frame.call_args_list
+ assert target_calls and all(call.kwargs["use_ar_rolling_path"] is False for call in target_calls)
+
+ @pytest.mark.L0
+ @pytest.mark.CPU
+ def test_ar_caches_use_the_model_buffer_pool_and_preallocated_rings(self) -> None:
+ """Finite AR caches get the model-owned pool, per-layer slots (CFG branch offset) and in-place rings."""
+ from cosmos_framework.model.generator.omni_mot_causal_model import OmniMoTCausalModel
+ from cosmos_framework.model.generator.utils.kv_cache import KVBufferPool
+
+ model = self._make_model_mock()
+ model._ar_kv_buffer_pool_for_generation = MethodType(
+ OmniMoTCausalModel._ar_kv_buffer_pool_for_generation, model
+ )
+ model.config.kv_cache_inference_size = 3
+ model.config.attention_sink_size = 1
+ model._ar_kv_buffer_pool = None # a real model starts without a pool; the loop creates one lazily
+ model.get_data_and_condition.return_value = self._make_gen_data("video_transfer")
+ model._get_inference_text_tokens.return_value = ([[1, 2, 3]], [[7, 8, 9]])
+ data_batch = {"caption": ["a prompt"], "neg_caption": ["avoid artifacts"]}
+
+ with patch(_PATCH_PACK, MagicMock()), patch(_PATCH_KV) as dual_cache_cls:
+ list(
+ OmniMoTCausalModel.iter_samples_from_batch_autoregressive(
+ model, data_batch, mode="video_transfer", guidance=7.0, has_negative_prompt=True
+ )
+ )
+ calls = dual_cache_cls.call_args_list
+ assert len(calls) == 2 * model.net.num_hidden_layers # cond + uncond caches
+ pools = {id(call.kwargs["buffer_pool"]) for call in calls}
+ assert len(pools) == 1 and isinstance(calls[0].kwargs["buffer_pool"], KVBufferPool)
+ assert model._ar_kv_buffer_pool is calls[0].kwargs["buffer_pool"]
+ assert [call.kwargs["pool_slot"] for call in calls] == list(range(2 * model.net.num_hidden_layers))
+ assert all(call.kwargs["preallocate_ring"] is True for call in calls)
+
+ # A second generation reuses the same pool object.
+ with patch(_PATCH_PACK, MagicMock()), patch(_PATCH_KV) as dual_cache_cls_2:
+ list(
+ OmniMoTCausalModel.iter_samples_from_batch_autoregressive(
+ model, data_batch, mode="video_transfer", guidance=7.0, has_negative_prompt=True
+ )
+ )
+ assert dual_cache_cls_2.call_args_list[0].kwargs["buffer_pool"] is model._ar_kv_buffer_pool
+
+ @pytest.mark.L0
+ @pytest.mark.CPU
+ def test_chunkwise_transfer_caches_do_not_preallocate_rings(self) -> None:
+ """Chunkwise transfer stores entries of different token counts; the ring needs uniform entries."""
+ from cosmos_framework.model.generator.omni_mot_causal_model import OmniMoTCausalModel
+
+ model = self._make_model_mock()
+ model._ar_kv_buffer_pool = None
+ model.config.teacher_forcing_frames_per_chunk = 4
+ num_frames = 9
+ control = torch.ones(1, self.C, num_frames, self.H, self.W) # [B,C,T,H,W]
+ model.get_data_and_condition.return_value = SimpleNamespace(
+ batch_size=1,
+ x0_tokens_vision=[control, torch.zeros_like(control)],
+ num_vision_items_per_sample=[2],
+ x0_tokens_action=None,
+ fps_vision=torch.tensor([24.0]), # [B]
+ fps_action=None,
+ action_domain_id=None,
+ raw_action_dim=None,
+ )
+ model._get_inference_text_tokens.return_value = ([[1, 2, 3]], None)
+ model.generate_next_frame.side_effect = lambda **kwargs: torch.zeros_like(kwargs["curr_vision_latent"])
+ with patch(_PATCH_PACK, MagicMock()), patch(_PATCH_KV) as dual_cache_cls:
+ list(OmniMoTCausalModel.iter_samples_from_batch_autoregressive(model, {}, mode="video_transfer"))
+ assert dual_cache_cls.call_args_list
+ assert all(call.kwargs["preallocate_ring"] is False for call in dual_cache_cls.call_args_list)
+
+ @pytest.mark.L0
+ @pytest.mark.CPU
+ def test_unbounded_ar_caches_do_not_preallocate_rings(self) -> None:
+ """text2video without kv_cache_inference_size keeps the lazily grown clone storage."""
+ from cosmos_framework.model.generator.omni_mot_causal_model import OmniMoTCausalModel
+
+ model = self._make_model_mock()
+ model._ar_kv_buffer_pool = None
+ with patch(_PATCH_PACK, MagicMock()), patch(_PATCH_KV) as dual_cache_cls:
+ model.get_data_and_condition.return_value = self._make_gen_data("text2video")
+ model._get_inference_text_tokens.return_value = ([[1, 2, 3]], None)
+ list(OmniMoTCausalModel.iter_samples_from_batch_autoregressive(model, {}, mode="text2video"))
+ assert dual_cache_cls.call_args_list
+ assert all(call.kwargs["gen_cache_size"] is None for call in dual_cache_cls.call_args_list)
+ assert all(call.kwargs["preallocate_ring"] is False for call in dual_cache_cls.call_args_list)
+
+ @pytest.mark.L0
+ @pytest.mark.CPU
+ def test_multiview_transfer_with_compile_still_requires_eager_attention(self) -> None:
+ """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
+
+ with pytest.raises(ValueError, match="Multiview transfer AR requires eager attention"):
+ self._run(model, "video_transfer")
+
@pytest.mark.L0
@pytest.mark.CPU
def test_video_transfer_chunkwise_generation_matches_training_partition(self) -> None:
@@ -3129,6 +3262,36 @@ def velocity_fn(x: torch.Tensor, timestep: torch.Tensor) -> torch.Tensor: # x:
if len(sample_seeds) > 1:
assert not torch.equal(reinjected_noise[0], reinjected_noise[1])
+ @pytest.mark.L0
+ @pytest.mark.CPU
+ def test_distilled_sde_sampler_per_row_frame_indices_restart_the_noise_stream(self) -> None:
+ """A restarted row (episode-local frame 0) draws the same SDE noise as a fresh single-row run."""
+ from cosmos_framework.model.generator.omni_mot_causal_model import OmniMoTCausalModel
+
+ model = SimpleNamespace(
+ config=SimpleNamespace(
+ fixed_step_sampler_config=SimpleNamespace(t_list=[0.5, 0.25], sample_type="sde"),
+ rectified_flow_inference_config=SimpleNamespace(num_train_timesteps=1000),
+ )
+ )
+ self._attach_distilled_schedule_helper(model, OmniMoTCausalModel)
+
+ def velocity_fn(x: torch.Tensor, timestep: torch.Tensor) -> torch.Tensor: # noqa: ARG001
+ return torch.zeros_like(x) # [B,N]
+
+ batched = OmniMoTCausalModel._run_distilled_ar_sampler(
+ model, velocity_fn, torch.ones(2, 3), seed=[10, 10], frame_idx=[3, 0]
+ ) # [2,N]
+ row0 = OmniMoTCausalModel._run_distilled_ar_sampler(
+ model, velocity_fn, torch.ones(1, 3), seed=[10], frame_idx=[3]
+ )
+ row1 = OmniMoTCausalModel._run_distilled_ar_sampler(
+ model, velocity_fn, torch.ones(1, 3), seed=[10], frame_idx=[0]
+ )
+ torch.testing.assert_close(batched[0:1], row0)
+ torch.testing.assert_close(batched[1:2], row1)
+ assert not torch.equal(batched[0], batched[1])
+
@pytest.mark.L0
@pytest.mark.CPU
def test_distilled_sde_sampler_uses_distinct_frame_seeds(self) -> None:
@@ -3485,3 +3648,162 @@ def test_teacher_forcing_replay_policy_still_rejects_unknown_real_fields() -> No
with pytest.raises(TypeError, match="control_visibilty"):
_resolve_teacher_forcing_replay_policy({"_type": "x", "control_visibilty": "current"})
+
+
+def _forward_graph_model_mock() -> MagicMock:
+ model = MagicMock()
+ model.config.causal_training_strategy = "none"
+ model.config.compile.enabled = True
+ model.config.compile.use_cuda_graphs = True
+ model.config.compile.cuda_graph_scope = "forward"
+ model.config.compile.ar_post_saturation_mode = "default"
+ model.config.rectified_flow_inference_config.scheduler_type = "unipc"
+ model.parallel_dims = None
+ from cosmos_framework.model.generator.omni_mot_causal_model import OmniMoTCausalModel
+
+ # Real validation/reset logic on the mock (the loop calls it through ``self``).
+ model._reset_ar_forward_cuda_graph_runtime_for_generation = MethodType(
+ OmniMoTCausalModel._reset_ar_forward_cuda_graph_runtime_for_generation, model
+ )
+ return model
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+@pytest.mark.parametrize("frame_idx", [0, 2])
+def test_seed_routes_static_frames_to_the_forward_cuda_graph(frame_idx: int) -> None:
+ """cuda_graph_scope="forward": cache idx >= 1 seeds replay a whole-forward graph, idx 0 stays eager."""
+ from cosmos_framework.model.generator.omni_mot_causal_model import OmniMoTCausalModel
+
+ model = _forward_graph_model_mock()
+ cond_pack = MagicMock(name="cond_pack")
+ with (
+ patch(_PATCH_PACK, MagicMock(return_value=cond_pack)),
+ patch("cosmos_framework.model.generator.omni_mot_causal_model.run_ar_forward_cuda_graph") as run_graph,
+ ):
+ OmniMoTCausalModel._seed_frame_into_kv_cache(
+ model,
+ frame_latent=torch.zeros(1, 4, 1, 2, 2), # [B,C,T,H,W]
+ frame_idx=frame_idx,
+ dual_kv_cache=[MagicMock()],
+ dual_kv_cache_uncond=None,
+ cond_text_tokens=None,
+ uncond_text_tokens=None,
+ cond_cached_text_offset=0,
+ uncond_cached_text_offset=0,
+ curr_action_latent=None,
+ action_domain_id=None,
+ gen_data_clean=SimpleNamespace(fps_vision=None, fps_action=None),
+ fps_vision_list=[24.0],
+ fps_action_list=[24.0],
+ seed=42,
+ cfg_active=False,
+ cfgp_enabled=False,
+ tcf=4,
+ patch_size=1,
+ action_dim=8,
+ video_tc=False,
+ enable_fps_mod=False,
+ base_fps=24.0,
+ modality_margin=0,
+ use_ar_rolling_path=True,
+ transfer_history_sink_tokens=8,
+ transfer_history_max_tokens=12,
+ )
+ if frame_idx == 0:
+ run_graph.assert_not_called()
+ model.denoise.assert_called_once()
+ else:
+ model.denoise.assert_not_called()
+ run_graph.assert_called_once()
+ kwargs = run_graph.call_args.kwargs
+ assert kwargs["kind"] == "refresh:12"
+ assert kwargs["branch"] == "conditional"
+ assert kwargs["packed_seq"] is cond_pack
+ assert kwargs["memory_info"]["use_ar_rolling"] is True
+ assert kwargs["memory_info"]["write_gen_cache"] is True
+ assert kwargs["memory_info"]["transfer_history_max_tokens"] == 12
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_denoise_routes_static_frames_to_the_forward_cuda_graph() -> None:
+ """cuda_graph_scope="forward": the sampler's velocity forwards replay the denoise graph."""
+ from cosmos_framework.model.generator.omni_mot_causal_model import OmniMoTCausalModel
+
+ model = _forward_graph_model_mock()
+ cond_pack = MagicMock(name="cond_pack")
+ cond_pack.vision = None
+ cond_pack.action = None
+
+ def sampler(velocity_fn: object, initial_noise: torch.Tensor, **_kwargs: object) -> torch.Tensor:
+ return velocity_fn(initial_noise, torch.ones(1, 1)) # type: ignore[operator]
+
+ model.sampler = sampler
+ graph_output = {"preds_vision": [torch.full((1, 1, 1, 1), 3.0)]} # [C,T,H,W]
+ with patch(
+ "cosmos_framework.model.generator.omni_mot_causal_model.run_ar_forward_cuda_graph",
+ return_value=graph_output,
+ ) as run_graph:
+ denoised = OmniMoTCausalModel.generate_next_frame(
+ model,
+ packed_seq=cond_pack,
+ packed_seq_uncond=None,
+ curr_vision_latent=torch.zeros(1, 1, 1, 1, 1), # [B,C,T,H,W]
+ curr_action_latent=None,
+ cond_text_tokens=[1],
+ uncond_text_tokens=[],
+ gen_data_clean=SimpleNamespace(),
+ dual_kv_cache=[MagicMock()],
+ dual_kv_cache_uncond=None,
+ guidance=1.0,
+ num_steps=1,
+ shift=1.0,
+ seed=7,
+ fps_vision_list=[24.0],
+ fps_action_list=[],
+ frame_idx=1,
+ cache_frame_idx=1,
+ use_ar_rolling_path=True,
+ transfer_history_sink_tokens=8,
+ transfer_history_max_tokens=12,
+ )
+ torch.testing.assert_close(denoised, torch.full((1, 1, 1, 1, 1), 3.0))
+ model.denoise.assert_not_called()
+ kwargs = run_graph.call_args.kwargs
+ assert kwargs["kind"] == "denoise:12"
+ assert kwargs["branch"] == "conditional"
+ assert kwargs["memory_info"]["use_ar_rolling"] is True
+ assert kwargs["memory_info"]["write_gen_cache"] is False
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_forward_cuda_graph_kind_variants_are_accepted() -> None:
+ """Graph keys may carry a variant suffix (e.g. the transfer history limit) after the base kind."""
+ from cosmos_framework.model.generator.mot.post_saturation.cuda_graph import ARPostSaturationCUDAGraphManager
+
+ assert ARPostSaturationCUDAGraphManager._base_kind("refresh:12") == "refresh"
+ assert ARPostSaturationCUDAGraphManager._base_kind("denoise") == "denoise"
+ with pytest.raises(ValueError, match="Unsupported post-saturation CUDA Graph kind"):
+ ARPostSaturationCUDAGraphManager._base_kind("prefill:3")
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_forward_cuda_graph_scope_requires_the_static_path() -> None:
+ """scope=forward without CUDA graphs (or with a post-saturation mode) is a configuration error."""
+ from cosmos_framework.model.generator.omni_mot_causal_model import OmniMoTCausalModel
+
+ model = _forward_graph_model_mock()
+ model.config.compile.ar_post_saturation_mode = "cuda-graph"
+ model._uses_multiview_flex_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.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 0aa0940c6..a70431234 100644
--- a/cosmos_framework/model/generator/omni_mot_model.py
+++ b/cosmos_framework/model/generator/omni_mot_model.py
@@ -68,6 +68,7 @@
GenerationDataNoised,
_expand_per_sample_to_per_vision_item,
build_dense_sound_schedule,
+ select_target_image_sizes,
unwrap_and_densify,
)
from cosmos_framework.model.generator.utils.load_balancing_stats import LBLConfig
@@ -83,6 +84,10 @@
from cosmos_framework.model.generator.utils.safetensors_loader import (
load_language_model as load_language_model_safetensors,
)
+from cosmos_framework.model.generator.utils.sr_latent_noise import (
+ apply_sr_latent_condition_noise,
+ sr_sample_mask,
+)
from cosmos_framework.model.generator.vision_encoder import (
VisionEncoder,
get_vae_pixel_shapes,
@@ -1138,8 +1143,13 @@ def _prepare_training_data(
# image_size[i] may be (1, 4) from IterativeJointDataLoader or (4,) from custom_collate_fn.
if "image_size" in data_batch:
data_resolutions: list[str] | str | None = []
+ # Multi-item samples (transfer, SR) carry one image_size per vision item; the noise
+ # schedule must follow the generated (last) item of each sample.
+ target_image_sizes = select_target_image_sizes(
+ data_batch["image_size"], gen_data_clean.num_vision_items_per_sample, gen_data_clean.batch_size
+ )
for i in range(gen_data_clean.batch_size):
- img_size = data_batch["image_size"][i]
+ img_size = target_image_sizes[i]
if img_size.dim() == 2:
img_size = img_size[0]
target_h = int(img_size[0].item())
@@ -1695,6 +1705,7 @@ 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"
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
@@ -1719,8 +1730,7 @@ def _compute_losses(
rectified_flow=rectified_flow_vision,
tensor_kwargs_fp32=self.tensor_kwargs_fp32,
normalize_by_active=normalize_by_active,
- exclude_fully_conditioned_items=getattr(self.config, "causal_training_strategy", None)
- == "teacher_forcing",
+ exclude_fully_conditioned_items=exclude_fully_conditioned_items,
)
loss_scale = (
rf_cfg.image_loss_scale if is_image_batch and rf_cfg.image_loss_scale is not None else rf_cfg.loss_scale
@@ -1739,6 +1749,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.
fm_loss_lidar, _ = compute_flow_matching_loss(
pred=out_net["preds_lidar"],
target=gen_data_noised.vt_target_lidar,
@@ -1748,6 +1760,7 @@ def _compute_losses(
rectified_flow=self.rectified_flow_video,
tensor_kwargs_fp32=self.tensor_kwargs_fp32,
normalize_by_active=normalize_by_active,
+ exclude_fully_conditioned_items=exclude_fully_conditioned_items,
)
lidar_loss_scale = rf_cfg.lidar_loss_scale if rf_cfg.lidar_loss_scale is not None else rf_cfg.loss_scale
total_loss += fm_loss_lidar * lidar_loss_scale # []
@@ -3076,7 +3089,7 @@ def _get_velocity(
offset += lidar_dim
lidar_offset += n_lidar
- if has_noisy_actions and noise_x_action is not None:
+ 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
action_dim = int(torch.prod(torch.tensor(action_shape)))
@@ -4768,6 +4781,7 @@ def get_data_and_condition(
sample_vision_list = data_batch[media_key]
+ # NOTE: as we assume that the vision items will be passed as a List[List[Tensor]],
# we should always get this information here during training. If we can read this field
# from data_batch it means we are in the visualization callback:
if "num_vision_items_per_sample" not in data_batch:
@@ -4780,6 +4794,7 @@ def get_data_and_condition(
num_vision_items_per_sample: list[int] | None = (
[len(v) for v in sample_vision_list] if has_multiple_vision_per_sample else None
)
+ # NOTE: we need to add this information back into the data_batch, because this
# information is only stored in the GenerationDataClean object which will be discarded
# outside the training loop. Error will be raised when the data batch is passed to the
# visualization callbacks.
@@ -4849,6 +4864,14 @@ def get_data_and_condition(
if frame_size is not None:
x0_tokens_vision = self._remove_padding_from_latent(x0_tokens_vision, frame_size)
+ sr_noise_cfg = getattr(self.config, "sr_latent_condition_noise", None)
+ if sr_noise_cfg is not None and self.training and torch.is_grad_enabled():
+ # L1: noise the LR conditioning latent of SR samples only (never the generated item).
+ eligible = sr_sample_mask(data_batch.get("dataset_name"), batch_size, sr_noise_cfg.dataset_names)
+ x0_tokens_vision, _ = apply_sr_latent_condition_noise(
+ x0_tokens_vision, num_vision_items_per_sample, eligible, sr_noise_cfg
+ )
+
temporal_positions_vision = self._get_temporal_positions_vision(
raw_state_vision=raw_state_vision,
x0_tokens_vision=x0_tokens_vision,
diff --git a/cosmos_framework/model/generator/tokenizers/lidar/__init__.py b/cosmos_framework/model/generator/tokenizers/lidar/__init__.py
new file mode 100644
index 000000000..09ef7ea81
--- /dev/null
+++ b/cosmos_framework/model/generator/tokenizers/lidar/__init__.py
@@ -0,0 +1,22 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+
+"""LiDAR range-map tokenizer."""
+
+from __future__ import annotations
+
+from typing import Any
+
+__all__ = [
+ "LidarTokenizerV1Interface",
+]
+
+
+def __getattr__(name: str) -> Any:
+ if name == "LidarTokenizerV1Interface":
+ from cosmos_framework.model.generator.tokenizers.lidar.lidar_tokenizer_v1 import (
+ LidarTokenizerV1Interface,
+ )
+
+ return LidarTokenizerV1Interface
+ raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
diff --git a/cosmos_framework/model/generator/tokenizers/lidar/checkpoint.py b/cosmos_framework/model/generator/tokenizers/lidar/checkpoint.py
new file mode 100644
index 000000000..7610d3b20
--- /dev/null
+++ b/cosmos_framework/model/generator/tokenizers/lidar/checkpoint.py
@@ -0,0 +1,318 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+
+"""Checkpoint helpers for the LiDAR TransformerVAE."""
+
+from __future__ import annotations
+
+from collections.abc import Mapping
+from pathlib import Path, PurePosixPath
+from typing import Any
+
+import torch
+
+from cosmos_framework.utils.easy_io import easy_io
+
+# =============================================================================
+# Key remapping
+# =============================================================================
+
+
+def remap_checkpoint_state(
+ state: Mapping[str, Any],
+ model_keys: set[str],
+) -> tuple[dict[str, torch.Tensor], set[str]]:
+ """Select network tensors and strip known training-wrapper prefixes.
+
+ Training checkpoints often nest weights under ``model.network.*`` /
+ ``network.*``. Anything that is not a tensor (optimizer state, iteration
+ counters, EMA copies, unused training-only params) is ignored.
+ """
+ remapped: dict[str, torch.Tensor] = {}
+ ignored: set[str] = set()
+ prefixes = ("model.network.", "network.", "model.vae.", "vae.")
+
+ for original_key, value in state.items():
+ if not isinstance(value, torch.Tensor):
+ ignored.add(original_key)
+ continue
+
+ # Strip the first matching training-wrapper prefix, if any.
+ key = original_key
+ for prefix in prefixes:
+ if key.startswith(prefix):
+ key = key.removeprefix(prefix)
+ break
+
+ if key in model_keys:
+ remapped[key] = value
+ else:
+ ignored.add(original_key)
+
+ return remapped, ignored
+
+
+# =============================================================================
+# Artifact I/O
+# =============================================================================
+
+
+_SAFETENSORS_SUFFIX = ".safetensors"
+
+
+def resolve_artifact_path(path: str, bucket_name: str = "") -> str:
+ """Join ``bucket_name`` for relative paths; leave URIs unchanged."""
+ if "://" in path or not bucket_name:
+ return path
+ return f"s3://{bucket_name}/{path.lstrip('/')}"
+
+
+def load_artifact(path: str, *, backend_args: dict[str, Any] | None = None) -> Any:
+ """Load a local or remote payload (checkpoint / latent stats), torch or safetensors.
+
+ A published tokenizer ships as ``.safetensors``: a flat tensor mapping carrying the weights
+ alongside ``latent_mean`` / ``latent_std``, which the readers below already accept.
+ """
+ path = _resolve_published_artifact(path)
+ if _artifact_suffix(path) == _SAFETENSORS_SUFFIX:
+ return _load_safetensors_artifact(path, backend_args=backend_args)
+ if "://" in path:
+ # Checkpoints embed OmegaConf containers; require full unpickling.
+ return easy_io.load(path, backend_args=backend_args, map_location="cpu", weights_only=False)
+
+ resolved = Path(path).expanduser()
+ if not resolved.is_file():
+ raise FileNotFoundError(f"LiDAR tokenizer artifact not found: {resolved}")
+ return torch.load(resolved, map_location="cpu", weights_only=False)
+
+
+def _artifact_suffix(path: str) -> str:
+ """Suffix of the path component; a presigned URL carries its signature in the query."""
+ return PurePosixPath(path.partition("#")[0].partition("?")[0]).suffix
+
+
+def _resolve_published_artifact(path: str) -> str:
+ """Resolve a published-artifact registry key to a real path; other URIs pass through.
+
+ A public export's ``vae_path`` becomes ``s3://bucket/...`` -- a registry key, not an address.
+ Failures are not caught: with ``check_exists=False`` the only one is a registered artifact
+ that would not download, and that cause is what the caller needs.
+ """
+ from cosmos_framework.utils.checkpoint_db import download_checkpoint_v2
+
+ return download_checkpoint_v2(path, check_exists=False)
+
+
+def _load_safetensors_artifact(path: str, *, backend_args: dict[str, Any] | None = None) -> dict[str, Any]:
+ """Read a ``.safetensors`` tokenizer into a flat ``{name: tensor}`` mapping."""
+ from safetensors.torch import load as load_safetensors_bytes
+ from safetensors.torch import load_file as load_safetensors_file
+
+ if "://" in path:
+ return load_safetensors_bytes(easy_io.get(path, backend_args=backend_args))
+
+ resolved = Path(path).expanduser()
+ if not resolved.is_file():
+ raise FileNotFoundError(f"LiDAR tokenizer artifact not found: {resolved}")
+ return load_safetensors_file(resolved, device="cpu")
+
+
+def parse_latent_stats(stats: Any) -> tuple[torch.Tensor, torch.Tensor]:
+ """Parse ``(mean, std)`` or ``{"mean": ..., "std": ...}`` latent statistics.
+
+ Returns:
+ mean, std: each shaped ``[C]`` (one value per latent channel).
+ """
+ if isinstance(stats, (tuple, list)) and len(stats) == 2:
+ mean, std = stats
+ elif isinstance(stats, Mapping) and "mean" in stats and "std" in stats:
+ mean, std = stats["mean"], stats["std"]
+ else:
+ raise ValueError("Latent stats must be a (mean, std) pair or mapping with mean/std keys")
+
+ return torch.as_tensor(mean), torch.as_tensor(std) # [C], [C]
+
+
+def parse_lidar_stats(
+ stats: Any,
+) -> tuple[torch.Tensor, torch.Tensor, float | None, float | None]:
+ """Parse latent statistics plus optional model-owned metric range bounds."""
+ mean, std = parse_latent_stats(stats)
+ if not isinstance(stats, Mapping):
+ return mean, std, None, None
+ min_range = float(stats["min_range"]) if "min_range" in stats else None
+ max_range = float(stats["max_range"]) if "max_range" in stats else None
+ if (min_range is None) != (max_range is None):
+ raise ValueError("LiDAR stats must contain both min_range and max_range")
+ if min_range is not None and max_range is not None and max_range <= min_range:
+ raise ValueError(f"max_range must exceed min_range, got {min_range=} and {max_range=}")
+ return mean, std, min_range, max_range
+
+
+def parse_lidar_checkpoint_stats(
+ checkpoint: Any,
+) -> tuple[torch.Tensor, torch.Tensor, float | None, float | None]:
+ """Read latent statistics and range bounds embedded by the training wrapper."""
+ if not isinstance(checkpoint, Mapping):
+ raise TypeError(f"Checkpoint must be a mapping, got {type(checkpoint).__name__}")
+ state = checkpoint.get("model", checkpoint)
+ if not isinstance(state, Mapping):
+ raise TypeError(f"Checkpoint model state must be a mapping, got {type(state).__name__}")
+
+ def _get(name: str, *, required: bool = True) -> Any:
+ for key in (name, f"model.{name}"):
+ if key in state:
+ return state[key]
+ if required:
+ raise ValueError(f"LiDAR tokenizer checkpoint does not contain {name!r}")
+ return None
+
+ stats = {
+ "mean": _get("latent_mean"),
+ "std": _get("latent_std"),
+ }
+ min_range = _get("min_range", required=False)
+ max_range = _get("max_range", required=False)
+ if (min_range is None) != (max_range is None):
+ raise ValueError("LiDAR tokenizer checkpoint must contain both min_range and max_range")
+ if min_range is not None:
+ stats.update(min_range=min_range, max_range=max_range)
+ return parse_lidar_stats(stats)
+
+
+# =============================================================================
+# Warm start
+# =============================================================================
+
+
+def resize_azimuth_embedding(embedding: torch.Tensor, width: int) -> torch.Tensor:
+ """Resample a learned range-view embedding ``[1,H,W,C]`` onto ``width`` azimuth bins.
+
+ Azimuth wraps, so the last column neighbours the first and source columns
+ are read modulo the source width. Sample positions sit at bin centers, which
+ keeps both grids spanning the same 360 degrees at any bin count.
+ """
+ source_width = int(embedding.shape[2])
+ if source_width < 1:
+ raise ValueError(f"embedding needs a non-empty azimuth axis, got shape {tuple(embedding.shape)}")
+ if width < 1:
+ raise ValueError(f"target width must be positive, got {width}")
+ if source_width == width:
+ return embedding
+
+ # float64 so that bin centers stay exact for large widths.
+ centers = (torch.arange(width, dtype=torch.float64) + 0.5) * source_width / width - 0.5 # [width]
+ lower = torch.floor(centers)
+ weight = (centers - lower).to(embedding.dtype)[None, None, :, None] # [1,1,width,1]
+ left = lower.to(torch.long) % source_width # [width]
+ right = (lower.to(torch.long) + 1) % source_width # [width]
+ return embedding[:, :, left, :] * (1.0 - weight) + embedding[:, :, right, :] * weight
+
+
+def _differs_only_along_azimuth(source: torch.Tensor, target: torch.Tensor) -> bool:
+ """True when two ``[1,H,W,C]`` range-view embeddings differ only in width."""
+ if source.ndim != 4 or target.ndim != 4:
+ return False
+ return (
+ source.shape[0] == target.shape[0]
+ and source.shape[1] == target.shape[1]
+ and source.shape[3] == target.shape[3]
+ and source.shape[2] != target.shape[2]
+ )
+
+
+def warm_start_network(
+ network: torch.nn.Module,
+ checkpoint_path: str,
+ *,
+ backend_args: dict[str, Any] | None = None,
+) -> dict[str, list[str]]:
+ """Initialize ``network`` from a checkpoint that need not match it exactly.
+
+ For starting a recipe from a published tokenizer whose range view or channel
+ layout differs. Tensors that already agree are copied, range-view positional
+ embeddings are resampled across azimuth, and anything still incompatible is
+ left at its freshly initialized value instead of raising -- unlike
+ :func:`load_model_checkpoint`, which serves published inference weights and
+ must reject any drift.
+
+ Returns the per-category key lists so callers can log what actually landed.
+ """
+ payload = load_artifact(checkpoint_path, backend_args=backend_args)
+ if not isinstance(payload, Mapping):
+ raise TypeError(f"Checkpoint must be a mapping, got {type(payload).__name__}")
+ raw_state = payload.get("model", payload)
+ if not isinstance(raw_state, Mapping):
+ raise TypeError(f"Checkpoint model state must be a mapping, got {type(raw_state).__name__}")
+
+ own_state = network.state_dict()
+ remapped, _ = remap_checkpoint_state(raw_state, set(own_state))
+
+ usable: dict[str, torch.Tensor] = {}
+ resized: list[str] = []
+ skipped: list[str] = []
+ for key, tensor in remapped.items():
+ target = own_state[key]
+ if tensor.shape == target.shape:
+ usable[key] = tensor
+ elif key.endswith("spatial_pe.embedding") and _differs_only_along_azimuth(tensor, target):
+ usable[key] = resize_azimuth_embedding(tensor, int(target.shape[2]))
+ resized.append(key)
+ else:
+ skipped.append(key)
+
+ # ``strict=False`` tolerates absent keys but still raises on a shape
+ # difference, so incompatible tensors have to be dropped above.
+ load_info = network.load_state_dict(usable, strict=False)
+ return {
+ "copied": sorted(set(usable) - set(resized)),
+ "resized": sorted(resized),
+ "skipped": sorted(skipped),
+ "missing": sorted(load_info.missing_keys),
+ }
+
+
+# =============================================================================
+# Model loading
+# =============================================================================
+
+
+def load_model_checkpoint(
+ model: torch.nn.Module,
+ checkpoint_path: str,
+ *,
+ backend_args: dict[str, Any] | None = None,
+ error_prefix: str = "Incompatible LiDAR tokenizer checkpoint",
+) -> Mapping[str, Any]:
+ """Remap and strictly load a published LiDAR tokenizer checkpoint.
+
+ Unknown training-only keys are filtered by ``remap_checkpoint_state``.
+ Missing or shape-mismatched keys required by ``model`` raise.
+ """
+ payload = load_artifact(checkpoint_path, backend_args=backend_args)
+ if not isinstance(payload, Mapping):
+ raise TypeError(f"Checkpoint must be a mapping, got {type(payload).__name__}")
+
+ # Prefer nested ``payload["model"]`` when present (training wrapper layout).
+ raw_state = payload.get("model", payload)
+ if not isinstance(raw_state, Mapping):
+ raise TypeError(f"Checkpoint model state must be a mapping, got {type(raw_state).__name__}")
+
+ own_state = model.state_dict()
+ remapped, _ = remap_checkpoint_state(raw_state, set(own_state))
+
+ missing = set(own_state) - set(remapped)
+ mismatched = {
+ key: (tuple(remapped[key].shape), tuple(own_state[key].shape))
+ for key in remapped.keys() & own_state.keys()
+ if remapped[key].shape != own_state[key].shape
+ }
+ if missing or mismatched:
+ missing_preview = sorted(missing)[:20]
+ raise RuntimeError(
+ f"{error_prefix}: missing={missing_preview}"
+ f"{'...' if len(missing) > len(missing_preview) else ''}, mismatched={mismatched}"
+ )
+
+ model.load_state_dict(remapped, strict=True)
+ return payload
diff --git a/cosmos_framework/model/generator/tokenizers/lidar/checkpoint_test.py b/cosmos_framework/model/generator/tokenizers/lidar/checkpoint_test.py
new file mode 100644
index 000000000..b850d4f46
--- /dev/null
+++ b/cosmos_framework/model/generator/tokenizers/lidar/checkpoint_test.py
@@ -0,0 +1,221 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+
+"""L0 tests for the LiDAR tokenizer checkpoint warm-start helpers."""
+
+from __future__ import annotations
+
+from pathlib import Path
+
+import pytest
+import torch
+import torch.nn as nn
+
+from cosmos_framework.model.generator.tokenizers.lidar.checkpoint import (
+ resize_azimuth_embedding,
+ warm_start_network,
+)
+
+
+class _Net(nn.Module):
+ """Stand-in for the parts of the TransformerVAE the warm start touches.
+
+ ``tokenizer`` stands for the patch embedding whose channel count changes,
+ ``spatial_pe.embedding`` for a range-view embedding whose azimuth width
+ changes, and ``trunk`` for the bulk of the network that transfers as-is.
+ """
+
+ def __init__(self, *, in_channels: int, width: int) -> None:
+ super().__init__()
+ self.tokenizer = nn.Conv2d(in_channels, 4, kernel_size=2, stride=2, bias=False)
+ self.trunk = nn.Linear(4, 4)
+ self.spatial_pe = nn.Module()
+ self.spatial_pe.embedding = nn.Parameter(torch.zeros(1, 2, width, 4))
+
+
+def _write_checkpoint(directory: Path, network: nn.Module) -> str:
+ path = directory / "source.pt"
+ torch.save({"model": network.state_dict()}, path)
+ return str(path)
+
+
+def _randomized(network: nn.Module) -> nn.Module:
+ with torch.no_grad():
+ for parameter in network.parameters():
+ parameter.copy_(torch.randn_like(parameter))
+ return network
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_resize_azimuth_embedding_is_a_no_op_at_equal_width() -> None:
+ embedding = torch.randn(1, 3, 8, 5)
+ assert resize_azimuth_embedding(embedding, 8) is embedding
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_resize_azimuth_embedding_preserves_a_constant_signal() -> None:
+ # A constant is the one signal any sane resampling must reproduce exactly,
+ # wrap-around included.
+ embedding = torch.full((1, 2, 16, 3), 0.75)
+ resized = resize_azimuth_embedding(embedding, 29)
+ assert resized.shape == (1, 2, 29, 3)
+ torch.testing.assert_close(resized, torch.full((1, 2, 29, 3), 0.75))
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_resize_azimuth_embedding_hits_source_bins_on_an_odd_upsample() -> None:
+ # Tripling places every third target center exactly on a source center
+ # (target j = 1 + 3k maps to source bin k), so those columns must come back
+ # untouched. Even ratios never coincide under the bin-center convention.
+ embedding = torch.randn(1, 1, 6, 2)
+ resized = resize_azimuth_embedding(embedding, 18)
+ torch.testing.assert_close(resized[:, :, 1::3, :], embedding)
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_resize_azimuth_embedding_wraps_rather_than_clamping() -> None:
+ # A ramp whose last bin is far from its first: clamping at the edge would
+ # extend the ramp, whereas wrapping blends the two ends back together.
+ embedding = torch.zeros(1, 1, 4, 1)
+ embedding[0, 0, :, 0] = torch.tensor([0.0, 1.0, 2.0, 3.0])
+ resized = resize_azimuth_embedding(embedding, 8)
+ # Target center 7 sits at source 3.25, a quarter of the way from bin 3 back
+ # round to bin 0, so 3 * 0.75 + 0 * 0.25.
+ torch.testing.assert_close(resized[0, 0, 7, 0], torch.tensor(2.25))
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_resize_azimuth_embedding_rejects_a_nonpositive_width() -> None:
+ with pytest.raises(ValueError, match="target width must be positive"):
+ resize_azimuth_embedding(torch.randn(1, 2, 4, 3), 0)
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_warm_start_applies_each_category(tmp_path: Path) -> None:
+ source = _randomized(_Net(in_channels=2, width=8))
+ checkpoint = _write_checkpoint(tmp_path, source)
+ # Wider azimuth plus an extra input channel: the V0 -> V1.2 shift in
+ # miniature.
+ target = _Net(in_channels=3, width=14)
+ stale_tokenizer = target.tokenizer.weight.detach().clone()
+
+ report = warm_start_network(target, checkpoint, backend_args=None)
+
+ # Shapes already agree, so the trunk transfers verbatim.
+ assert report["copied"] == ["trunk.bias", "trunk.weight"]
+ torch.testing.assert_close(target.trunk.weight, source.trunk.weight)
+ # Differs only along azimuth, so it is resampled rather than dropped.
+ assert report["resized"] == ["spatial_pe.embedding"]
+ assert target.spatial_pe.embedding.shape == (1, 2, 14, 4)
+ # Gained an input channel, so it keeps its fresh initialization.
+ assert report["skipped"] == ["tokenizer.weight"]
+ assert report["missing"] == ["tokenizer.weight"]
+ torch.testing.assert_close(target.tokenizer.weight, stale_tokenizer)
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_warm_start_transfers_everything_when_shapes_already_agree(tmp_path: Path) -> None:
+ source = _randomized(_Net(in_channels=3, width=14))
+ checkpoint = _write_checkpoint(tmp_path, source)
+ target = _Net(in_channels=3, width=14)
+
+ report = warm_start_network(target, checkpoint, backend_args=None)
+
+ assert report["resized"] == []
+ assert report["skipped"] == []
+ assert report["missing"] == []
+ for key, tensor in source.state_dict().items():
+ torch.testing.assert_close(target.state_dict()[key], tensor)
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_warm_start_accepts_a_bare_state_dict_without_a_model_key(tmp_path: Path) -> None:
+ # V0 stores a training wrapper with a "model" key; published weights may be
+ # a bare state dict.
+ source = _randomized(_Net(in_channels=3, width=14))
+ path = tmp_path / "bare.pt"
+ torch.save(source.state_dict(), path)
+ target = _Net(in_channels=3, width=14)
+
+ report = warm_start_network(target, str(path), backend_args=None)
+
+ assert report["missing"] == []
+ torch.testing.assert_close(target.trunk.weight, source.trunk.weight)
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_warm_start_leaves_the_network_usable_when_nothing_matches(tmp_path: Path) -> None:
+ path = tmp_path / "unrelated.pt"
+ torch.save({"model": {"totally.unrelated": torch.randn(3, 3)}}, path)
+ target = _Net(in_channels=3, width=14)
+
+ report = warm_start_network(target, str(path), backend_args=None)
+
+ # Reported as absent rather than raising, so a bad path cannot silently
+ # look like a successful warm start.
+ assert report["copied"] == []
+ assert report["resized"] == []
+ assert "trunk.weight" in report["missing"]
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_warm_start_rejects_a_payload_that_is_not_a_mapping(tmp_path: Path) -> None:
+ path = tmp_path / "bad.pt"
+ torch.save([1, 2, 3], path)
+ with pytest.raises(TypeError, match="Checkpoint must be a mapping"):
+ warm_start_network(_Net(in_channels=3, width=14), str(path), backend_args=None)
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_query_string_url_still_routes_to_the_safetensors_loader(monkeypatch) -> None:
+ """Assert the routing, not just the helper.
+
+ Testing ``_artifact_suffix`` alone leaves the decision that uses it uncovered: reverting the
+ call site to ``path.endswith(...)`` would send a presigned URL down the pickle path with
+ every helper test still green.
+ """
+ from cosmos_framework.model.generator.tokenizers.lidar import checkpoint as ckpt
+
+ taken: list[str] = []
+ monkeypatch.setattr(ckpt, "_resolve_published_artifact", lambda path: path)
+ monkeypatch.setattr(
+ ckpt, "_load_safetensors_artifact", lambda path, backend_args=None: taken.append("safetensors") or {}
+ )
+ monkeypatch.setattr(ckpt.easy_io, "load", lambda *a, **k: taken.append("pickle") or {})
+
+ ckpt.load_artifact("https://host/a/model.safetensors?X-Amz-Signature=abc")
+ assert taken == ["safetensors"], taken
+
+ taken.clear()
+ ckpt.load_artifact("s3://bucket/checkpoints/iter_000030000.pt")
+ assert taken == ["pickle"], taken
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_registered_artifact_download_failure_keeps_its_cause(monkeypatch) -> None:
+ """A failed download must surface, not degrade into a misleading storage error.
+
+ Catching it would hand back the ``s3://bucket/...`` registry alias, which the loader then
+ tries to read as a real object-store address -- reporting a deserialization failure instead
+ of the missing credential or network error behind it.
+ """
+ from cosmos_framework.model.generator.tokenizers.lidar import checkpoint as ckpt
+
+ def _boom(uri, *, check_exists=True):
+ raise RuntimeError("credential file not found")
+
+ monkeypatch.setattr("cosmos_framework.utils.checkpoint_db.download_checkpoint_v2", _boom)
+ with pytest.raises(RuntimeError, match="credential file not found"):
+ ckpt._resolve_published_artifact("s3://bucket/pretrained/tokenizers/lidar/x.safetensors")
diff --git a/cosmos_framework/model/generator/tokenizers/lidar/dtypes.py b/cosmos_framework/model/generator/tokenizers/lidar/dtypes.py
new file mode 100644
index 000000000..cb00c4c08
--- /dev/null
+++ b/cosmos_framework/model/generator/tokenizers/lidar/dtypes.py
@@ -0,0 +1,39 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+
+"""Dtype coercion for tokenizer configs that have been through a config round trip."""
+
+import torch
+
+__all__ = ["as_torch_dtype"]
+
+
+def as_torch_dtype(dtype: str | torch.dtype) -> torch.dtype:
+ """Return ``dtype`` as a ``torch.dtype``, accepting the serialized string form.
+
+ The LiDAR tokenizer configs pin a real ``torch.dtype`` (``_LIDAR_INFERENCE_DTYPE``
+ in configs/base/defaults/tokenizer.py). Serializing a model config to ``config.json``
+ turns that into the string ``"float32"``, and nothing on the way back in turns it
+ into a dtype again -- so an exported checkpoint reaches ``Module.to(dtype=...)``
+ with a ``str`` and dies with::
+
+ TypeError: to() received an invalid combination of arguments -
+ got (dtype=str, device=torch.device, )
+
+ ``device`` never had this problem because the constructors already normalize it
+ through ``torch.device(...)``, which accepts a string. This is the dtype equivalent.
+
+ Note ``torch.dtype("float32")`` does NOT work -- it raises
+ ``TypeError: cannot create 'torch.dtype' instances`` -- so the name is resolved as
+ an attribute instead. Both ``"float32"`` and ``"torch.float32"`` are accepted,
+ because serializers differ on whether they keep the module prefix.
+ """
+ if isinstance(dtype, torch.dtype):
+ return dtype
+ if not isinstance(dtype, str):
+ raise TypeError(f"Expected a torch.dtype or its name, got {type(dtype).__name__}: {dtype!r}")
+
+ resolved = getattr(torch, dtype.split(".")[-1], None)
+ if not isinstance(resolved, torch.dtype):
+ raise ValueError(f"Not a torch dtype name: {dtype!r}")
+ return resolved
diff --git a/cosmos_framework/model/generator/tokenizers/lidar/dtypes_test.py b/cosmos_framework/model/generator/tokenizers/lidar/dtypes_test.py
new file mode 100644
index 000000000..b1d98d29c
--- /dev/null
+++ b/cosmos_framework/model/generator/tokenizers/lidar/dtypes_test.py
@@ -0,0 +1,41 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+
+import pytest
+import torch
+
+from cosmos_framework.model.generator.tokenizers.lidar.dtypes import as_torch_dtype
+
+pytestmark = [pytest.mark.L0, pytest.mark.CPU]
+
+
+def test_passes_through_a_real_dtype() -> None:
+ assert as_torch_dtype(torch.bfloat16) is torch.bfloat16
+
+
+@pytest.mark.parametrize("name", ["float32", "torch.float32"])
+def test_accepts_the_serialized_string_form(name: str) -> None:
+ """An exported config.json carries the dtype as a string, with or without the prefix."""
+ assert as_torch_dtype(name) is torch.float32
+
+
+def test_rejects_a_name_that_is_not_a_dtype() -> None:
+ """torch.nn is a real attribute of torch but not a dtype; it must not slip through."""
+ with pytest.raises(ValueError, match="Not a torch dtype name"):
+ as_torch_dtype("nn")
+
+
+def test_rejects_a_non_string_non_dtype() -> None:
+ with pytest.raises(TypeError):
+ as_torch_dtype(32) # type: ignore[arg-type]
+
+
+def test_torch_dtype_is_not_constructible_from_a_string() -> None:
+ """Pins why this helper resolves the name by attribute lookup.
+
+ cosmos_framework/model/generator/diffusion/rectified_flow.py does
+ `torch.dtype(dtype) if isinstance(dtype, str) else dtype`, which cannot work --
+ this asserts the reason so the idiom is not copied here.
+ """
+ with pytest.raises(TypeError, match="cannot create 'torch.dtype' instances"):
+ torch.dtype("float32") # type: ignore[call-arg]
diff --git a/cosmos_framework/model/generator/tokenizers/lidar/lidar_tokenizer_v1.py b/cosmos_framework/model/generator/tokenizers/lidar/lidar_tokenizer_v1.py
new file mode 100644
index 000000000..f72edcd81
--- /dev/null
+++ b/cosmos_framework/model/generator/tokenizers/lidar/lidar_tokenizer_v1.py
@@ -0,0 +1,554 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+
+"""LiDAR TransformerVAE inference interface (range + intensity + mask)."""
+
+from __future__ import annotations
+
+import warnings
+from collections.abc import Mapping, Sequence
+from typing import Any
+
+import torch
+
+from cosmos_framework.model.generator.tokenizers.interface import VideoTokenizerInterface
+from cosmos_framework.model.generator.tokenizers.lidar.checkpoint import (
+ load_model_checkpoint,
+ parse_lidar_checkpoint_stats,
+ resolve_artifact_path,
+)
+from cosmos_framework.model.generator.tokenizers.lidar.dtypes import as_torch_dtype
+from cosmos_framework.model.generator.tokenizers.lidar.network.transformer_vae import TransformerVAE
+from cosmos_framework.model.generator.tokenizers.lidar.normalization import (
+ metric_lidar_to_network,
+ network_lidar_to_metric_clip,
+)
+from cosmos_framework.model.generator.tokenizers.lidar.paths import (
+ is_remote_uri,
+ resolve_credential_path,
+)
+from cosmos_framework.model.generator.tokenizers.lidar.postprocessing import (
+ resolved_validity_channel,
+ validate_validity_threshold,
+ validity_probability,
+)
+from cosmos_framework.model.generator.tokenizers.lidar.preprocessing import (
+ INVALID_NORMALIZED_VALUE,
+ MAX_RANGE_METERS,
+ MIN_RANGE_METERS,
+)
+from cosmos_framework.model.generator.tokenizers.lidar.range_projection import LidarRangeProjectionConfig
+
+# =============================================================================
+# Defaults (3-channel range + intensity + mask tokenizer)
+# =============================================================================
+
+DEFAULT_EXPERIMENT_NAME = "lidar_tokenizer_v1"
+DEFAULT_INPUT_RESOLUTION = (128, 3600) # (H, W) native range-map width
+DEFAULT_LATENT_SPATIAL = (16, 225) # H/8, W/16
+DEFAULT_SPATIAL_COMPRESSION = (8, 16) # (H, W)
+DEFAULT_IN_CHANNELS = 3 # range, intensity, mask
+DEFAULT_OUT_CHANNELS = 3 # range, intensity, mask (mask channel trained as logits)
+DEFAULT_LATENT_CH = 128
+
+# Spatial: patch 2x4 + 2 merges -> 8x16. Temporal: no downsample -> 1x.
+DEFAULT_NETWORK_CONFIG: dict[str, Any] = {
+ "resolution": list(DEFAULT_INPUT_RESOLUTION),
+ "in_channels": DEFAULT_IN_CHANNELS,
+ "out_channels": DEFAULT_OUT_CHANNELS,
+ "z_dim": DEFAULT_LATENT_CH,
+ "base_channels": 128,
+ "patch_size": [2, 4],
+ "window_size": [5, 45],
+ "depths": [3, 3, 3],
+ "num_heads": [4, 4, 8],
+ "dilation": [1, 1, 1],
+ "temporal_downsample": [False, False],
+ "temporal_upsample": [False, False],
+ "mlp_ratio": 3.0,
+ "mapping_depth": 2,
+ "positional_embedding": "learnable_embedding",
+ "formulation": "VAE",
+ # Mask is already channel 2 of the input tensor (not derived inside encode).
+ # Reconstructed as output channel 2 (logits); no extra validity head.
+ "mask_as_input": False,
+ "bottleneck_3d": True,
+ "bottleneck_3d_max_t": 32,
+ "bottleneck_3d_causal_time": True,
+ "bottleneck_3d_rope": True,
+}
+
+
+def infer_lidar_compression(network_config: Mapping[str, Any]) -> tuple[tuple[int, int], int]:
+ """Infer ``((height, width), time)`` compression from a TransformerVAE config."""
+ patch_size = tuple(int(value) for value in network_config["patch_size"])
+ depths = tuple(network_config["depths"])
+ temporal_downsample = tuple(bool(value) for value in network_config["temporal_downsample"])
+ if len(patch_size) != 2 or any(value <= 0 for value in patch_size):
+ raise ValueError(f"patch_size must contain two positive values, got {patch_size}")
+ if not depths:
+ raise ValueError("depths must contain at least one level")
+ expected_merges = len(depths) - 1
+ if len(temporal_downsample) != expected_merges:
+ raise ValueError(
+ "temporal_downsample must have one entry per encoder merge, got "
+ f"{len(temporal_downsample)} entries for {len(depths)} levels"
+ )
+ spatial_merge_factor = 2**expected_merges
+ spatial = (patch_size[0] * spatial_merge_factor, patch_size[1] * spatial_merge_factor)
+ temporal = 2 ** sum(temporal_downsample)
+ return spatial, temporal
+
+
+# =============================================================================
+# Interface
+# =============================================================================
+
+
+class LidarTokenizerV1Interface(VideoTokenizerInterface):
+ """V1 LiDAR range+intensity+mask tokenizer (H/8 × W/16 spatial / 1× temporal).
+
+ Input: ``[B, 3, T, 128, 3600]`` (metric range, unit intensity, mask).
+ Latent: ``[B, 128, T, 16, 225]`` (H/8, W/16).
+ Decode: ``[B, 3, T, 128, 3600]`` (metric range, unit intensity, mask).
+
+ The decoded mask channel matches the input convention: ``{0, 1}`` once the
+ validity cut is applied, and the underlying probability when it is not. The
+ raw logits never leave this class.
+ """
+
+ def __init__(
+ self,
+ vae_path: str | None = None,
+ *,
+ object_store_credential_path_pretrained: str | None = None,
+ bucket_name: str = "",
+ device: str | torch.device = "cuda",
+ dtype: str | torch.dtype = torch.float32,
+ sample_posterior: bool = False,
+ apply_validity_mask: bool = True,
+ network_config: Mapping[str, Any] | None = None,
+ load_checkpoint: bool = True,
+ spatial_compression_factor: int | None = None,
+ spatial_compression: Sequence[int] | None = None,
+ temporal_compression_factor: int | None = None,
+ range_projection: LidarRangeProjectionConfig | Mapping[str, Any] | None = None,
+ streaming_chunk_frames: int | None = None,
+ streaming_context_frames: int | None = None,
+ ) -> None:
+ # Resolve object-store credentials when loading remote artifacts.
+ needs_remote = load_checkpoint and vae_path is not None and is_remote_uri(vae_path)
+ if object_store_credential_path_pretrained is None and needs_remote:
+ object_store_credential_path_pretrained = resolve_credential_path()
+ super().__init__(object_store_credential_path_pretrained)
+
+ self.device = torch.device(device)
+ self.dtype = as_torch_dtype(dtype)
+ self.sample_posterior = sample_posterior
+ self.apply_validity_mask = apply_validity_mask
+ if streaming_chunk_frames is not None and streaming_chunk_frames < 1:
+ raise ValueError(f"streaming_chunk_frames must be positive, got {streaming_chunk_frames}")
+ if streaming_context_frames is not None and streaming_context_frames < 1:
+ raise ValueError(f"streaming_context_frames must be positive, got {streaming_context_frames}")
+ if (
+ streaming_chunk_frames is not None
+ and streaming_context_frames is not None
+ and streaming_context_frames < streaming_chunk_frames
+ ):
+ raise ValueError("streaming_context_frames must be at least streaming_chunk_frames")
+ self.streaming_chunk_frames = streaming_chunk_frames
+ self.streaming_context_frames = streaming_context_frames
+
+ # Derive compression from the effective architecture. Explicit metadata
+ # remains available to config interpolation, but cannot silently drift.
+ config = dict(DEFAULT_NETWORK_CONFIG)
+ if network_config is not None:
+ config.update(network_config)
+ inferred_spatial, inferred_temporal = infer_lidar_compression(config)
+ input_resolution = tuple(int(value) for value in config["resolution"])
+ if len(input_resolution) != 2 or any(value <= 0 for value in input_resolution):
+ raise ValueError(f"resolution must contain two positive values, got {input_resolution}")
+ if any(size % factor for size, factor in zip(input_resolution, inferred_spatial, strict=True)):
+ raise ValueError(
+ f"resolution {input_resolution} must be divisible by spatial compression {inferred_spatial}"
+ )
+ self._input_resolution = input_resolution
+ self._latent_spatial = tuple(
+ size // factor for size, factor in zip(input_resolution, inferred_spatial, strict=True)
+ )
+ spatial_compression = inferred_spatial if spatial_compression is None else tuple(spatial_compression)
+ spatial_compression_factor = (
+ inferred_spatial[0] if spatial_compression_factor is None else int(spatial_compression_factor)
+ )
+ temporal_compression_factor = (
+ inferred_temporal if temporal_compression_factor is None else int(temporal_compression_factor)
+ )
+ if (
+ spatial_compression_factor != inferred_spatial[0]
+ or spatial_compression != inferred_spatial
+ or temporal_compression_factor != inferred_temporal
+ ):
+ raise ValueError(
+ "Compression metadata does not match the network architecture: "
+ f"expected factor={inferred_spatial[0]}, spatial={inferred_spatial}, temporal={inferred_temporal}; "
+ f"got factor={spatial_compression_factor}, spatial={spatial_compression}, "
+ f"temporal={temporal_compression_factor}"
+ )
+
+ # VideoTokenizerInterface / DiT bookkeeping.
+ self._causal = True
+ self._spatial_compression_factor = spatial_compression_factor
+ self._spatial_compression = spatial_compression
+ self._temporal_compression_factor = temporal_compression_factor
+ self._pixel_chunk_duration = 9
+
+ # Build network (optionally override architecture for experiments).
+ self.model = TransformerVAE(**config)
+
+ # Checkpoint is optional so smoke tests can run with random weights.
+ checkpoint_payload: Mapping[str, Any] | None = None
+ if load_checkpoint:
+ if vae_path is None:
+ raise ValueError("vae_path is required when load_checkpoint=True")
+ checkpoint_payload = load_model_checkpoint(
+ self.model,
+ resolve_artifact_path(vae_path, bucket_name),
+ backend_args=self.backend_args,
+ error_prefix="Incompatible LiDAR tokenizer checkpoint",
+ )
+ self.model.eval().requires_grad_(False)
+ self.model.to(device=self.device, dtype=self.dtype)
+
+ # Latent mean/std ride in the checkpoint, so the affine and the weights
+ # it belongs to can never be paired wrongly. Without a checkpoint the
+ # affine is the identity, for untrained / smoke-test runs.
+ if load_checkpoint:
+ assert checkpoint_payload is not None
+ latent_mean, latent_std, stored_min_range, stored_max_range = parse_lidar_checkpoint_stats(
+ checkpoint_payload
+ )
+ else:
+ latent_mean = torch.zeros(self.latent_ch)
+ latent_std = torch.ones(self.latent_ch)
+ stored_min_range = stored_max_range = None
+
+ if isinstance(range_projection, Mapping):
+ range_projection = LidarRangeProjectionConfig.from_dict(range_projection)
+ if range_projection is None:
+ min_range = stored_min_range
+ max_range = stored_max_range
+ else:
+ min_range = range_projection.min_range_m
+ max_range = range_projection.max_range_m
+ if stored_min_range is not None and stored_min_range != min_range:
+ raise ValueError(
+ f"Configured min_range_m {min_range} does not match checkpoint value {stored_min_range}"
+ )
+ if stored_max_range is not None and stored_max_range != max_range:
+ raise ValueError(
+ f"Configured max_range_m {max_range} does not match checkpoint value {stored_max_range}"
+ )
+ if min_range is None or max_range is None:
+ if load_checkpoint:
+ warnings.warn(
+ f"LiDAR range metadata is absent; assuming the default "
+ f"[{MIN_RANGE_METERS}, {MAX_RANGE_METERS}] m span. "
+ "Pass range_projection for checkpoints trained on another span.",
+ stacklevel=2,
+ )
+ min_range = MIN_RANGE_METERS if min_range is None else min_range
+ max_range = MAX_RANGE_METERS if max_range is None else max_range
+ if range_projection is None:
+ range_projection = LidarRangeProjectionConfig(
+ semantic_width=input_resolution[1],
+ model_width=input_resolution[1],
+ min_range_m=min_range,
+ max_range_m=max_range,
+ )
+ if (range_projection.native_height, range_projection.model_width) != input_resolution:
+ raise ValueError(
+ "Range projection dimensions do not match the tokenizer architecture: "
+ f"projection={(range_projection.native_height, range_projection.model_width)}, "
+ f"network={input_resolution}"
+ )
+ self.range_projection = range_projection
+ self.min_range = range_projection.min_range_m
+ self.max_range = range_projection.max_range_m
+ self.validity_threshold = validate_validity_threshold(range_projection.validity_threshold)
+
+ if latent_mean.numel() != self.latent_ch or latent_std.numel() != self.latent_ch:
+ raise ValueError(
+ "Latent statistics must have one value per channel: "
+ f"expected {self.latent_ch}, got mean={latent_mean.numel()} and std={latent_std.numel()}"
+ )
+ if torch.any(latent_std <= 0):
+ raise ValueError("Latent standard deviations must be positive")
+
+ stats_shape = (1, self.latent_ch, 1, 1, 1)
+ self.latent_mean = latent_mean.to(device=self.device, dtype=self.dtype).reshape(stats_shape)
+ self.latent_std = latent_std.to(device=self.device, dtype=self.dtype).reshape(stats_shape)
+
+ # -------------------------------------------------------------------------
+ # Encode / decode
+ # -------------------------------------------------------------------------
+
+ def reset_dtype(self) -> None:
+ self.model.to(device=self.device, dtype=self.dtype)
+ self.latent_mean = self.latent_mean.to(device=self.device, dtype=self.dtype)
+ self.latent_std = self.latent_std.to(device=self.device, dtype=self.dtype)
+
+ @torch.inference_mode()
+ def encode(self, state: torch.Tensor) -> torch.Tensor:
+ """Encode metric range, unit intensity, and mask to normalized latents."""
+ if self.streaming_chunk_frames is not None:
+ return self.encode_streaming(state, chunk_frames=self.streaming_chunk_frames)
+ state = state.to(device=self.device, dtype=self.dtype)
+ normalized, _ = metric_lidar_to_network(state, min_range=self.min_range, max_range=self.max_range)
+ return self.encode_normalized(normalized)
+
+ @torch.inference_mode()
+ def encode_normalized(self, state: torch.Tensor) -> torch.Tensor:
+ """Encode legacy normalized ``[B,3,T,H,W]`` tensors."""
+ if (
+ state.ndim != 5
+ or state.shape[1] != DEFAULT_IN_CHANNELS
+ or tuple(state.shape[-2:]) != self._input_resolution
+ ):
+ raise ValueError(
+ f"Expected [B,{DEFAULT_IN_CHANNELS},T,"
+ f"{self._input_resolution[0]},{self._input_resolution[1]}], "
+ f"got {tuple(state.shape)}"
+ )
+
+ state = state.to(device=self.device, dtype=self.dtype)
+ sample, (posterior_mean, _) = self.model.encode(state)
+ latent = sample if self.sample_posterior else posterior_mean
+ latent = latent.to(dtype=self.dtype)
+
+ return (latent - self.latent_mean) / self.latent_std
+
+ @torch.inference_mode()
+ def encode_streaming(self, state: torch.Tensor, *, chunk_frames: int | None = None) -> torch.Tensor:
+ """Chunked metric-input encode for long videos."""
+ state = state.to(device=self.device, dtype=self.dtype)
+ normalized, _ = metric_lidar_to_network(state, min_range=self.min_range, max_range=self.max_range)
+ return self.encode_streaming_normalized(normalized, chunk_frames=chunk_frames)
+
+ @torch.inference_mode()
+ def encode_streaming_normalized(self, state: torch.Tensor, *, chunk_frames: int | None = None) -> torch.Tensor:
+ """Chunked encode for legacy normalized tensors."""
+ if (
+ state.ndim != 5
+ or state.shape[1] != DEFAULT_IN_CHANNELS
+ or tuple(state.shape[-2:]) != self._input_resolution
+ ):
+ raise ValueError(
+ f"Expected [B,{DEFAULT_IN_CHANNELS},T,"
+ f"{self._input_resolution[0]},{self._input_resolution[1]}], "
+ f"got {tuple(state.shape)}"
+ )
+ chunk = self.pixel_chunk_duration if chunk_frames is None else int(chunk_frames)
+ state = state.to(device=self.device, dtype=self.dtype)
+ latent = self.model.encode_streaming(
+ state,
+ chunk_frames=chunk,
+ context_frames=self.streaming_context_frames,
+ sample_posterior=self.sample_posterior,
+ )
+ return (latent.to(dtype=self.dtype) - self.latent_mean) / self.latent_std
+
+ def _resolve_validity_threshold(self, override: float | None) -> float:
+ """Per-call probability cut, falling back to the configured one."""
+ return self.validity_threshold if override is None else validate_validity_threshold(override)
+
+ def _split_decode_output(self, output: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
+ """Split decoder output into data(+mask logits) and sigmoid validity."""
+ if output.shape[1] < 3:
+ raise ValueError(f"Expected >=3 decode channels, got {tuple(output.shape)}")
+ validity = validity_probability(output[:, 2:3].to(dtype=output.dtype))
+ return output, validity
+
+ @torch.inference_mode()
+ def decode_normalized(
+ self,
+ latent: torch.Tensor,
+ *,
+ return_validity: bool = False,
+ apply_mask: bool | None = None,
+ validity_threshold: float | None = None,
+ ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor | None]:
+ """Decode to the legacy normalized range/intensity representation."""
+ if latent.ndim != 5 or latent.shape[1] != self.latent_ch:
+ raise ValueError(f"Expected [B,{self.latent_ch},T,H,W], got {tuple(latent.shape)}")
+
+ latent = latent.to(device=self.device, dtype=self.dtype)
+ denormalized = latent * self.latent_std + self.latent_mean
+
+ # Align decode crop target with this latent's T (1x temporal compression).
+ self.model._input_T = int(latent.shape[2])
+ output = self.model.decode(denormalized, return_validity=False)
+ output, validity = self._split_decode_output(output)
+
+ should_mask = self.apply_validity_mask if apply_mask is None else apply_mask
+ threshold = self._resolve_validity_threshold(validity_threshold)
+ output = output.clone()
+ if should_mask:
+ output[:, :2] = output[:, :2].masked_fill(validity < threshold, INVALID_NORMALIZED_VALUE)
+ output[:, 2:3] = resolved_validity_channel(validity, should_mask=should_mask, threshold=threshold)
+
+ if return_validity:
+ return output, validity
+ return output
+
+ @torch.inference_mode()
+ def decode(
+ self,
+ latent: torch.Tensor,
+ *,
+ return_validity: bool = False,
+ apply_mask: bool | None = None,
+ validity_threshold: float | None = None,
+ ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor | None]:
+ """Decode to metric range, unit intensity, and resolved mask."""
+ if self.streaming_chunk_frames is not None:
+ return self.decode_streaming(
+ latent,
+ chunk_frames=self.streaming_chunk_frames,
+ return_validity=return_validity,
+ apply_mask=apply_mask,
+ validity_threshold=validity_threshold,
+ )
+ normalized, validity = self.decode_normalized(latent, return_validity=True, apply_mask=False)
+ assert validity is not None
+ should_mask = self.apply_validity_mask if apply_mask is None else apply_mask
+ threshold = self._resolve_validity_threshold(validity_threshold)
+ output = network_lidar_to_metric_clip(
+ normalized,
+ validity,
+ min_range=self.min_range,
+ max_range=self.max_range,
+ apply_validity_mask=should_mask,
+ validity_threshold=threshold,
+ )
+ if return_validity:
+ return output, validity
+ return output
+
+ @torch.inference_mode()
+ def decode_streaming_normalized(
+ self,
+ latent: torch.Tensor,
+ *,
+ chunk_frames: int | None = None,
+ return_validity: bool = False,
+ apply_mask: bool | None = None,
+ validity_threshold: float | None = None,
+ ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor | None]:
+ """Chunked decode to the legacy normalized representation."""
+ if latent.ndim != 5 or latent.shape[1] != self.latent_ch:
+ raise ValueError(f"Expected [B,{self.latent_ch},T,H,W], got {tuple(latent.shape)}")
+ chunk = self.latent_chunk_duration if chunk_frames is None else int(chunk_frames)
+ latent = latent.to(device=self.device, dtype=self.dtype)
+ denormalized = latent * self.latent_std + self.latent_mean
+ self.model._input_T = int(latent.shape[2])
+ output = self.model.decode_streaming(
+ denormalized,
+ chunk_frames=chunk,
+ context_frames=self.streaming_context_frames,
+ return_validity=False,
+ )
+ output, validity = self._split_decode_output(output)
+ should_mask = self.apply_validity_mask if apply_mask is None else apply_mask
+ threshold = self._resolve_validity_threshold(validity_threshold)
+ output = output.clone()
+ if should_mask:
+ output[:, :2] = output[:, :2].masked_fill(validity < threshold, INVALID_NORMALIZED_VALUE)
+ output[:, 2:3] = resolved_validity_channel(validity, should_mask=should_mask, threshold=threshold)
+ if return_validity:
+ return output, validity
+ return output
+
+ @torch.inference_mode()
+ def decode_streaming(
+ self,
+ latent: torch.Tensor,
+ *,
+ chunk_frames: int | None = None,
+ return_validity: bool = False,
+ apply_mask: bool | None = None,
+ validity_threshold: float | None = None,
+ ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor | None]:
+ """Chunked decode to metric range, unit intensity, and resolved mask."""
+ normalized, validity = self.decode_streaming_normalized(
+ latent,
+ chunk_frames=chunk_frames,
+ return_validity=True,
+ apply_mask=False,
+ )
+ assert validity is not None
+ should_mask = self.apply_validity_mask if apply_mask is None else apply_mask
+ threshold = self._resolve_validity_threshold(validity_threshold)
+ output = network_lidar_to_metric_clip(
+ normalized,
+ validity,
+ min_range=self.min_range,
+ max_range=self.max_range,
+ apply_validity_mask=should_mask,
+ validity_threshold=threshold,
+ )
+ if return_validity:
+ return output, validity
+ return output
+
+ # -------------------------------------------------------------------------
+ # Frame / compression bookkeeping
+ # -------------------------------------------------------------------------
+
+ def get_latent_num_frames(self, num_pixel_frames: int) -> int:
+ return num_pixel_frames
+
+ def get_pixel_num_frames(self, num_latent_frames: int, **kwargs: Any) -> int:
+ del kwargs
+ return num_latent_frames
+
+ @property
+ def spatial_compression_factor(self) -> int:
+ """Height compression (interface API is a single int). Prefer :meth:`spatial_compression`."""
+ return self._spatial_compression_factor
+
+ @property
+ def spatial_compression(self) -> tuple[int, int]:
+ """``(H, W)`` spatial compression factors (8, 16)."""
+ return self._spatial_compression
+
+ @property
+ def temporal_compression_factor(self) -> int:
+ return self._temporal_compression_factor
+
+ @property
+ def spatial_resolution(self) -> int:
+ return self._input_resolution[0]
+
+ @property
+ def input_resolution(self) -> tuple[int, int]:
+ return self._input_resolution
+
+ @property
+ def latent_spatial(self) -> tuple[int, int]:
+ return self._latent_spatial
+
+ @property
+ def pixel_chunk_duration(self) -> int:
+ return self._pixel_chunk_duration
+
+ @property
+ def latent_chunk_duration(self) -> int:
+ return self._pixel_chunk_duration
+
+ @property
+ def latent_ch(self) -> int:
+ return DEFAULT_LATENT_CH
+
+ @property
+ def name(self) -> str:
+ return DEFAULT_EXPERIMENT_NAME
diff --git a/cosmos_framework/model/generator/tokenizers/lidar/network/__init__.py b/cosmos_framework/model/generator/tokenizers/lidar/network/__init__.py
new file mode 100644
index 000000000..1697f3eaf
--- /dev/null
+++ b/cosmos_framework/model/generator/tokenizers/lidar/network/__init__.py
@@ -0,0 +1,18 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+
+"""Checkpoint-compatible LiDAR TransformerVAE network."""
+
+from __future__ import annotations
+
+from typing import Any
+
+__all__ = ["TransformerVAE"]
+
+
+def __getattr__(name: str) -> Any:
+ if name != "TransformerVAE":
+ raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
+ from cosmos_framework.model.generator.tokenizers.lidar.network.transformer_vae import TransformerVAE
+
+ return TransformerVAE
diff --git a/cosmos_framework/model/generator/tokenizers/lidar/network/distributions.py b/cosmos_framework/model/generator/tokenizers/lidar/network/distributions.py
new file mode 100644
index 000000000..77a6264c0
--- /dev/null
+++ b/cosmos_framework/model/generator/tokenizers/lidar/network/distributions.py
@@ -0,0 +1,49 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+
+"""The distribution modes to use for continuous image tokenizers."""
+
+from enum import Enum
+
+import torch
+
+
+class IdentityDistribution(torch.nn.Module):
+ def __init__(self) -> None:
+ super().__init__()
+
+ def forward(
+ self, parameters: torch.Tensor, split: bool = False
+ ) -> tuple[torch.Tensor, tuple[torch.Tensor, torch.Tensor]]: # parameters: [B,C,...], returns sample: [B,C,...]
+ if not split:
+ zero = parameters.new_zeros(1) # [1]
+ return parameters, (zero, zero)
+ mean, logvar = parameters.chunk(2, dim=1) # [B,C/2,...], [B,C/2,...]
+ mean, logvar = mean.contiguous(), logvar.contiguous() # [B,C/2,...], [B,C/2,...]
+ zero = parameters.new_zeros(1) # [1]
+ return mean, (zero, zero)
+
+
+class GaussianDistribution(torch.nn.Module):
+ def __init__(self, min_logvar: float = -30.0, max_logvar: float = 20.0) -> None:
+ super().__init__()
+ self.min_logvar = min_logvar
+ self.max_logvar = max_logvar
+
+ def sample(self, mean: torch.Tensor, logvar: torch.Tensor) -> torch.Tensor: # [B,C,...], [B,C,...] -> [B,C,...]
+ std = torch.exp(0.5 * logvar) # [B,C,...]
+ return mean + std * torch.randn_like(mean) # [B,C,...]
+
+ def forward(
+ self, parameters: torch.Tensor, split: bool = False
+ ) -> tuple[torch.Tensor, tuple[torch.Tensor, torch.Tensor]]: # parameters: [B,2C,...], returns sample: [B,C,...]
+ del split
+ mean, logvar = torch.chunk(parameters, 2, dim=1) # [B,C,...], [B,C,...]
+ mean, logvar = mean.contiguous(), logvar.contiguous() # [B,C,...], [B,C,...]
+ logvar = torch.clamp(logvar, self.min_logvar, self.max_logvar) # [B,C,...]
+ return self.sample(mean, logvar), (mean, logvar)
+
+
+class ContinuousFormulation(Enum):
+ VAE = GaussianDistribution
+ AE = IdentityDistribution
diff --git a/cosmos_framework/model/generator/tokenizers/lidar/network/encoding.py b/cosmos_framework/model/generator/tokenizers/lidar/network/encoding.py
new file mode 100644
index 000000000..912c3cd4e
--- /dev/null
+++ b/cosmos_framework/model/generator/tokenizers/lidar/network/encoding.py
@@ -0,0 +1,26 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+
+"""Coordinate helpers for the LiDAR TransformerVAE.
+
+Only polar range-map coordinates are used (RoPE / neighborhood attention).
+Alternate absolute PE modes (spherical harmonics, Fourier features) are not
+part of the shipped checkpoint and were removed.
+"""
+
+from __future__ import annotations
+
+import torch
+
+
+def generate_polar_coords(H: int, W: int, device: torch.device | str = "cpu") -> torch.Tensor: # returns [1,2,H,W]
+ """Build polar angles for a range map.
+
+ theta: azimuthal angle in [-pi, pi]
+ phi: polar angle in [0, pi]
+ """
+ phi = (0.5 - torch.arange(H, device=device) / H) * torch.pi # [H]
+ theta = (1 - torch.arange(W, device=device) / W) * 2 * torch.pi - torch.pi # [W]
+ phi, theta = torch.meshgrid(phi, theta, indexing="ij") # [H,W], [H,W]
+ angles = torch.stack([phi, theta]) # [2,H,W]
+ return angles[None] # [1,2,H,W]
diff --git a/cosmos_framework/model/generator/tokenizers/lidar/network/ops.py b/cosmos_framework/model/generator/tokenizers/lidar/network/ops.py
new file mode 100644
index 000000000..2c23435f5
--- /dev/null
+++ b/cosmos_framework/model/generator/tokenizers/lidar/network/ops.py
@@ -0,0 +1,212 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+
+from collections.abc import Sequence
+
+import numpy as np
+import torch
+import torch.nn.functional as F
+from einops.layers.torch import Rearrange
+from torch import nn
+from torch.nn.modules.utils import _pair, _quadruple
+
+
+def zero_out(m: nn.Module) -> None:
+ for p in m.parameters():
+ p.data.zero_()
+
+
+class SinusoidalPositionalEmbedding(nn.Module):
+ def __init__(self, channels: int, max_period: int = 10_000) -> None:
+ super().__init__()
+ self.channels = channels
+ self.max_period = max_period
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor: # x: [B], returns [B,C]
+ assert len(x.shape) == 1
+ h = -np.log(self.max_period) / (self.channels // 2 - 1)
+ h = torch.exp(h * torch.arange(self.channels // 2, device=x.device)) # [C/2]
+ h = x[:, None] * h[None, :] # [B,C/2]
+ h = torch.cat([h.sin(), h.cos()], dim=-1) # [B,C]
+ return h.to(x) # [B,C]
+
+ def extra_repr(self) -> str:
+ return f"dim={self.channels} max_period={self.max_period}"
+
+
+class Pad(nn.Module):
+ def __init__(
+ self,
+ padding: int | tuple[int, int] | tuple[int, int, int, int],
+ ring: bool = False,
+ mode: str = "constant",
+ ) -> None:
+ super().__init__()
+ self.padding = _quadruple(padding)
+ self.horizontal = "circular" if ring else mode
+ self.vertical = mode
+
+ def forward(self, h: torch.Tensor) -> torch.Tensor: # h: [B,C,H,W], returns [B,C,H',W']
+ left, right, top, bottom = self.padding
+ h = F.pad(h, (left, right, 0, 0), mode=self.horizontal) # [B,C,H,W']
+ h = F.pad(h, (0, 0, top, bottom), mode=self.vertical) # [B,C,H',W']
+ return h
+
+ def extra_repr(self) -> str:
+ return f"padding={self.padding}, " + f"horizontal={self.horizontal}, vertical={self.vertical}"
+
+
+class Resample(nn.Module):
+ def __init__(
+ self,
+ up: int | tuple[int, int] = 1,
+ down: int | tuple[int, int] = 1,
+ window: Sequence[float] = (1, 3, 3, 1), # bilinear
+ ring: bool = True,
+ normalize: bool = True,
+ direction: str = "hw",
+ mode: str = "constant",
+ ) -> None:
+ super().__init__()
+ self.up = np.asarray(_pair(up))
+ self.down = np.asarray(_pair(down))
+ self.window = window
+ self.n_taps = len(window)
+ self.ring = ring
+ self.pad_mode_w = "circular" if ring else mode
+ self.pad_mode_h = mode
+ self.normalize = normalize
+ self.direction = direction
+ assert self.direction in ("h", "w", "hw")
+
+ # setup sizes
+ if "h" in self.direction:
+ self.k_h = self.n_taps
+ self.up_h = self.up[0]
+ self.down_h = self.down[0]
+ else:
+ self.k_h = self.up_h = self.down_h = 1
+
+ if "w" in self.direction:
+ self.k_w = self.n_taps
+ self.up_w = self.up[1]
+ self.down_w = self.down[1]
+ else:
+ self.k_w = self.up_w = self.down_w = 1
+
+ # setup filter
+ kernel = torch.tensor(self.window, dtype=torch.float32) # [K]
+ if self.normalize:
+ kernel /= kernel.sum()
+ kernel *= (self.up_h * self.up_w) ** (kernel.ndim / 2)
+ self.register_buffer("kernel", kernel)
+
+ # setup padding
+ if self.up[0] > 1:
+ self.ph0 = (self.k_h - self.up_h + 1) // 2 + self.up_h - 1
+ self.ph1 = (self.k_h - self.up_h) // 2
+ elif self.down[0] >= 1:
+ self.ph0 = (self.k_h - self.down_h + 1) // 2
+ self.ph1 = (self.k_h - self.down_h) // 2
+ if self.up[1] > 1:
+ self.pw0 = (self.k_w - self.up_w + 1) // 2 + self.up_w - 1
+ self.pw1 = (self.k_w - self.up_w) // 2
+ elif self.down[1] >= 1:
+ self.pw0 = (self.k_w - self.down_w + 1) // 2
+ self.pw1 = (self.k_w - self.down_w) // 2
+
+ self.margin = int(max(self.ph0, self.ph1, self.pw0, self.pw1))
+
+ def forward(self, h: torch.Tensor) -> torch.Tensor: # h: [B,C,H,W], returns [B,C,H',W']
+ # margin
+ h = F.pad(h, (self.margin, self.margin, 0, 0), mode=self.pad_mode_w) # [B,C,H,Wm]
+ h = F.pad(h, (0, 0, self.margin, self.margin), mode=self.pad_mode_h) # [B,C,Hm,Wm]
+ # up by zero-insertion
+ B, C, H, W = h.shape
+ h = h.view(B, C, H, 1, W, 1) # [B,C,H,1,W,1]
+ h = F.pad(h, [0, self.up_w - 1, 0, 0, 0, self.up_h - 1]) # [B,C,H,UH,W,UW]
+ h = h.view(B, C, H * self.up_h, W * self.up_w) # [B,C,Hu,Wu]
+ # crop
+ h = h[
+ ...,
+ self.margin * self.up_h - self.ph0 : (H - self.margin) * self.up_h + self.ph1,
+ self.margin * self.up_w - self.pw0 : (W - self.margin) * self.up_w + self.pw1,
+ ]
+ # fir
+ kernel = self.kernel[None, None].repeat(C, 1, 1).to(dtype=h.dtype) # [C,1,K]
+ if self.direction == "hw":
+ h = F.conv2d(h, kernel[..., None, :], groups=C)
+ h = F.conv2d(h, kernel[..., :, None], groups=C)
+ elif self.direction == "h":
+ h = F.conv2d(h, kernel[..., :, None], groups=C)
+ elif self.direction == "w":
+ h = F.conv2d(h, kernel[..., None, :], groups=C)
+ # down
+ h = h[:, :, :: self.down_h, :: self.down_w]
+ return h
+
+ def extra_repr(self) -> str:
+ return f"up={tuple(self.up)}, down={tuple(self.down)}, ring={self.ring}"
+
+
+class Conv2d(nn.Conv2d):
+ def __init__(
+ self,
+ in_channels: int,
+ out_channels: int,
+ kernel_size: int | tuple[int, int],
+ stride: int | tuple[int, int],
+ padding: int | tuple[int, int],
+ bias: bool = True,
+ ring: bool = False,
+ ) -> None:
+ super().__init__(
+ in_channels=in_channels,
+ out_channels=out_channels,
+ kernel_size=kernel_size,
+ stride=stride,
+ padding=0,
+ bias=bias,
+ )
+ self.pad = Pad(padding=padding, ring=ring) if padding != 0 else None
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor: # x: [B,Cin,H,W], returns [B,Cout,H',W']
+ if self.pad is not None:
+ x = self.pad(x)
+ return super().forward(x)
+
+
+class AdaGN(nn.GroupNorm):
+ def __init__(
+ self,
+ emb_channels: int,
+ out_channels: int,
+ num_groups: int,
+ eps: float = 1e-5,
+ ) -> None:
+ super().__init__(
+ num_groups=num_groups,
+ num_channels=out_channels,
+ eps=eps,
+ affine=False,
+ )
+ self.proj = nn.Sequential(
+ nn.SiLU(),
+ nn.Linear(emb_channels, out_channels * 2),
+ Rearrange("B C -> B C 1 1"),
+ )
+
+ def forward(
+ self, x: torch.Tensor, emb: torch.Tensor
+ ) -> torch.Tensor: # x: [B,C,H,W], emb: [B,E], returns [B,C,H,W]
+ h = super().forward(x) # [B,C,H,W]
+ scale, shift = self.proj(emb).chunk(2, dim=1) # [B,C,1,1], [B,C,1,1]
+ h = h * (1 + scale) + shift # [B,C,H,W]
+ return h
+
+
+class ConditionalSequential(nn.Sequential):
+ def forward(self, x: torch.Tensor, condition: torch.Tensor) -> torch.Tensor:
+ for module in self:
+ x = module(x, condition)
+ return x
diff --git a/cosmos_framework/model/generator/tokenizers/lidar/network/rope3d.py b/cosmos_framework/model/generator/tokenizers/lidar/network/rope3d.py
new file mode 100644
index 000000000..7487841c7
--- /dev/null
+++ b/cosmos_framework/model/generator/tokenizers/lidar/network/rope3d.py
@@ -0,0 +1,141 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+
+"""3D rotary positional embedding for ``(T, H, W)`` token grids.
+
+Adapted (PyTorch-only, no TransformerEngine) from
+the Cosmos video-DiT ``VideoRopePosition3DEmb``: the head dim
+is split into three chunks ``(dim_t, dim_h, dim_w)``, each axis contributes
+its own rotary frequency band, and the resulting angles are duplicated
+across the two rotation halves (GPT-NeoX style) so that
+:func:`apply_rotary_emb` can rotate Q/K with a single tensor.
+
+For a query at position ``p_q = (t_q, h_q, w_q)`` and a key at
+``p_k = (t_k, h_k, w_k)``, the dot product after rotation depends only on
+``p_q - p_k``. If the query and key share the same value along some axis
+(e.g. a factored *spatial* attention block where both are on the same
+frame), the rotation along that axis cancels out -- so this single 3D
+RoPE specializes correctly to 2D RoPE for spatial-only blocks and 1D RoPE
+for temporal-only blocks.
+"""
+
+from __future__ import annotations
+
+import torch
+from torch import nn
+
+
+class VideoRopePosition3DEmb(nn.Module):
+ """3D RoPE producing angles for a ``(T, H, W)`` token grid.
+
+ Args:
+ head_dim: Per-head channel count. Must satisfy
+ ``head_dim - 2 * (head_dim // 6 * 2) > 0``.
+ len_t, len_h, len_w: Maximum supported grid sizes (frequencies are
+ cached for these). ``forward`` accepts any smaller ``(T, H, W)``.
+ base_theta: RoPE base (``10000.0`` matches LLaMA/Cosmos).
+ t_extrapolation_ratio, h_extrapolation_ratio, w_extrapolation_ratio:
+ NTK-aware extrapolation factors, applied as
+ ``theta *= ratio ** (dim_axis / (dim_axis - 2))``.
+ """
+
+ def __init__(
+ self,
+ *,
+ head_dim: int,
+ len_t: int,
+ len_h: int,
+ len_w: int,
+ base_theta: float = 10000.0,
+ t_extrapolation_ratio: float = 1.0,
+ h_extrapolation_ratio: float = 1.0,
+ w_extrapolation_ratio: float = 1.0,
+ ) -> None:
+ super().__init__()
+ dim = head_dim
+ dim_h = (dim // 6) * 2
+ dim_w = dim_h
+ dim_t = dim - dim_h - dim_w
+ assert dim_t >= 2, f"head_dim={dim} too small for 3D RoPE (dim_t={dim_t}, dim_h={dim_h}, dim_w={dim_w})"
+ self.head_dim = dim
+ self.dim_t = dim_t
+ self.dim_h = dim_h
+ self.dim_w = dim_w
+ self.max_t = len_t
+ self.max_h = len_h
+ self.max_w = len_w
+
+ # NTK-aware extrapolation exponent ``dim_axis / (dim_axis - 2)`` is
+ # only well-defined for ``dim_axis > 2``. Smaller axes are degenerate
+ # (only one frequency bucket) and would also raise ZeroDivisionError
+ # below regardless of the ratio. Fall back to base_theta in that
+ # case -- when the ratio is 1.0 (the default) the expression is a
+ # no-op anyway, so we lose nothing.
+ def _theta(ratio: float, dim: int) -> float:
+ if ratio == 1.0 or dim <= 2:
+ return base_theta
+ return base_theta * (ratio ** (dim / (dim - 2)))
+
+ h_theta = _theta(h_extrapolation_ratio, dim_h)
+ w_theta = _theta(w_extrapolation_ratio, dim_w)
+ t_theta = _theta(t_extrapolation_ratio, dim_t)
+
+ idx_h = torch.arange(0, dim_h, 2, dtype=torch.float32) / dim_h
+ idx_w = torch.arange(0, dim_w, 2, dtype=torch.float32) / dim_w
+ idx_t = torch.arange(0, dim_t, 2, dtype=torch.float32) / dim_t
+ self.register_buffer("freqs_h", 1.0 / (h_theta**idx_h), persistent=False)
+ self.register_buffer("freqs_w", 1.0 / (w_theta**idx_w), persistent=False)
+ self.register_buffer("freqs_t", 1.0 / (t_theta**idx_t), persistent=False)
+ self.register_buffer("seq_h", torch.arange(len_h, dtype=torch.float32), persistent=False)
+ self.register_buffer("seq_w", torch.arange(len_w, dtype=torch.float32), persistent=False)
+ self.register_buffer("seq_t", torch.arange(len_t, dtype=torch.float32), persistent=False)
+
+ def forward(self, T: int, H: int, W: int) -> torch.Tensor: # returns [THW,D]
+ """Returns rope angles of shape ``(T*H*W, head_dim)``.
+
+ Each half of the last dim duplicates the per-axis angles
+ ``[t_half | h_half | w_half | t_half | h_half | w_half]`` so that
+ :func:`apply_rotary_emb` (GPT-NeoX style) can rotate Q/K via
+ ``[x_lo, x_hi] @ rot``.
+ """
+ assert T <= self.max_t and H <= self.max_h and W <= self.max_w, (
+ f"Requested ({T},{H},{W}) exceeds cached ({self.max_t},{self.max_h},{self.max_w})"
+ )
+ ang_t = torch.outer(self.seq_t[:T], self.freqs_t) # [T,Dt/2]
+ ang_h = torch.outer(self.seq_h[:H], self.freqs_h) # [H,Dh/2]
+ ang_w = torch.outer(self.seq_w[:W], self.freqs_w) # [W,Dw/2]
+ ang_t_b = ang_t.view(T, 1, 1, -1).expand(T, H, W, -1) # [T,H,W,Dt/2]
+ ang_h_b = ang_h.view(1, H, 1, -1).expand(T, H, W, -1) # [T,H,W,Dh/2]
+ ang_w_b = ang_w.view(1, 1, W, -1).expand(T, H, W, -1) # [T,H,W,Dw/2]
+ half = torch.cat([ang_t_b, ang_h_b, ang_w_b], dim=-1) # [T,H,W,D/2]
+ full = torch.cat([half, half], dim=-1) # [T,H,W,D]
+ return full.reshape(T * H * W, -1) # [THW,D]
+
+
+def apply_rotary_emb(
+ x: torch.Tensor, rope_emb: torch.Tensor
+) -> torch.Tensor: # x: [...,S,N,D], rope_emb: [S,D], returns [...,S,N,D]
+ """Rotate the last dim of ``x`` via GPT-NeoX rotary embeddings.
+
+ Args:
+ x: ``(..., S, num_heads, head_dim)`` (any leading shape).
+ rope_emb: ``(S, head_dim)`` angles tensor produced by
+ :class:`VideoRopePosition3DEmb`.
+
+ Notes:
+ Standard GPT-NeoX rotation pairs channel ``i`` with channel
+ ``i + head_dim/2`` (not adjacent channels). With ``rope_emb``'s
+ ``[half | half]`` duplication, ``cos(rope_emb[:half]) ==
+ cos(rope_emb[half:])`` so both halves of a pair rotate by the same
+ angle.
+ """
+ rope_emb = rope_emb.to(x.dtype).unsqueeze(-2) # [S,1,D]
+ cos = rope_emb.cos() # [S,1,D]
+ sin = rope_emb.sin() # [S,1,D]
+ half = x.shape[-1] // 2
+ x_lo, x_hi = x[..., :half], x[..., half:] # [...,S,N,D/2], [...,S,N,D/2]
+ cos_lo, cos_hi = cos[..., :half], cos[..., half:] # [S,1,D/2], [S,1,D/2]
+ sin_lo, sin_hi = sin[..., :half], sin[..., half:] # [S,1,D/2], [S,1,D/2]
+ out_lo = x_lo * cos_lo - x_hi * sin_lo # [...,S,N,D/2]
+ out_hi = x_hi * cos_hi + x_lo * sin_hi # [...,S,N,D/2]
+ return torch.cat([out_lo, out_hi], dim=-1) # [...,S,N,D]
diff --git a/cosmos_framework/model/generator/tokenizers/lidar/network/transformer_vae.py b/cosmos_framework/model/generator/tokenizers/lidar/network/transformer_vae.py
new file mode 100644
index 000000000..77636f071
--- /dev/null
+++ b/cosmos_framework/model/generator/tokenizers/lidar/network/transformer_vae.py
@@ -0,0 +1,2283 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+
+from __future__ import annotations
+
+# =============================================================================
+# Transformer VAE with hierarchical hourglass architecture,
+# neighborhood attention (natten) for spatial processing, and
+# causal temporal attention for video.
+#
+# Architecture:
+# Encoder: Tokenizer → [spatial blocks + PatchMerging] (level 0, no temporal)
+# → [spatial→temporal blocks + PatchMerging (+ TemporalMerging)] × (N-1)
+# → global spatial+temporal → z
+# Decoder: z → global temporal+spatial
+# → [TemporalExpanding + PatchExpanding + temporal→spatial blocks] × (N-1)
+# → [PatchExpanding + spatial blocks] (level 0, no temporal) → Detokenizer
+#
+# No skip connections between encoder and decoder.
+#
+# Based on: https://github.com/crowsonkb/k-diffusion/blob/master/
+# k_diffusion/models/image_transformer_v2.py
+# Call `natten.use_fused_na(True)` for acceleration on GPUs.
+# =============================================================================
+import functools
+import math
+from typing import Any, Literal
+
+import einops
+import torch
+import torch.nn.functional as F
+from einops.layers.torch import Rearrange
+from torch import nn
+from torch.nn.modules.utils import _pair
+from torch.utils.checkpoint import checkpoint
+
+from cosmos_framework.model.generator.tokenizers.lidar.network import encoding, ops
+from cosmos_framework.model.generator.tokenizers.lidar.network.distributions import ContinuousFormulation
+from cosmos_framework.model.generator.tokenizers.lidar.network.rope3d import (
+ VideoRopePosition3DEmb,
+ apply_rotary_emb,
+)
+
+__all__ = ["TransformerVAE"]
+
+
+@functools.lru_cache(maxsize=1)
+def _natten_na2d():
+ """Resolve ``natten.functional.na2d`` on first use.
+
+ NATTEN is an optional, platform-gated wheel -- the ``cu128`` install group marks it
+ ``platform_machine == 'x86_64'`` while torch carries no marker -- so it can be missing from an
+ otherwise working install. Only a forward pass through a local-attention block needs its
+ kernels, so resolving it here keeps this module importable, constructible and
+ checkpoint-loadable everywhere, and turns the failure into an actionable message at the point
+ that actually needs it. Follows the frontend's explain-before-importing idiom rather than a
+ bare deferred import, which would surface a plain ModuleNotFoundError from inside an
+ activation-checkpoint recompute.
+ """
+ try:
+ import natten
+ except ImportError as exc:
+ raise ImportError(
+ "The LiDAR TransformerVAE needs NATTEN for its neighborhood-attention blocks. "
+ "Install a natten wheel for this platform (e.g. `uv sync --group=cu130-train`)."
+ ) from exc
+ return natten.functional.na2d
+
+
+# =============================================================================
+# Normalization
+# =============================================================================
+
+
+class RMSNorm(torch.nn.Module):
+ def __init__(self, in_dim: int, scale: bool = True, eps: float = 1e-6):
+ super().__init__()
+ self.in_dim = in_dim
+ self.eps = eps
+ self.scale = nn.Parameter(torch.ones(in_dim)) if scale else 1.0
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ x_fp32 = x.float()
+ x_normed = x_fp32 * torch.rsqrt(x_fp32.pow(2).mean(-1, keepdim=True) + self.eps)
+ return (x_normed * self.scale).to(x)
+
+ def extra_repr(self) -> str:
+ return f"in_dim={self.in_dim}"
+
+
+# =============================================================================
+# Rotary positional encoding (spatial)
+# =============================================================================
+
+
+class AxialRoPE(nn.Module):
+ def __init__(self, dim: int, num_heads: int, max_harmonics: list[int]):
+ super().__init__()
+ freqs_h = self.setup_freqs(num_heads * dim // 4, max_harmonics[0])
+ freqs_w = self.setup_freqs(num_heads * dim // 4, max_harmonics[1])
+ self.register_buffer("freqs_h", freqs_h.view(dim // 4, num_heads).T)
+ self.register_buffer("freqs_w", freqs_w.view(dim // 4, num_heads).T)
+
+ def setup_freqs(self, dim: int, max_harmonics: int):
+ return torch.linspace(math.log(1), math.log(max_harmonics), dim).exp().round()
+
+ def forward(self, coords: torch.Tensor) -> torch.Tensor:
+ coords = einops.rearrange(coords, "b c h w -> b h w c")
+ radian_h = coords[..., None, [0]] * self.freqs_h
+ radian_w = coords[..., None, [1]] * self.freqs_w
+ return torch.cat((radian_h, radian_w), dim=-1)
+
+ @staticmethod
+ def rotate(x: torch.Tensor, theta: torch.Tensor) -> torch.Tensor:
+ x1, x2 = x.chunk(2, dim=-1)
+ y1 = x1 * theta.cos() - x2 * theta.sin()
+ y2 = x1 * theta.sin() + x2 * theta.cos()
+ return torch.cat((y1, y2), dim=-1)
+
+ def extra_repr(self) -> str:
+ return f"freqs_h={tuple(self.freqs_h.shape)}, freqs_w={tuple(self.freqs_w.shape)}"
+
+
+# =============================================================================
+# Spatial self-attention blocks
+# =============================================================================
+
+
+class GlobalSelfAttentionBlock(nn.Module):
+ def __init__(
+ self,
+ dim: int,
+ num_heads: int,
+ dropout: float = 0.0,
+ rope_max_harmonics: list[int] = (1, 1),
+ bias=False,
+ eps=1e-6,
+ ):
+ super().__init__()
+ self.num_heads = num_heads
+ self.head_dim = dim // self.num_heads
+ self.eps = eps
+
+ self.norm = RMSNorm(dim)
+ self.scale = nn.Parameter(torch.full([self.num_heads, 1], math.log(10.0)))
+ self.qkv_proj = nn.Linear(dim, dim * 3, bias=bias)
+ self.rope = AxialRoPE(self.head_dim, num_heads, rope_max_harmonics)
+ self.dropout = nn.Dropout(dropout)
+ self.out_proj = nn.Linear(dim, dim, bias=bias).apply(ops.zero_out)
+
+ def scale_qk(self, q: torch.Tensor, k: torch.Tensor) -> torch.Tensor:
+ scale = self.scale.clamp(max=math.log(100)).exp().sqrt()
+ q = (F.normalize(q, p=2, dim=-1, eps=self.eps) * scale).to(q.dtype)
+ k = (F.normalize(k, p=2, dim=-1, eps=self.eps) * scale).to(k.dtype)
+ return q, k
+
+ def apply_rope_qk(self, q: torch.Tensor, k: torch.Tensor, coords: torch.Tensor):
+ theta = self.rope(coords)
+ d = theta.shape[-1] * 2
+ assert (q.shape[-1] >= d) and (k.shape[-1] >= d)
+ q[..., :d] = self.rope.rotate(q[..., :d], theta).to(q.dtype)
+ k[..., :d] = self.rope.rotate(k[..., :d], theta).to(k.dtype)
+ return q, k
+
+ def residual(self, x: torch.Tensor, coords: torch.Tensor) -> torch.Tensor:
+ B, H, W, C = x.shape
+ h = self.norm(x)
+ qkv = self.qkv_proj(h)
+ q, k, v = einops.rearrange(qkv, "B H W (T N D) -> T B H W N D", T=3, D=self.head_dim)
+ q, k = self.scale_qk(q, k)
+ q, k = self.apply_rope_qk(q, k, coords)
+ q = einops.rearrange(q, "B H W N D -> B N (H W) D")
+ k = einops.rearrange(k, "B H W N D -> B N (H W) D")
+ v = einops.rearrange(v, "B H W N D -> B N (H W) D")
+ h = F.scaled_dot_product_attention(q, k, v, scale=1.0)
+ h = einops.rearrange(h, "B N (H W) D -> B H W (N D)", H=H, W=W)
+ h = self.dropout(h)
+ h = self.out_proj(h)
+ return h
+
+ def forward(self, x: torch.Tensor, coords: torch.Tensor) -> torch.Tensor:
+ return x + self.residual(x, coords)
+
+ def extra_repr(self) -> str:
+ return f"head_dim={self.head_dim}, num_heads={self.num_heads}"
+
+
+class CircularNeighborhoodSelfAttentionBlock(GlobalSelfAttentionBlock):
+ def __init__(
+ self,
+ dim: int,
+ num_heads: int,
+ kernel_size: list[int],
+ dilation: list[int] = 1,
+ dropout: float = 0.0,
+ rope_max_harmonics: list[int] = (1, 1),
+ # When True (default), the W (last spatial) axis is treated as
+ # periodic: it is circularly padded before the neighborhood attention
+ # so the seam wraps (range-image azimuth, column 0 ~ column W-1). When
+ # False, no wrap -- ``natten`` clamps the neighborhood at the W edges
+ # (the same boundary handling it already uses for H). Set False for a
+ # Cartesian grid (e.g. BEV x-axis) where the two W edges are unrelated.
+ circular: bool = True,
+ ):
+ super().__init__(
+ dim=dim,
+ num_heads=num_heads,
+ dropout=dropout,
+ rope_max_harmonics=rope_max_harmonics,
+ )
+ self.kernel_size = _pair(kernel_size)
+ self.dilation = _pair(dilation)
+ self.circular = circular
+
+ def before_attn(self, q, k, v):
+ if not self.circular:
+ return q, k, v
+ padding = self.kernel_size[1] // 2
+ q = F.pad(q, (0, 0, 0, 0, padding, padding), mode="circular")
+ k = F.pad(k, (0, 0, 0, 0, padding, padding), mode="circular")
+ v = F.pad(v, (0, 0, 0, 0, padding, padding), mode="circular")
+ return q, k, v
+
+ def after_attn(self, x):
+ if not self.circular:
+ return x
+ padding = self.kernel_size[1] // 2
+ x = x[:, :, padding:-padding]
+ return x
+
+ def residual(self, x: torch.Tensor, coords: torch.Tensor) -> torch.Tensor:
+ h = self.norm(x)
+ qkv = self.qkv_proj(h)
+ q, k, v = einops.rearrange(qkv, "B H W (T N D) -> T B H W N D", T=3, D=self.head_dim)
+ q, k = self.scale_qk(q, k)
+ q, k = self.apply_rope_qk(q, k, coords)
+ q, k, v = self.before_attn(q, k, v)
+ h = _natten_na2d()(
+ query=q,
+ key=k,
+ value=v,
+ kernel_size=self.kernel_size,
+ dilation=self.dilation,
+ scale=1.0,
+ )
+ h = einops.rearrange(h, "B H W N D -> B H W (N D)")
+ h = self.after_attn(h)
+ h = self.dropout(h)
+ h = self.out_proj(h)
+ return h
+
+ def forward(self, x: torch.Tensor, coords: torch.Tensor) -> torch.Tensor:
+ return x + self.residual(x, coords)
+
+ def extra_repr(self) -> str:
+ return (
+ f"head_dim={self.head_dim}, num_heads={self.num_heads}, "
+ f"kernel_size={self.kernel_size}, dilation={self.dilation}"
+ )
+
+
+# =============================================================================
+# Causal temporal self-attention
+# =============================================================================
+
+
+class CausalTemporalAttention(nn.Module):
+ """Causal self-attention along the temporal dimension.
+
+ Input/output shape: (N, T, C) where N = B*H*W.
+ Each position attends only to itself and earlier positions.
+ """
+
+ def __init__(self, dim: int, num_heads: int, dropout: float = 0.0, bias=False):
+ super().__init__()
+ self.num_heads = num_heads
+ self.head_dim = dim // num_heads
+
+ self.norm = RMSNorm(dim)
+ self.qkv_proj = nn.Linear(dim, dim * 3, bias=bias)
+ self.out_proj = nn.Linear(dim, dim, bias=bias).apply(ops.zero_out)
+ self.dropout = nn.Dropout(dropout)
+ self._SDPA_MAX_BATCH = 65535
+
+ def _project_qkv(self, x: torch.Tensor):
+ h = self.norm(x)
+ qkv = self.qkv_proj(h)
+ return einops.rearrange(
+ qkv,
+ "N T (three heads D) -> three N heads T D",
+ three=3,
+ D=self.head_dim,
+ )
+
+ def _sdpa(self, q, k, v, *, attn_mask=None, is_causal=False):
+ N = q.shape[0]
+ if N <= self._SDPA_MAX_BATCH:
+ return F.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask, is_causal=is_causal)
+ chunks = []
+ for i in range(0, N, self._SDPA_MAX_BATCH):
+ j = min(i + self._SDPA_MAX_BATCH, N)
+ chunks.append(
+ F.scaled_dot_product_attention(
+ q[i:j],
+ k[i:j],
+ v[i:j],
+ attn_mask=attn_mask,
+ is_causal=is_causal,
+ )
+ )
+ return torch.cat(chunks, dim=0)
+
+ def _finish(self, x, h):
+ h = einops.rearrange(h, "N heads T D -> N T (heads D)")
+ h = self.dropout(h)
+ h = self.out_proj(h)
+ return x + h
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ q, k, v = self._project_qkv(x)
+ return self._finish(x, self._sdpa(q, k, v, is_causal=True))
+
+ def forward_stream(
+ self,
+ x: torch.Tensor,
+ kv_cache: tuple[torch.Tensor, torch.Tensor] | None = None,
+ ) -> tuple[torch.Tensor, tuple[torch.Tensor, torch.Tensor]]:
+ """Attend new temporal tokens to cached history and this causal block."""
+ q, k_new, v_new = self._project_qkv(x)
+ past = 0
+ if kv_cache is not None:
+ k_past, v_past = kv_cache
+ if (
+ k_past.shape[:2] != k_new.shape[:2]
+ or k_past.shape[3:] != k_new.shape[3:]
+ or v_past.shape != k_past.shape
+ ):
+ raise ValueError(
+ "temporal KV cache shape mismatch: "
+ f"cached k/v={tuple(k_past.shape)}/{tuple(v_past.shape)}, "
+ f"new k={tuple(k_new.shape)}"
+ )
+ past = k_past.shape[2]
+ k = torch.cat([k_past, k_new], dim=2)
+ v = torch.cat([v_past, v_new], dim=2)
+ else:
+ k, v = k_new, v_new
+
+ T_new = q.shape[2]
+ q_index = past + torch.arange(T_new, device=q.device)[:, None]
+ k_index = torch.arange(past + T_new, device=q.device)[None, :]
+ causal_offset_mask = k_index <= q_index
+ out = self._finish(
+ x,
+ self._sdpa(
+ q,
+ k,
+ v,
+ attn_mask=causal_offset_mask,
+ is_causal=False,
+ ),
+ )
+ return out, (k.detach(), v.detach())
+
+ def extra_repr(self) -> str:
+ return f"head_dim={self.head_dim}, num_heads={self.num_heads}"
+
+
+# =============================================================================
+# Feed-forward
+# =============================================================================
+
+
+class GEGLU(nn.Linear):
+ def __init__(self, in_features, out_features, bias):
+ super().__init__(in_features, out_features * 2, bias=bias)
+
+ def forward(self, x):
+ h = super().forward(x)
+ h, gate = h.chunk(2, dim=-1)
+ return h * F.gelu(gate)
+
+
+class FeedForwardNetwork(nn.Module):
+ def __init__(self, dim, mid_dim, dropout=0.0):
+ super().__init__()
+ self.norm = RMSNorm(dim)
+ self.gegelu = GEGLU(dim, mid_dim, bias=False)
+ self.dropout = nn.Dropout(dropout)
+ self.linear = nn.Linear(mid_dim, dim, bias=False).apply(ops.zero_out)
+
+ def residual(self, x: torch.Tensor) -> torch.Tensor:
+ x = self.norm(x)
+ x = self.gegelu(x)
+ x = self.dropout(x)
+ x = self.linear(x)
+ return x
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ return x + self.residual(x)
+
+
+# =============================================================================
+# Spatial block (spatial attention + FFN)
+# =============================================================================
+
+
+class SpatialBlock(nn.Module):
+ """Spatial attention + FFN. Operates on (BT, H, W, C) per-frame."""
+
+ def __init__(
+ self,
+ in_dim: int,
+ num_heads: int,
+ attn_type: Literal["local", "global"] = "global",
+ kernel_size: list[int] | None = None,
+ dilation: list[int] = 1,
+ rope_max_harmonics: list[int] = (1, 1),
+ mlp_ratio: float = 3.0,
+ dropout: float = 0.0,
+ # Forwarded to the local (neighborhood) attention: whether the W axis
+ # wraps (periodic azimuth) or clamps (Cartesian). No effect for global.
+ circular: bool = True,
+ ):
+ super().__init__()
+
+ if attn_type == "global":
+ self.residual_attn = GlobalSelfAttentionBlock(
+ dim=in_dim,
+ num_heads=num_heads,
+ dropout=dropout,
+ rope_max_harmonics=rope_max_harmonics,
+ )
+ elif attn_type == "local":
+ self.residual_attn = CircularNeighborhoodSelfAttentionBlock(
+ dim=in_dim,
+ num_heads=num_heads,
+ kernel_size=kernel_size,
+ dilation=dilation,
+ dropout=dropout,
+ rope_max_harmonics=rope_max_harmonics,
+ circular=circular,
+ )
+
+ self.residual_ffn = FeedForwardNetwork(
+ dim=in_dim,
+ mid_dim=int(in_dim * mlp_ratio),
+ dropout=dropout,
+ )
+
+ def forward(self, x: torch.Tensor, coords: torch.Tensor) -> torch.Tensor:
+ x = self.residual_attn(x, coords)
+ x = self.residual_ffn(x)
+ return x
+
+
+class TemporalBlock(nn.Module):
+ """Causal temporal attention + FFN. Operates on (BHW, T, C)."""
+
+ def __init__(
+ self,
+ in_dim: int,
+ num_heads: int,
+ mlp_ratio: float = 3.0,
+ dropout: float = 0.0,
+ ):
+ super().__init__()
+ self.temporal_attn = CausalTemporalAttention(in_dim, num_heads, dropout)
+ self.ffn = FeedForwardNetwork(in_dim, int(in_dim * mlp_ratio), dropout)
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ x = self.temporal_attn(x)
+ x = self.ffn(x)
+ return x
+
+ def forward_stream(
+ self,
+ x: torch.Tensor,
+ kv_cache: tuple[torch.Tensor, torch.Tensor] | None = None,
+ ) -> tuple[torch.Tensor, tuple[torch.Tensor, torch.Tensor]]:
+ x, kv_cache = self.temporal_attn.forward_stream(x, kv_cache)
+ x = self.ffn(x)
+ return x, kv_cache
+
+
+class CausalTemporalConv(nn.Module):
+ """Causal temporal convolution along T (WAN-style).
+
+ The conv analog of :class:`CausalTemporalAttention` -- same
+ (N, T, C) interface (N = B*H*W), same zero-init residual so the
+ block is a residual add at init. Implemented as a depth-preserving
+ ``Conv1d`` over the T axis with left-only causal padding (kernel-1),
+ so each output only depends on the current and past frames -- the
+ ``kernel_size x 1 x 1`` case of WanVAE's ``CausalConv3d`` (spatial
+ mixing is left to the surrounding neighborhood-attention blocks, so
+ swapping only the temporal mixer stays a clean ablation).
+ """
+
+ def __init__(self, dim: int, kernel_size: int = 3, dropout: float = 0.0):
+ super().__init__()
+ self.norm = RMSNorm(dim)
+ self.conv = nn.Conv1d(dim, dim, kernel_size=kernel_size, bias=False)
+ self.out_proj = nn.Linear(dim, dim, bias=False).apply(ops.zero_out)
+ self.dropout = nn.Dropout(dropout)
+ self.causal_pad = kernel_size - 1
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ # x: (N, T, C)
+ h = self.norm(x)
+ h = einops.rearrange(h, "N T C -> N C T")
+ h = F.pad(h, (self.causal_pad, 0))
+ h = self.conv(h)
+ h = einops.rearrange(h, "N C T -> N T C")
+ h = self.dropout(h)
+ h = self.out_proj(h)
+ return x + h
+
+
+class TemporalConvBlock(nn.Module):
+ """Causal temporal conv + FFN. Operates on (BHW, T, C).
+
+ Drop-in replacement for :class:`TemporalBlock` (identical constructor
+ signature and (N, T, C) forward) that swaps the temporal *attention*
+ for a WAN-style causal temporal *convolution*. ``num_heads`` is
+ accepted only for signature parity (unused).
+ """
+
+ def __init__(
+ self,
+ in_dim: int,
+ num_heads: int,
+ mlp_ratio: float = 3.0,
+ dropout: float = 0.0,
+ kernel_size: int = 3,
+ ):
+ super().__init__()
+ self.temporal_conv = CausalTemporalConv(in_dim, kernel_size, dropout)
+ self.ffn = FeedForwardNetwork(in_dim, int(in_dim * mlp_ratio), dropout)
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ x = self.temporal_conv(x)
+ x = self.ffn(x)
+ return x
+
+
+# =============================================================================
+# Joint 3D self-attention block
+# =============================================================================
+
+
+class Joint3DSelfAttention(nn.Module):
+ """Joint self-attention over flattened (T*H*W) tokens.
+
+ Mirrors :class:`GlobalSelfAttentionBlock` but operates jointly over time
+ and space instead of per-frame.
+
+ Two positional-encoding modes are supported via ``use_3d_rope``:
+
+ * **Default (``use_3d_rope=False``)** -- spatial position via the
+ existing :class:`AxialRoPE` on ``(h, w)`` (broadcast across time, since
+ polar coords don't vary with time), and time position via a *learnable*
+ 1D embedding added to the pre-attention features. This is what
+ original training runs used; kept as the default for backward compat.
+ Caveat: the learnable temporal slots past the largest ``T`` seen in
+ training are uninitialized noise, so this mode does *not* extrapolate
+ cleanly to longer-than-trained sequence lengths at inference.
+ * **``use_3d_rope=True``** -- replace both the spatial :class:`AxialRoPE`
+ and the learnable temporal PE with a single 3D RoPE
+ (:class:`VideoRopePosition3DEmb`) over ``(t, h, w)``. RoPE is a
+ rotation (length-preserving), and is fully parameter-free, so the
+ angles extrapolate to longer sequences with no untrained parameters.
+ The spatial axes of the 3D RoPE play the same role :class:`AxialRoPE`
+ did, but without coupling to polar coords -- the bottleneck operates
+ on a regular (H_z, W_z) grid so a Cartesian RoPE is well-defined and
+ consistent with how the rest of the literature does video VAEs.
+
+ Causality is *orthogonal* to the PE choice: ``causal_time=True`` masks
+ the joint attention so each ``(t, h, w)`` token only attends to
+ ``(t', h', w')`` with ``t' <= t``. The causal mask is materialized as
+ a ``(T*H*W, T*H*W)`` bool tensor so memory scales with ``(T*H*W)^2``.
+ """
+
+ def __init__(
+ self,
+ dim: int,
+ num_heads: int,
+ max_t: int,
+ rope_max_harmonics: list[int] = (1, 1),
+ dropout: float = 0.0,
+ causal_time: bool = False,
+ bias: bool = False,
+ eps: float = 1e-6,
+ # If True, encode both spatial and temporal position with a single
+ # 3D RoPE (length-preserving rotation, no learnable PE buffer) and
+ # skip the original 2D AxialRoPE + learnable `temporal_pe` path.
+ # ``len_h``/``len_w`` set the spatial cap of the 3D RoPE cache.
+ use_3d_rope: bool = False,
+ len_h: int = 1,
+ len_w: int = 1,
+ ):
+ super().__init__()
+ self.num_heads = num_heads
+ self.head_dim = dim // num_heads
+ self.eps = eps
+ self.causal_time = causal_time
+ self.use_3d_rope = use_3d_rope
+ self.max_t = max_t
+ self.max_h = len_h if use_3d_rope else 0
+ self.max_w = len_w if use_3d_rope else 0
+
+ self.norm = RMSNorm(dim)
+ self.scale = nn.Parameter(torch.full([num_heads, 1], math.log(10.0)))
+ self.qkv_proj = nn.Linear(dim, dim * 3, bias=bias)
+ if use_3d_rope:
+ self.rope = None
+ self.temporal_pe = None
+ self.rope3d = VideoRopePosition3DEmb(
+ head_dim=self.head_dim,
+ len_t=max_t,
+ len_h=len_h,
+ len_w=len_w,
+ )
+ else:
+ self.rope = AxialRoPE(self.head_dim, num_heads, rope_max_harmonics)
+ self.temporal_pe = nn.Parameter(torch.zeros(max_t, dim))
+ nn.init.trunc_normal_(self.temporal_pe, std=0.02)
+ self.rope3d = None
+ self.dropout = nn.Dropout(dropout)
+ self.out_proj = nn.Linear(dim, dim, bias=bias).apply(ops.zero_out)
+
+ def _scale_qk(self, q: torch.Tensor, k: torch.Tensor):
+ scale = self.scale.clamp(max=math.log(100)).exp().sqrt()
+ q = (F.normalize(q, p=2, dim=-1, eps=self.eps) * scale).to(q.dtype)
+ k = (F.normalize(k, p=2, dim=-1, eps=self.eps) * scale).to(k.dtype)
+ return q, k
+
+ def _apply_spatial_rope(self, q: torch.Tensor, k: torch.Tensor, coords: torch.Tensor):
+ theta = self.rope(coords) # (B_c, H, W, N, D//2)
+ d = theta.shape[-1] * 2
+ assert (q.shape[-1] >= d) and (k.shape[-1] >= d)
+ theta = theta.unsqueeze(1) # broadcast over T → (B_c, 1, H, W, N, D//2)
+ q[..., :d] = self.rope.rotate(q[..., :d], theta).to(q.dtype)
+ k[..., :d] = self.rope.rotate(k[..., :d], theta).to(k.dtype)
+ return q, k
+
+ def residual(self, x: torch.Tensor, coords: torch.Tensor) -> torch.Tensor:
+ B, T, H, W, C = x.shape
+ if T > self.max_t:
+ raise ValueError(
+ f"Joint3DSelfAttention received T={T} but was constructed with "
+ f"max_t={self.max_t}. Increase bottleneck_3d_max_t."
+ )
+ if self.use_3d_rope and (H > self.max_h or W > self.max_w):
+ raise ValueError(
+ f"Joint3DSelfAttention received (H, W)=({H}, {W}) but the 3D "
+ f"RoPE was sized for (len_h, len_w)=({self.max_h}, {self.max_w}). "
+ f"This usually means the bottleneck spatial shape changed -- "
+ f"rebuild the decoder with larger len_h/len_w."
+ )
+ h = self.norm(x)
+
+ if self.use_3d_rope:
+ # 3D RoPE path. We flatten (T, H, W) to a single sequence axis
+ # so `apply_rotary_emb` can broadcast its (S, head_dim) rope
+ # tensor over batch + num_heads in one shot.
+ qkv = self.qkv_proj(h)
+ qkv = einops.rearrange(
+ qkv,
+ "B T H W (three N D) -> three B (T H W) N D",
+ three=3,
+ D=self.head_dim,
+ )
+ q, k, v = qkv[0], qkv[1], qkv[2]
+ q, k = self._scale_qk(q, k) # rotation preserves L2 norm
+ rope_emb = self.rope3d(T, H, W) # (T*H*W, head_dim)
+ q = apply_rotary_emb(q, rope_emb)
+ k = apply_rotary_emb(k, rope_emb)
+ q = einops.rearrange(q, "B S N D -> B N S D")
+ k = einops.rearrange(k, "B S N D -> B N S D")
+ v = einops.rearrange(v, "B S N D -> B N S D")
+ else:
+ # Original: learnable temporal PE + 2D AxialRoPE over (h, w).
+ h = h + self.temporal_pe[:T].view(1, T, 1, 1, C)
+ qkv = self.qkv_proj(h)
+ q, k, v = einops.rearrange(
+ qkv,
+ "B T H W (three N D) -> three B T H W N D",
+ three=3,
+ D=self.head_dim,
+ )
+ q, k = self._scale_qk(q, k)
+ q, k = self._apply_spatial_rope(q, k, coords)
+ q = einops.rearrange(q, "B T H W N D -> B N (T H W) D")
+ k = einops.rearrange(k, "B T H W N D -> B N (T H W) D")
+ v = einops.rearrange(v, "B T H W N D -> B N (T H W) D")
+
+ if self.causal_time:
+ S = H * W
+ t_idx = torch.arange(T, device=q.device).repeat_interleave(S)
+ # mask[i, j] = True iff query t_i can attend to key t_j (t_j <= t_i)
+ mask = t_idx.unsqueeze(0) <= t_idx.unsqueeze(1) # (T*S, T*S)
+ out = F.scaled_dot_product_attention(q, k, v, attn_mask=mask, scale=1.0)
+ else:
+ out = F.scaled_dot_product_attention(q, k, v, scale=1.0)
+
+ out = einops.rearrange(out, "B N (T H W) D -> B T H W (N D)", T=T, H=H, W=W)
+ out = self.dropout(out)
+ out = self.out_proj(out)
+ return out
+
+ def forward(self, x: torch.Tensor, coords: torch.Tensor) -> torch.Tensor:
+ return x + self.residual(x, coords)
+
+ def forward_stream(
+ self,
+ x: torch.Tensor,
+ coords: torch.Tensor,
+ kv_cache: tuple[torch.Tensor, torch.Tensor] | None = None,
+ ) -> tuple[torch.Tensor, tuple[torch.Tensor, torch.Tensor]]:
+ """Attend a new temporal chunk to a sliding causal KV cache.
+
+ WAN-style long-video path: only ``x`` (the new frames) is projected to
+ Q; K/V are concatenated with a cached prefix and trimmed to
+ ``max_t`` frames. RoPE / temporal PE are applied over the *window*
+ indices ``0 .. T_win-1``, so a pure prefix stream with
+ ``T_total <= max_t`` matches :meth:`forward` exactly.
+
+ Cache tensors are **unrotated** ``(B, N, T, H*W, D)`` so the window
+ can be re-indexed when older frames are dropped.
+
+ Requires ``causal_time=True``.
+ """
+ if not self.causal_time:
+ raise NotImplementedError("Joint3DSelfAttention.forward_stream requires causal_time=True")
+ B, T_new, H, W, C = x.shape
+ if T_new > self.max_t:
+ raise ValueError(
+ f"streaming chunk T={T_new} exceeds bottleneck_3d_max_t={self.max_t}; use a smaller chunk_frames"
+ )
+ del coords # 3D RoPE path does not use polar coords
+ if self.use_3d_rope and (H > self.max_h or W > self.max_w):
+ raise ValueError(
+ f"Joint3DSelfAttention received (H, W)=({H}, {W}) but the 3D "
+ f"RoPE was sized for (len_h, len_w)=({self.max_h}, {self.max_w})."
+ )
+
+ S = H * W
+ h = self.norm(x)
+ if not self.use_3d_rope:
+ raise NotImplementedError(
+ "Joint3DSelfAttention.forward_stream currently requires use_3d_rope=True "
+ "(learnable temporal PE streaming is not supported)"
+ )
+ qkv = self.qkv_proj(h)
+ qkv = einops.rearrange(
+ qkv,
+ "B T H W (three N D) -> three B (T H W) N D",
+ three=3,
+ D=self.head_dim,
+ )
+ q_flat, k_flat, v_flat = qkv[0], qkv[1], qkv[2]
+ q_flat, k_flat = self._scale_qk(q_flat, k_flat)
+ # Cache / stream in (B, N, T, S, D) so frames can be trimmed.
+ q_new = einops.rearrange(q_flat, "B (T S) N D -> B N T S D", T=T_new, S=S)
+ k_new = einops.rearrange(k_flat, "B (T S) N D -> B N T S D", T=T_new, S=S)
+ v_new = einops.rearrange(v_flat, "B (T S) N D -> B N T S D", T=T_new, S=S)
+
+ if kv_cache is not None:
+ k_past, v_past = kv_cache
+ k_all = torch.cat([k_past, k_new], dim=2)
+ v_all = torch.cat([v_past, v_new], dim=2)
+ else:
+ k_all, v_all = k_new, v_new
+
+ T_total = k_all.shape[2]
+ if T_total > self.max_t:
+ keep = self.max_t
+ k_all = k_all[:, :, -keep:].contiguous()
+ v_all = v_all[:, :, -keep:].contiguous()
+ T_total = keep
+
+ # Window-local temporal positions 0 .. T_total-1 (matches one-shot when
+ # the stream never exceeds max_t).
+ rope_full = self.rope3d(T_total, H, W) # (T_total*S, D)
+ rope_full = einops.rearrange(rope_full, "(T S) D -> T S D", T=T_total, S=S)
+ rope_q = rope_full[-T_new:].reshape(T_new * S, -1)
+ rope_k = rope_full.reshape(T_total * S, -1)
+ q = einops.rearrange(q_new, "B N T S D -> B (T S) N D")
+ k = einops.rearrange(k_all, "B N T S D -> B (T S) N D")
+ v = einops.rearrange(v_all, "B N T S D -> B (T S) N D")
+ q = apply_rotary_emb(q, rope_q)
+ k = apply_rotary_emb(k, rope_k)
+ q = einops.rearrange(q, "B S N D -> B N S D")
+ k = einops.rearrange(k, "B S N D -> B N S D")
+ v = einops.rearrange(v, "B S N D -> B N S D")
+
+ # With one newly appended frame, every cached key is at or before the
+ # query frame. Avoiding an explicit all-true mask lets SDPA use its
+ # memory-efficient backend; at LiDAR resolution the dense mask alone can
+ # otherwise consume several GiB. Multi-frame chunks still need temporal
+ # causality between their newly appended frames.
+ mask = None
+ if T_new > 1:
+ t_q = torch.arange(T_total - T_new, T_total, device=q.device).repeat_interleave(S)
+ t_k = torch.arange(T_total, device=q.device).repeat_interleave(S)
+ mask = t_k.unsqueeze(0) <= t_q.unsqueeze(1) # (T_new*S, T_total*S)
+ out = F.scaled_dot_product_attention(q, k, v, attn_mask=mask, scale=1.0)
+ out = einops.rearrange(out, "B N (T H W) D -> B T H W (N D)", T=T_new, H=H, W=W)
+ out = self.dropout(out)
+ out = self.out_proj(out)
+ return x + out, (k_all.detach(), v_all.detach())
+
+ def extra_repr(self) -> str:
+ pe = (
+ f"rope3d(len_t={self.max_t}, len_h={self.max_h}, len_w={self.max_w})"
+ if self.use_3d_rope
+ else f"max_t={self.max_t}, axial_rope_2d"
+ )
+ return f"head_dim={self.head_dim}, num_heads={self.num_heads}, {pe}, causal_time={self.causal_time}"
+
+
+class Bottleneck3DBlock(nn.Module):
+ """Joint 3D self-attention + FFN. Operates on (B, T, H, W, C).
+
+ Set ``use_3d_rope=True`` (with ``len_h`` / ``len_w`` set to the
+ bottleneck's spatial shape) to swap the default learnable temporal PE
+ + 2D AxialRoPE for a single 3D RoPE over ``(t, h, w)`` -- see
+ :class:`Joint3DSelfAttention` for the tradeoffs. Backward-compatible:
+ leaving ``use_3d_rope`` at its default keeps the original behavior.
+ """
+
+ def __init__(
+ self,
+ in_dim: int,
+ num_heads: int,
+ max_t: int,
+ rope_max_harmonics: list[int] = (1, 1),
+ mlp_ratio: float = 3.0,
+ dropout: float = 0.0,
+ causal_time: bool = False,
+ use_3d_rope: bool = False,
+ len_h: int = 1,
+ len_w: int = 1,
+ ):
+ super().__init__()
+ self.attn = Joint3DSelfAttention(
+ dim=in_dim,
+ num_heads=num_heads,
+ max_t=max_t,
+ rope_max_harmonics=rope_max_harmonics,
+ dropout=dropout,
+ causal_time=causal_time,
+ use_3d_rope=use_3d_rope,
+ len_h=len_h,
+ len_w=len_w,
+ )
+ self.ffn = FeedForwardNetwork(in_dim, int(in_dim * mlp_ratio), dropout)
+
+ def forward(self, x: torch.Tensor, coords: torch.Tensor) -> torch.Tensor:
+ x = self.attn(x, coords)
+ x = self.ffn(x)
+ return x
+
+ def forward_stream(
+ self,
+ x: torch.Tensor,
+ coords: torch.Tensor,
+ kv_cache: tuple[torch.Tensor, torch.Tensor] | None = None,
+ ) -> tuple[torch.Tensor, tuple[torch.Tensor, torch.Tensor]]:
+ x, kv_cache = self.attn.forward_stream(x, coords, kv_cache)
+ x = self.ffn(x)
+ return x, kv_cache
+
+
+# =============================================================================
+# Patch / temporal resampling
+# =============================================================================
+
+
+class PatchMerging(nn.Sequential):
+ """2x2 spatial downsample: (B, 2H, 2W, C) -> (B, H, W, 2C)"""
+
+ def __init__(self, dim: int):
+ super().__init__(
+ Rearrange("B (H P1) (W P2) C -> B H W (P1 P2 C)", P1=2, P2=2),
+ nn.Linear(4 * dim, 2 * dim, bias=False),
+ )
+
+
+class PatchExpanding(nn.Module):
+ """2x2 spatial upsample: (B, H, W, C) -> (B, 2H, 2W, C//2)"""
+
+ def __init__(self, dim: int):
+ super().__init__()
+ self.linear = nn.Linear(dim, dim * 2, bias=False)
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ x = self.linear(x)
+ x = einops.rearrange(x, "B H W (P1 P2 C) -> B (H P1) (W P2) C", P1=2, P2=2)
+ return x
+
+
+class TemporalMerging(nn.Module):
+ """Causal 2x temporal downsample (mimics WanVAE downsample3d).
+
+ Uses Conv1d(kernel=3, stride=2) with left-only causal padding.
+ Each output only depends on the current and past frames.
+ Input: (B, T, H, W, C)
+ Output: (B, ceil(T/2), H, W, C)
+ """
+
+ def __init__(self, dim: int):
+ super().__init__()
+ self.conv = nn.Conv1d(dim, dim, kernel_size=3, stride=2, padding=0, bias=False)
+ self.causal_pad = 2 # kernel_size - 1
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ B, T, H, W, C = x.shape
+ x = einops.rearrange(x, "B T H W C -> (B H W) C T")
+ x = F.pad(x, (self.causal_pad, 0))
+ x = self.conv(x)
+ x = einops.rearrange(x, "(B H W) C T -> B T H W C", B=B, H=H, W=W)
+ return x
+
+
+class TemporalMergingFirstFrameSpecial(nn.Module):
+ """First-frame-special causal 2x temporal downsample (WAN-style).
+
+ Mirrors :class:`WanVAE2P1`'s singleton-first-frame behavior at the
+ transformer-VAE merge stage: the first input slot becomes ``z[0]``
+ via a per-frame ``Linear`` (no temporal mixing), and the remaining
+ ``T_in - 1`` frames are merged in *disjoint* 2-frame chunks via a
+ stride-2 ``Conv1d`` with ``kernel=2`` and **no** causal padding --
+ each pair stays inside its own chunk:
+
+ ``in[0:1]`` -> z[0] per-frame Linear
+ ``in[2i-1:2i+1]`` -> z[i>=1] stride-2 kernel-2 conv
+
+ Pair this with :class:`TemporalExpandingWanStyle` in the decoder so
+ ``T_decoded = 2*T_z - 1`` exactly matches ``T_in = 1 + 2*(T_z - 1)``.
+ No front-crop needed, no boundary information loss.
+
+ Compared to plain :class:`TemporalMerging` (kernel=3, stride=2,
+ causal-pad-2):
+ * z[0] genuinely depends on in[0] alone (not "in[0] + 2 pad"
+ activations); the Linear is the cleanest analog of WAN's CausalConv
+ with a fully zero-cached history.
+ * z[i>=1] sees both frames of its 2-frame chunk (vs. only the first
+ frame + 2 history frames in plain merging) -- so each latent has
+ *full* info on the chunk it represents, eliminating the
+ "z[i] never sees in[2i+1]" information loss.
+
+ Requires ``T_in`` to be ODD (i.e. ``T_in - 1`` divisible by 2).
+
+ Input: (B, T_in, H, W, C), T_in = 2*T_z - 1
+ Output: (B, T_z, H, W, C)
+ """
+
+ def __init__(self, dim: int):
+ super().__init__()
+ # First-frame projection: like WAN's CausalConv3d on a length-1
+ # input with zero-padded kernel history -- a linear function of
+ # the single frame's features.
+ self.first_frame_proj = nn.Linear(dim, dim, bias=False)
+ # Disjoint pair merger: kernel=2, stride=2, no padding. Each
+ # output index i corresponds to ``rest[2i, 2i+1]`` exclusively
+ # (no kernel overlap across chunks, no causal-pad zero leakage).
+ self.conv = nn.Conv1d(dim, dim, kernel_size=2, stride=2, padding=0, bias=False)
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ B, T_in, H, W, C = x.shape
+ if T_in == 1:
+ return self.first_frame_proj(x)
+ first = self.first_frame_proj(x[:, :1]) # (B, 1, H, W, C)
+ rest = x[:, 1:]
+ if rest.shape[1] % 2 != 0:
+ raise ValueError(
+ f"TemporalMergingFirstFrameSpecial requires (T_in - 1) "
+ f"divisible by 2 at this level, got T_in={T_in}. With "
+ f"``temporal_first_frame_special=True`` and N levels of "
+ f"temporal merge, the input T must satisfy "
+ f"T = 1 + (2 ** N) * (T_z - 1) -- e.g. T=9 for N=1, T_z=5."
+ )
+ rest = einops.rearrange(rest, "B T H W C -> (B H W) C T")
+ rest = self.conv(rest) # (N, C, (T_in - 1) // 2)
+ rest = einops.rearrange(rest, "(B H W) C T -> B T H W C", B=B, H=H, W=W)
+ return torch.cat([first, rest], dim=1)
+
+
+class TemporalExpanding(nn.Module):
+ """Causal 2x temporal upsample (mimics WanVAE upsample3d).
+
+ Uses a causal Conv1d that doubles channels, then interleaves the two
+ channel halves along time to achieve 2x upsampling. No future leakage
+ because the underlying conv is causal (left-pad only, stride=1).
+ Input: (B, T//2, H, W, C)
+ Output: (B, T, H, W, C)
+ """
+
+ def __init__(self, dim: int):
+ super().__init__()
+ self.conv = nn.Conv1d(dim, dim * 2, kernel_size=3, stride=1, padding=0, bias=False)
+ self.causal_pad = 2 # kernel_size - 1
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ B, T_in, H, W, C = x.shape
+ x = einops.rearrange(x, "B T H W C -> (B H W) C T")
+ x = F.pad(x, (self.causal_pad, 0))
+ x = self.conv(x) # (N, 2C, T_in)
+ x = einops.rearrange(
+ x,
+ "N (two C) T -> N C (T two)",
+ two=2,
+ ) # interleave: [t0_a, t0_b, t1_a, t1_b, ...] → 2*T_in frames
+ x = einops.rearrange(x, "(B H W) C T -> B T H W C", B=B, H=H, W=W)
+ return x
+
+
+class TemporalExpandingWanStyle(nn.Module):
+ """Causal 2x temporal upsample with WAN-style first-frame handling.
+
+ Mirrors :class:`WanVAE2P1.decode`: the first input frame produces a
+ single output frame (just the "a" half of the causal conv's doubled
+ output -- i.e. no temporal duplication), while every subsequent input
+ frame produces 2 output frames. Composed twice (matching the WAN
+ encoder's 4x temporal compression with its special first-frame chunk),
+ a latent of length ``T_z`` decodes to ``4*T_z - 3`` output frames -- the
+ same length as the original input when ``T_in = 1 + 4*(T_z - 1)``, so
+ no leading-frame crop is needed.
+
+ Input: (B, T_in, H, W, C)
+ Output: (B, 2*T_in - 1, H, W, C) (or 1 frame if T_in == 1)
+ """
+
+ def __init__(self, dim: int):
+ super().__init__()
+ self.conv = nn.Conv1d(dim, dim * 2, kernel_size=3, stride=1, padding=0, bias=False)
+ self.causal_pad = 2 # kernel_size - 1
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ B, T_in, H, W, C = x.shape
+ h = einops.rearrange(x, "B T H W C -> (B H W) C T")
+ h = F.pad(h, (self.causal_pad, 0))
+ h = self.conv(h) # (N, 2C, T_in)
+ h = einops.rearrange(
+ h,
+ "N (two C) T -> N C (T two)",
+ two=2,
+ ) # interleave: [t0_a, t0_b, t1_a, t1_b, ...] → 2*T_in frames along T
+ h = einops.rearrange(h, "(B H W) C T -> B T H W C", B=B, H=H, W=W)
+ # Keep t0_a as the single "first frame" output; drop t0_b. Keep both
+ # halves for every other input frame.
+ if T_in == 1:
+ return h[:, :1]
+ return torch.cat([h[:, :1], h[:, 2:]], dim=1)
+
+
+class Tokenizer(nn.Sequential):
+ def __init__(
+ self,
+ in_channels: int,
+ out_channels: int,
+ patch_size: list[int],
+ ):
+ patch_size = _pair(patch_size)
+ super().__init__(
+ nn.Conv2d(
+ in_channels,
+ out_channels,
+ kernel_size=patch_size,
+ stride=patch_size,
+ padding=0,
+ bias=False,
+ ),
+ Rearrange("B C H W -> B H W C"),
+ )
+
+
+class Detokenizer(nn.Sequential):
+ def __init__(
+ self,
+ in_channels: int,
+ out_channels: int,
+ patch_size: list[int],
+ ):
+ patch_size = _pair(patch_size)
+ super().__init__(
+ RMSNorm(in_channels),
+ nn.Linear(
+ in_channels,
+ out_channels * patch_size[0] * patch_size[1],
+ bias=False,
+ ),
+ Rearrange(
+ "B H W (P1 P2 C) -> B C (H P1) (W P2)",
+ P1=patch_size[0],
+ P2=patch_size[1],
+ ),
+ )
+
+
+# =============================================================================
+# Positional embedding
+# =============================================================================
+
+
+class LearnablePositionalEmbedding(nn.Module):
+ def __init__(self, out_dim: int, resolution: list[int]):
+ super().__init__()
+ self.embedding = nn.Parameter(torch.zeros(1, *resolution, out_dim))
+ nn.init.trunc_normal_(self.embedding, std=0.02)
+
+ def forward(self, coords=None):
+ return self.embedding
+
+ def extra_repr(self):
+ return f"resolution={tuple(self.embedding.shape[1:3])}, out_dim={self.embedding.shape[3]}"
+
+
+# =============================================================================
+# Encoder
+# =============================================================================
+
+
+class Encoder(nn.Module):
+ def __init__(
+ self,
+ resolution: list[int],
+ in_channels: int,
+ z_dim: int,
+ base_channels: int = 128,
+ patch_size: list[int] = (1, 4),
+ window_size: list[int] = (3, 9),
+ depths: list[int] = (3, 3, 3, 3),
+ num_heads: list[int] = (2, 4, 8, 16),
+ dilation: list[int] = (1, 1, 1, 1),
+ temporal_downsample: list[bool] = (True, False, False),
+ mlp_ratio: float = 3.0,
+ dropout: float = 0.0,
+ mapping_depth: int = 2,
+ positional_embedding: str = "learnable_embedding",
+ # If True, every level with ``temporal_downsample[i]=True`` uses
+ # :class:`TemporalMergingFirstFrameSpecial` instead of the plain
+ # :class:`TemporalMerging`. Yields clean WAN-style chunk
+ # boundaries (z[0]=enc(in[0]), z[i>=1]=enc(disjoint 2-frame
+ # chunk)) at the cost of requiring T_in - 1 to be divisible
+ # by ``2 ** sum(temporal_downsample)``.
+ temporal_first_frame_special: bool = False,
+ # ---- Bottleneck options (mirror Decoder) ----
+ # If True, replace the encoder bottleneck's factorized (global spatial
+ # + temporal) attention with a single joint 3D self-attention over
+ # (T*H*W) tokens. Down-levels are unchanged. Default False keeps the
+ # original factorized bottleneck (bit-identical to existing configs).
+ bottleneck_3d: bool = False,
+ bottleneck_3d_max_t: int = 32,
+ bottleneck_3d_causal_time: bool = False,
+ bottleneck_3d_rope: bool = False,
+ # If False, the local (neighborhood) attention does NOT wrap the W
+ # axis (clamps at the edges instead). Default True preserves the
+ # range-image azimuth-periodic behavior. Set False for a Cartesian
+ # BEV grid where the left/right W edges are unrelated.
+ circular_padding: bool = True,
+ # Temporal-mixing mechanism for the down/mid temporal blocks:
+ # "attention" (:class:`CausalTemporalAttention`, default) or "conv"
+ # (WAN-style :class:`CausalTemporalConv`, kernel ``temporal_conv_kernel``).
+ temporal_mixer: Literal["attention", "conv"] = "attention",
+ temporal_conv_kernel: int = 3,
+ ):
+ super().__init__()
+ self.resolution = _pair(resolution)
+ self.patch_size = _pair(patch_size)
+ self.depths = depths
+ self.temporal_downsample = temporal_downsample
+ self.bottleneck_3d = bottleneck_3d
+ self._temporal_cls = (
+ functools.partial(TemporalConvBlock, kernel_size=temporal_conv_kernel)
+ if temporal_mixer == "conv"
+ else TemporalBlock
+ )
+
+ token_size = torch.tensor(self.resolution) // torch.tensor(self.patch_size)
+ max_harmonics = (token_size / 2).int()
+
+ self.tokenizer = Tokenizer(in_channels, base_channels, patch_size)
+
+ if positional_embedding != "learnable_embedding":
+ raise ValueError(
+ "Only positional_embedding='learnable_embedding' is supported "
+ f"(got {positional_embedding!r}). Alternate absolute PE modes were removed."
+ )
+ # mapping_depth is retained for config compatibility with the shipped checkpoint.
+ _ = mapping_depth
+ self.spatial_pe = LearnablePositionalEmbedding(
+ out_dim=base_channels,
+ resolution=token_size.tolist(),
+ )
+
+ # Down levels: spatial blocks + temporal blocks + merging
+ # Skip temporal attention at level 0 (highest resolution) for efficiency.
+ self.down_levels = nn.ModuleDict()
+ for i, num_blocks in enumerate(depths[:-1]):
+ dim_i = base_channels << i
+ spatial_blocks = nn.ModuleList()
+ for j in range(num_blocks):
+ spatial_blocks.append(
+ SpatialBlock(
+ in_dim=dim_i,
+ num_heads=num_heads[i],
+ attn_type="local",
+ kernel_size=window_size,
+ dilation=1 if j % 2 == 0 else dilation[i],
+ mlp_ratio=mlp_ratio,
+ dropout=dropout,
+ rope_max_harmonics=(max_harmonics >> i).clamp(min=1),
+ circular=circular_padding,
+ )
+ )
+ self.down_levels[f"spatial_{i}"] = spatial_blocks
+ if i > 0:
+ temporal_blocks = nn.ModuleList()
+ for j in range(num_blocks):
+ temporal_blocks.append(
+ self._temporal_cls(
+ in_dim=dim_i,
+ num_heads=num_heads[i],
+ mlp_ratio=mlp_ratio,
+ dropout=dropout,
+ )
+ )
+ self.down_levels[f"temporal_{i}"] = temporal_blocks
+ self.down_levels[f"merge_{i}"] = PatchMerging(dim_i)
+ if temporal_downsample[i]:
+ merge_cls = TemporalMergingFirstFrameSpecial if temporal_first_frame_special else TemporalMerging
+ self.down_levels[f"temporal_merge_{i}"] = merge_cls(dim_i * 2)
+
+ # Bottleneck: either factorized (global spatial + temporal), or a
+ # single joint 3D self-attention over (T*H*W) tokens.
+ i = len(depths) - 1
+ bottleneck_dim = base_channels << i
+ bottleneck_size = (token_size >> i).tolist()
+ if bottleneck_3d:
+ self.mid_3d = nn.ModuleList()
+ for _ in range(depths[-1]):
+ self.mid_3d.append(
+ Bottleneck3DBlock(
+ in_dim=bottleneck_dim,
+ num_heads=num_heads[-1],
+ max_t=bottleneck_3d_max_t,
+ rope_max_harmonics=(max_harmonics >> i).clamp(min=1),
+ mlp_ratio=mlp_ratio,
+ dropout=dropout,
+ causal_time=bottleneck_3d_causal_time,
+ use_3d_rope=bottleneck_3d_rope,
+ len_h=bottleneck_size[0],
+ len_w=bottleneck_size[1],
+ )
+ )
+ else:
+ self.mid_spatial = nn.ModuleList()
+ self.mid_temporal = nn.ModuleList()
+ for j in range(depths[-1]):
+ self.mid_spatial.append(
+ SpatialBlock(
+ in_dim=bottleneck_dim,
+ num_heads=num_heads[-1],
+ attn_type="global",
+ mlp_ratio=mlp_ratio,
+ dropout=dropout,
+ rope_max_harmonics=(max_harmonics >> i).clamp(min=1),
+ )
+ )
+ self.mid_temporal.append(
+ self._temporal_cls(
+ in_dim=bottleneck_dim,
+ num_heads=num_heads[-1],
+ mlp_ratio=mlp_ratio,
+ dropout=dropout,
+ )
+ )
+
+ # Project to z_dim * 2 (mean + logvar)
+ self.head = nn.Sequential(
+ RMSNorm(bottleneck_dim),
+ nn.Linear(bottleneck_dim, z_dim * 2, bias=False),
+ Rearrange("B H W C -> B C H W"),
+ )
+
+ def forward(self, x: torch.Tensor, coords: torch.Tensor) -> torch.Tensor:
+ """
+ Args:
+ x: (B, C, T, H, W) video tensor
+ coords: (1, 2, H, W) polar coordinates
+ Returns:
+ (B, z_dim*2, T', H_z, W_z) latent parameters
+ """
+ B, C, T, H, W = x.shape
+ c = F.avg_pool2d(coords, kernel_size=self.patch_size, stride=self.patch_size)
+
+ # Tokenize per-frame
+ x = einops.rearrange(x, "B C T H W -> (B T) C H W")
+ h = self.tokenizer(x) + self.spatial_pe(c) # (BT, H', W', C0)
+ h = einops.rearrange(h, "(B T) H W C -> B T H W C", B=B, T=T)
+
+ T_cur = T
+ for i in range(len(self.depths) - 1):
+ _, _, Hc, Wc, _ = h.shape
+ has_temporal = f"temporal_{i}" in self.down_levels
+ t_blocks = (
+ self.down_levels[f"temporal_{i}"] if has_temporal else [None] * len(self.down_levels[f"spatial_{i}"])
+ )
+ for s_block, t_block in zip(self.down_levels[f"spatial_{i}"], t_blocks):
+ h = einops.rearrange(h, "B T H W C -> (B T) H W C")
+ h = s_block(h, c)
+ if t_block is not None:
+ h = einops.rearrange(h, "(B T) H W C -> (B H W) T C", B=B, T=T_cur)
+ h = t_block(h)
+ h = einops.rearrange(h, "(B H W) T C -> B T H W C", B=B, H=Hc, W=Wc)
+ else:
+ h = einops.rearrange(h, "(B T) H W C -> B T H W C", B=B, T=T_cur)
+
+ # Spatial downsample per-frame
+ h = einops.rearrange(h, "B T H W C -> (B T) H W C")
+ h = self.down_levels[f"merge_{i}"](h)
+ h = einops.rearrange(h, "(B T) H W C -> B T H W C", B=B, T=T_cur)
+ c = F.avg_pool2d(c, kernel_size=2, stride=2)
+
+ # Temporal downsample
+ if self.temporal_downsample[i]:
+ h = self.down_levels[f"temporal_merge_{i}"](h)
+ T_cur = h.shape[1]
+
+ # Bottleneck: joint 3D attention, or factorized spatial-then-temporal.
+ _, _, Hc, Wc, _ = h.shape
+ if self.bottleneck_3d:
+ for block in self.mid_3d:
+ h = block(h, c)
+ else:
+ for s_block, t_block in zip(self.mid_spatial, self.mid_temporal):
+ h = einops.rearrange(h, "B T H W C -> (B T) H W C")
+ h = s_block(h, c)
+ h = einops.rearrange(h, "(B T) H W C -> (B H W) T C", B=B, T=T_cur)
+ h = t_block(h)
+ h = einops.rearrange(h, "(B H W) T C -> B T H W C", B=B, H=Hc, W=Wc)
+
+ # Head per-frame
+ h = einops.rearrange(h, "B T H W C -> (B T) H W C")
+ h = self.head(h) # (BT, z_dim*2, H_z, W_z)
+ h = einops.rearrange(h, "(B T) C H W -> B C T H W", B=B, T=T_cur)
+ return h
+
+ def forward_stream(
+ self,
+ x: torch.Tensor,
+ coords: torch.Tensor,
+ temporal_kv_cache: dict[str, tuple[torch.Tensor, torch.Tensor]] | None = None,
+ ) -> tuple[torch.Tensor, dict[str, tuple[torch.Tensor, torch.Tensor]]]:
+ """Encode a new temporal chunk with causal KV caches (WAN-style).
+
+ Spatial ops run only on ``x``; temporal / joint-3D layers attend to the
+ cached prefix. Requires no temporal downsampling (1x temporal latent).
+ """
+ if any(self.temporal_downsample):
+ raise NotImplementedError("encoder streaming requires temporal_downsample=False everywhere")
+ B, C, T, H, W = x.shape
+ old_cache = temporal_kv_cache or {}
+ new_cache: dict[str, tuple[torch.Tensor, torch.Tensor]] = {}
+ c = F.avg_pool2d(coords, kernel_size=self.patch_size, stride=self.patch_size)
+
+ x = einops.rearrange(x, "B C T H W -> (B T) C H W")
+ h = self.tokenizer(x) + self.spatial_pe(c)
+ h = einops.rearrange(h, "(B T) H W C -> B T H W C", B=B, T=T)
+
+ T_cur = T
+ for i in range(len(self.depths) - 1):
+ _, _, Hc, Wc, _ = h.shape
+ has_temporal = f"temporal_{i}" in self.down_levels
+ t_blocks = (
+ self.down_levels[f"temporal_{i}"] if has_temporal else [None] * len(self.down_levels[f"spatial_{i}"])
+ )
+ for block_idx, (s_block, t_block) in enumerate(zip(self.down_levels[f"spatial_{i}"], t_blocks)):
+ h = einops.rearrange(h, "B T H W C -> (B T) H W C")
+ h = s_block(h, c)
+ if t_block is not None:
+ if not isinstance(t_block, TemporalBlock):
+ raise NotImplementedError(
+ "encoder streaming currently supports temporal attention, not temporal convolution"
+ )
+ h = einops.rearrange(h, "(B T) H W C -> (B H W) T C", B=B, T=T_cur)
+ cache_key = f"down_levels.temporal_{i}.{block_idx}"
+ h, new_cache[cache_key] = t_block.forward_stream(h, old_cache.pop(cache_key, None))
+ h = einops.rearrange(h, "(B H W) T C -> B T H W C", B=B, H=Hc, W=Wc)
+ else:
+ h = einops.rearrange(h, "(B T) H W C -> B T H W C", B=B, T=T_cur)
+
+ h = einops.rearrange(h, "B T H W C -> (B T) H W C")
+ h = self.down_levels[f"merge_{i}"](h)
+ h = einops.rearrange(h, "(B T) H W C -> B T H W C", B=B, T=T_cur)
+ c = F.avg_pool2d(c, kernel_size=2, stride=2)
+
+ _, _, Hc, Wc, _ = h.shape
+ if self.bottleneck_3d:
+ for block_idx, block in enumerate(self.mid_3d):
+ cache_key = f"mid_3d.{block_idx}"
+ h, new_cache[cache_key] = block.forward_stream(h, c, old_cache.pop(cache_key, None))
+ else:
+ for block_idx, (s_block, t_block) in enumerate(zip(self.mid_spatial, self.mid_temporal)):
+ if not isinstance(t_block, TemporalBlock):
+ raise NotImplementedError(
+ "encoder streaming currently supports temporal attention, not temporal convolution"
+ )
+ h = einops.rearrange(h, "B T H W C -> (B T) H W C")
+ h = s_block(h, c)
+ h = einops.rearrange(h, "(B T) H W C -> (B H W) T C", B=B, T=T_cur)
+ cache_key = f"mid_temporal.{block_idx}"
+ h, new_cache[cache_key] = t_block.forward_stream(h, old_cache.pop(cache_key, None))
+ h = einops.rearrange(h, "(B H W) T C -> B T H W C", B=B, H=Hc, W=Wc)
+
+ h = einops.rearrange(h, "B T H W C -> (B T) H W C")
+ h = self.head(h)
+ h = einops.rearrange(h, "(B T) C H W -> B C T H W", B=B, T=T_cur)
+ missing = set(old_cache) - set(new_cache)
+ if missing:
+ raise ValueError(f"unused temporal KV cache keys: {sorted(missing)}")
+ return h, new_cache
+
+
+class Decoder(nn.Module):
+ def __init__(
+ self,
+ resolution: list[int],
+ out_channels: int,
+ z_dim: int,
+ base_channels: int = 128,
+ patch_size: list[int] = (1, 4),
+ window_size: list[int] = (3, 9),
+ depths: list[int] = (3, 3, 3, 3),
+ num_heads: list[int] = (2, 4, 8, 16),
+ dilation: list[int] = (1, 1, 1, 1),
+ temporal_downsample: list[bool] = (True, False, False),
+ temporal_upsample: list[bool] | None = None,
+ mlp_ratio: float = 3.0,
+ dropout: float = 0.0,
+ mapping_depth: int = 2,
+ positional_embedding: str = "learnable_embedding",
+ stem_patchify: bool = False,
+ # ---- Bottleneck options ----
+ # If True, replace the bottleneck's factorized (temporal + global
+ # spatial) attention with a single joint 3D self-attention over
+ # (T*H*W) tokens. The up-levels are unchanged (still skip temporal
+ # attn at the highest resolution as before).
+ bottleneck_3d: bool = False,
+ # Max T accepted by the bottleneck-3D's learnable temporal PE. Only
+ # used when bottleneck_3d=True. Set this >= the largest T_z (encoder
+ # latent frames) you'll ever feed the decoder.
+ bottleneck_3d_max_t: int = 32,
+ # If True, the joint 3D attention is causal in time (token (t, h, w)
+ # only attends to (t', h', w') with t' <= t). The causal mask is
+ # materialized as a (T*H*W, T*H*W) bool tensor, so memory scales as
+ # O((T*H*W)^2). Bidirectional is the default for non-streaming use.
+ bottleneck_3d_causal_time: bool = False,
+ # If True (and bottleneck_3d=True), encode position in the joint 3D
+ # attention with a parameter-free 3D RoPE (`VideoRopePosition3DEmb`)
+ # over (t, h, w) instead of the default (learnable temporal PE + 2D
+ # spatial AxialRoPE). Recommended for length-extrapolation: the
+ # learnable temporal PE has uninitialized slots past the largest T
+ # seen in training (= noise at inference); 3D RoPE has none. The
+ # spatial axes of the 3D RoPE replace AxialRoPE -- the bottleneck
+ # operates on a regular (H_z, W_z) grid so Cartesian RoPE is fine
+ # and decouples positional encoding from polar `coords`.
+ # Orthogonal to `bottleneck_3d_causal_time` -- both can be combined.
+ bottleneck_3d_rope: bool = False,
+ # If True, replace each TemporalExpanding with TemporalExpandingWanStyle,
+ # which mirrors the WAN decoder's first-frame-special behavior:
+ # T_in -> 2*T_in - 1 per level. Composed twice this maps T_z latents
+ # to 4*T_z - 3 frames, so when paired with the WAN encoder there's no
+ # over-production and no leading-frame crop needed downstream.
+ temporal_expand_wan_style: bool = False,
+ # ---- Detokenizer (output) options ----
+ # Patch size for the final unpatchify (pixel-shuffle) layer. Defaults
+ # to `patch_size` (symmetric with the stem, the original behavior).
+ # Setting this larger than `patch_size` lets the decoder run the top
+ # (highest-resolution) up-levels at a coarser feature grid and rely
+ # on the final per-pixel Linear + rearrange to upsample to the output
+ # resolution -- substantially cheaper since the top level dominates
+ # the activation memory. The user is responsible for making sure
+ # `bottleneck_size = resolution // out_patch_size >> n_down` matches
+ # the encoder's latent spatial shape (i.e. reduce `depths` by
+ # `log2(out_patch_size / patch_size)` to compensate).
+ out_patch_size: list[int] | None = None,
+ # If False, the local (neighborhood) attention does NOT wrap the W
+ # axis (clamps at the edges). Default True keeps the range-image
+ # azimuth-periodic behavior. Set False for a Cartesian BEV grid.
+ circular_padding: bool = True,
+ # Temporal-mixing mechanism for the mid/up temporal blocks:
+ # "attention" (:class:`CausalTemporalAttention`, default) or "conv"
+ # (WAN-style :class:`CausalTemporalConv`, kernel ``temporal_conv_kernel``).
+ temporal_mixer: Literal["attention", "conv"] = "attention",
+ temporal_conv_kernel: int = 3,
+ ):
+ super().__init__()
+ self.resolution = _pair(resolution)
+ self.patch_size = _pair(patch_size)
+ self.out_patch_size = _pair(out_patch_size) if out_patch_size is not None else self.patch_size
+ self.depths = depths
+ self.temporal_upsample = temporal_upsample if temporal_upsample is not None else temporal_downsample
+ self.stem_patchify = stem_patchify
+ self.bottleneck_3d = bottleneck_3d
+ self._temporal_cls = (
+ functools.partial(TemporalConvBlock, kernel_size=temporal_conv_kernel)
+ if temporal_mixer == "conv"
+ else TemporalBlock
+ )
+ self.temporal_expand_wan_style = temporal_expand_wan_style
+ temporal_expand_cls = TemporalExpandingWanStyle if temporal_expand_wan_style else TemporalExpanding
+
+ # `token_size` is the feature grid at the top up-level, i.e. *before*
+ # the final detokenizer's unpatchify -- this is what RoPE harmonics
+ # and the coordinate pyramid index. With `out_patch_size > patch_size`
+ # the top up-level sits at a coarser grid than the output resolution.
+ token_size = torch.tensor(self.resolution) // torch.tensor(self.out_patch_size)
+ max_harmonics = (token_size / 2).int()
+ n_down = len(depths) - 1
+ bottleneck_dim = base_channels << n_down
+ bottleneck_size = (token_size >> n_down).tolist()
+
+ # Project from z_dim to bottleneck channels.
+ # By default the stem is a per-pixel channel projection (used in the
+ # full TransformerVAE, where the Encoder's Tokenizer has already
+ # patchified the input by `patch_size`). When `stem_patchify=True`,
+ # we instead apply a Conv2d with stride=patch_size, so the decoder
+ # can be paired with an external encoder (e.g. the WAN encoder) whose
+ # latent spatial shape is patch_size larger than the bottleneck.
+ if stem_patchify and tuple(self.patch_size) != (1, 1):
+ self.stem = nn.Sequential(
+ nn.Conv2d(
+ z_dim,
+ bottleneck_dim,
+ kernel_size=self.patch_size,
+ stride=self.patch_size,
+ bias=False,
+ ),
+ Rearrange("B C H W -> B H W C"),
+ )
+ else:
+ self.stem = nn.Sequential(
+ Rearrange("B C H W -> B H W C"),
+ nn.Linear(z_dim, bottleneck_dim, bias=False),
+ )
+
+ # Positional embedding at bottleneck resolution
+ if positional_embedding != "learnable_embedding":
+ raise ValueError(
+ "Only positional_embedding='learnable_embedding' is supported "
+ f"(got {positional_embedding!r}). Alternate absolute PE modes were removed."
+ )
+ # mapping_depth is retained for config compatibility with the shipped checkpoint.
+ _ = mapping_depth
+ self.spatial_pe = LearnablePositionalEmbedding(
+ out_dim=bottleneck_dim,
+ resolution=bottleneck_size,
+ )
+
+ # Bottleneck: either factorized (temporal + global spatial), or a
+ # single joint 3D self-attention over (T*H*W) tokens.
+ if bottleneck_3d:
+ self.mid_3d = nn.ModuleList()
+ for _ in range(depths[-1]):
+ self.mid_3d.append(
+ Bottleneck3DBlock(
+ in_dim=bottleneck_dim,
+ num_heads=num_heads[-1],
+ max_t=bottleneck_3d_max_t,
+ rope_max_harmonics=(max_harmonics >> n_down).clamp(min=1),
+ mlp_ratio=mlp_ratio,
+ dropout=dropout,
+ causal_time=bottleneck_3d_causal_time,
+ use_3d_rope=bottleneck_3d_rope,
+ # 3D RoPE caches per-axis frequencies up to (len_t,
+ # len_h, len_w). Spatial caps come from the
+ # bottleneck's own grid; the temporal cap is shared
+ # with the learnable-PE path via `bottleneck_3d_max_t`.
+ len_h=bottleneck_size[0],
+ len_w=bottleneck_size[1],
+ )
+ )
+ else:
+ self.mid_temporal = nn.ModuleList()
+ self.mid_spatial = nn.ModuleList()
+ for _ in range(depths[-1]):
+ self.mid_temporal.append(
+ self._temporal_cls(
+ in_dim=bottleneck_dim,
+ num_heads=num_heads[-1],
+ mlp_ratio=mlp_ratio,
+ dropout=dropout,
+ )
+ )
+ self.mid_spatial.append(
+ SpatialBlock(
+ in_dim=bottleneck_dim,
+ num_heads=num_heads[-1],
+ attn_type="global",
+ mlp_ratio=mlp_ratio,
+ dropout=dropout,
+ rope_max_harmonics=(max_harmonics >> n_down).clamp(min=1),
+ )
+ )
+
+ # Up levels: temporal expand + spatial expand + temporal→spatial blocks
+ # Skip temporal attention at level 0 (highest resolution) for efficiency.
+ self.up_levels = nn.ModuleDict()
+ for i in reversed(range(n_down)):
+ dim_i = base_channels << i
+ dim_above = base_channels << (i + 1)
+
+ if self.temporal_upsample[i]:
+ self.up_levels[f"temporal_expand_{i}"] = temporal_expand_cls(dim_above)
+
+ self.up_levels[f"expand_{i}"] = PatchExpanding(dim_above)
+
+ spatial_blocks = nn.ModuleList()
+ for j in range(depths[i]):
+ spatial_blocks.append(
+ SpatialBlock(
+ in_dim=dim_i,
+ num_heads=num_heads[i],
+ attn_type="local",
+ kernel_size=window_size,
+ dilation=1 if j % 2 == 0 else dilation[i],
+ mlp_ratio=mlp_ratio,
+ dropout=dropout,
+ rope_max_harmonics=(max_harmonics >> i).clamp(min=1),
+ circular=circular_padding,
+ )
+ )
+ self.up_levels[f"spatial_{i}"] = spatial_blocks
+ # Temporal blocks are needed at every level that does a temporal
+ # expand (to mix the newly-interpolated frames into a coherent
+ # sequence), plus -- as a cheap-but-useful compute heuristic --
+ # at every level deeper than the top one. The previous condition
+ # was just `i > 0`, which silently skipped temporal attention at
+ # the top level even when temporal_upsample[0] was True, leaving
+ # the just-expanded frames un-mixed.
+ if i > 0 or self.temporal_upsample[i]:
+ temporal_blocks = nn.ModuleList()
+ for j in range(depths[i]):
+ temporal_blocks.append(
+ self._temporal_cls(
+ in_dim=dim_i,
+ num_heads=num_heads[i],
+ mlp_ratio=mlp_ratio,
+ dropout=dropout,
+ )
+ )
+ self.up_levels[f"temporal_{i}"] = temporal_blocks
+
+ self.detokenizer = Detokenizer(base_channels, out_channels, self.out_patch_size)
+
+ def forward(
+ self,
+ z: torch.Tensor,
+ coords: torch.Tensor,
+ temporal_kv_cache: dict[str, tuple[torch.Tensor, torch.Tensor]] | None = None,
+ return_temporal_kv_cache: bool = False,
+ ) -> torch.Tensor | tuple[torch.Tensor, dict[str, tuple[torch.Tensor, torch.Tensor]]]:
+ """
+ Args:
+ z: (B, z_dim, T', H_z, W_z) latent
+ coords: (1, 2, H, W) polar coordinates
+ Returns:
+ (B, C, T, H, W) reconstructed video
+ """
+ B = z.shape[0]
+ T_cur = z.shape[2]
+ n_down = len(self.depths) - 1
+ streaming = return_temporal_kv_cache
+ if temporal_kv_cache is not None and not streaming:
+ raise ValueError("temporal_kv_cache requires return_temporal_kv_cache=True")
+ if streaming:
+ if any(self.temporal_upsample):
+ raise NotImplementedError(
+ "streaming decode currently requires temporal_upsample=False at every decoder level"
+ )
+ if self.bottleneck_3d:
+ if not all(
+ getattr(block.attn, "causal_time", False) and getattr(block.attn, "use_3d_rope", False)
+ for block in self.mid_3d
+ ):
+ raise NotImplementedError(
+ "streaming decode with bottleneck_3d requires causal_time=True and use_3d_rope=True"
+ )
+ elif not all(isinstance(block, TemporalBlock) for block in self.mid_temporal):
+ raise NotImplementedError(
+ "streaming decode currently supports temporal attention, not temporal convolution"
+ )
+ old_cache = temporal_kv_cache or {}
+ new_cache: dict[str, tuple[torch.Tensor, torch.Tensor]] = {}
+
+ # Build coordinate pyramid. The base sits at the *top up-level* token
+ # grid (= resolution // out_patch_size), which is the highest-res grid
+ # the attention blocks actually see -- the final detokenizer's
+ # pixel-shuffle never gets coords.
+ c_base = F.avg_pool2d(coords, kernel_size=self.out_patch_size, stride=self.out_patch_size)
+ c_levels = [c_base]
+ for i in range(n_down):
+ c_levels.append(F.avg_pool2d(c_levels[-1], kernel_size=2, stride=2))
+
+ # Stem per-frame + spatial PE
+ z = einops.rearrange(z, "B C T H W -> (B T) C H W")
+ h = self.stem(z) + self.spatial_pe(c_levels[n_down]) # (BT, H_z, W_z, C)
+ h = einops.rearrange(h, "(B T) H W C -> B T H W C", B=B, T=T_cur)
+
+ # Bottleneck
+ _, _, Hc, Wc, _ = h.shape
+ if self.bottleneck_3d:
+ for block_idx, block in enumerate(self.mid_3d):
+ if streaming:
+ cache_key = f"mid_3d.{block_idx}"
+ h, new_cache[cache_key] = block.forward_stream(h, c_levels[n_down], old_cache.pop(cache_key, None))
+ else:
+ h = block(h, c_levels[n_down])
+ else:
+ # Factorized: temporal then global-spatial (decoder order).
+ for block_idx, (t_block, s_block) in enumerate(zip(self.mid_temporal, self.mid_spatial)):
+ h = einops.rearrange(h, "B T H W C -> (B H W) T C")
+ if streaming:
+ cache_key = f"mid_temporal.{block_idx}"
+ h, new_cache[cache_key] = t_block.forward_stream(h, old_cache.pop(cache_key, None))
+ else:
+ h = t_block(h)
+ h = einops.rearrange(h, "(B H W) T C -> (B T) H W C", B=B, H=Hc, W=Wc)
+ h = s_block(h, c_levels[n_down])
+ h = einops.rearrange(h, "(B T) H W C -> B T H W C", B=B, T=T_cur)
+
+ # Up levels
+ for i in reversed(range(n_down)):
+ # Temporal upsample
+ if self.temporal_upsample[i]:
+ h = self.up_levels[f"temporal_expand_{i}"](h)
+ T_cur = h.shape[1]
+
+ # Spatial upsample per-frame
+ h = einops.rearrange(h, "B T H W C -> (B T) H W C")
+ h = self.up_levels[f"expand_{i}"](h)
+ h = einops.rearrange(h, "(B T) H W C -> B T H W C", B=B, T=T_cur)
+
+ _, _, Hc, Wc, _ = h.shape
+
+ has_temporal = f"temporal_{i}" in self.up_levels
+ t_blocks = self.up_levels[f"temporal_{i}"] if has_temporal else [None] * len(self.up_levels[f"spatial_{i}"])
+ for block_idx, (t_block, s_block) in enumerate(zip(t_blocks, self.up_levels[f"spatial_{i}"])):
+ if t_block is not None:
+ h = einops.rearrange(h, "B T H W C -> (B H W) T C")
+ if streaming:
+ if not isinstance(t_block, TemporalBlock):
+ raise NotImplementedError(
+ "streaming decode currently supports temporal attention, not temporal convolution"
+ )
+ cache_key = f"up_levels.temporal_{i}.{block_idx}"
+ h, new_cache[cache_key] = t_block.forward_stream(h, old_cache.pop(cache_key, None))
+ else:
+ h = t_block(h)
+ h = einops.rearrange(h, "(B H W) T C -> (B T) H W C", B=B, H=Hc, W=Wc)
+ else:
+ h = einops.rearrange(h, "B T H W C -> (B T) H W C")
+ h = s_block(h, c_levels[i])
+ h = einops.rearrange(h, "(B T) H W C -> B T H W C", B=B, T=T_cur)
+
+ # Detokenize per-frame
+ h = einops.rearrange(h, "B T H W C -> (B T) H W C")
+ h = self.detokenizer(h) # (BT, C_out, H, W)
+ h = einops.rearrange(h, "(B T) C H W -> B C T H W", B=B, T=T_cur)
+ if streaming:
+ missing = set(old_cache) - set(new_cache)
+ if missing:
+ raise ValueError(f"unused temporal KV cache keys: {sorted(missing)}")
+ return h, new_cache
+ return h
+
+
+# =============================================================================
+# TransformerVAE
+# =============================================================================
+
+
+class TransformerVAE(nn.Module):
+ def __init__(
+ self,
+ resolution: list[int],
+ in_channels: int,
+ out_channels: int | None = None,
+ z_dim: int = 4,
+ base_channels: int = 128,
+ patch_size: list[int] = (1, 4),
+ window_size: list[int] = (3, 9),
+ depths: list[int] = (3, 3, 3, 3),
+ num_heads: list[int] = (2, 4, 8, 16),
+ dilation: list[int] = (1, 1, 1, 1),
+ temporal_downsample: list[bool] = (True, False, False),
+ temporal_upsample: list[bool] | None = None,
+ # ---- Asymmetric decoder (opt-in) ----
+ # By default the decoder mirrors the encoder (same depths / num_heads /
+ # dilation / temporal schedule, out_patch_size = patch_size). To build a
+ # *plain-ViT* encoder (large patch + few merges) paired with a heavier
+ # *progressive* decoder (small out_patch + more expand levels), set
+ # these. ``spatial_compression`` / ``temporal_compression`` are always
+ # derived from the ENCODER side. The only hard constraint is that the
+ # decoder's bottleneck spatial grid
+ # resolution // out_patch_size >> (len(decoder_depths) - 1)
+ # equals the encoder's latent grid
+ # resolution // patch_size >> (len(depths) - 1).
+ # When ``decoder_depths`` is set you MUST also pass
+ # ``decoder_temporal_upsample`` (length len(decoder_depths) - 1) since
+ # the encoder's temporal schedule no longer has the matching length.
+ decoder_depths: list[int] | None = None,
+ decoder_num_heads: list[int] | None = None,
+ decoder_dilation: list[int] | None = None,
+ decoder_temporal_upsample: list[bool] | None = None,
+ out_patch_size: list[int] | None = None,
+ mlp_ratio: float = 3.0,
+ dropout: float = 0.0,
+ mapping_depth: int = 2,
+ positional_embedding: Literal["learnable_embedding"] = "learnable_embedding",
+ formulation: str = ContinuousFormulation.AE.name,
+ # WAN-style first-frame-special temporal merging in the encoder.
+ # See :class:`TemporalMergingFirstFrameSpecial` for details.
+ # Pair with ``temporal_expand_wan_style=True`` so the decoder's
+ # ``TemporalExpandingWanStyle`` mirrors the encoder's chunk
+ # boundaries: ``T_in = 1 + (2 ** N) * (T_z - 1)`` exactly,
+ # no front-crop, no boundary information loss.
+ temporal_first_frame_special: bool = False,
+ # WAN-style first-frame-special temporal expanding in the decoder.
+ # When True every level with ``temporal_upsample[i]=True`` uses
+ # :class:`TemporalExpandingWanStyle` (T_in -> 2*T_in - 1) instead
+ # of plain :class:`TemporalExpanding` (T_in -> 2*T_in). Designed
+ # to be paired with ``temporal_first_frame_special=True``.
+ temporal_expand_wan_style: bool = False,
+ # ---- Temporal-mixing mechanism (ablation) ----
+ # "attention" (default): temporal mixing via
+ # :class:`CausalTemporalAttention` in the down/mid/up temporal blocks.
+ # "conv": WAN-style causal temporal convolution
+ # (:class:`CausalTemporalConv`, kernel ``temporal_conv_kernel``,
+ # spatial 1x1) instead. Spatial mixing (neighborhood attention) is
+ # unchanged, so this is a clean temporal-attn-vs-temporal-conv swap.
+ temporal_mixer: Literal["attention", "conv"] = "attention",
+ temporal_conv_kernel: int = 3,
+ # ---- Validity (raydrop) mask prediction (arxiv:2605.22809) ----
+ # When True, the decoder emits ONE extra output channel of validity
+ # logits (1 = valid return, 0 = no return). The data channels are
+ # unchanged; ``decode()`` returns only the data channels by default,
+ # so all existing call sites are unaffected. The validity logits can
+ # be retrieved at inference via ``decode(..., return_validity=True)``.
+ predict_validity: bool = False,
+ # ---- Validity mask as INPUT channel ----
+ # When True, the *caller* (e.g. the 16x tokenizer interface) is
+ # expected to concatenate a binary validity mask as an extra input
+ # channel before ``encode()``. Kept as attributes so interfaces can
+ # read thresholding constants (``mask_input_min_value`` /
+ # ``mask_input_eps``). ``in_channels`` MUST include this extra channel
+ # (e.g. mono depth -> in_channels=2).
+ mask_as_input: bool = False,
+ mask_input_min_value: float = -1.0,
+ mask_input_eps: float = 0.02,
+ # ---- Bottleneck 3D attention (encoder AND decoder) ----
+ # When True, both the encoder and decoder replace their factorized
+ # bottleneck (global spatial + temporal attention) with a single
+ # joint 3D self-attention over (T*H*W) tokens (:class:`Bottleneck3DBlock`).
+ # Default False -> both keep the factorized bottleneck, bit-identical
+ # to every existing config. ``bottleneck_3d_max_t`` caps the temporal
+ # PE / RoPE length (set >= largest latent T_z fed to the model);
+ # ``bottleneck_3d_causal_time`` masks the joint attention to be causal
+ # in time; ``bottleneck_3d_rope`` swaps the (learnable temporal PE +
+ # 2D AxialRoPE) for a parameter-free 3D RoPE over (t, h, w).
+ bottleneck_3d: bool = False,
+ bottleneck_3d_max_t: int = 32,
+ bottleneck_3d_causal_time: bool = False,
+ bottleneck_3d_rope: bool = False,
+ # If False, the local (neighborhood) attention does NOT wrap the W
+ # axis (clamps at the edges). Default True keeps the range-image
+ # azimuth-periodic behavior.
+ circular_padding: bool = True,
+ # Decoder-only fine-tuning: freeze ``encoder`` + ``quant_conv``.
+ freeze_encoder: bool = False,
+ # Recompute encoder/decoder activations during backward. Recommended
+ # for the production 128x3600, T=9 training shape.
+ gradient_checkpointing: bool = False,
+ ):
+ super().__init__()
+ self.resolution = _pair(resolution)
+ self.in_channels = in_channels
+ self.out_channels = out_channels if out_channels else in_channels
+ self.predict_validity = predict_validity
+ self.mask_as_input = mask_as_input
+ self.mask_input_min_value = float(mask_input_min_value)
+ self.mask_input_eps = float(mask_input_eps)
+ # Alias of ``mask_as_input`` retained for callers that inspect this flag.
+ self.uses_valid_mask = mask_as_input
+ self.freeze_encoder = bool(freeze_encoder)
+ self.gradient_checkpointing = bool(gradient_checkpointing)
+ # Decoder produces data channels (+1 validity logit channel if enabled).
+ self._decoder_out_channels = self.out_channels + (1 if predict_validity else 0)
+ self.z_dim = z_dim
+ self.patch_size = _pair(patch_size)
+
+ n_down = len(depths) - 1
+ self.spatial_compression = (
+ self.patch_size[0] * (2**n_down),
+ self.patch_size[1] * (2**n_down),
+ )
+ self.temporal_compression = 2 ** sum(temporal_downsample)
+
+ coords = encoding.generate_polar_coords(*self.resolution)
+ self.register_buffer("coords", coords)
+
+ self.encoder = Encoder(
+ resolution=resolution,
+ in_channels=in_channels,
+ z_dim=z_dim,
+ base_channels=base_channels,
+ patch_size=patch_size,
+ window_size=window_size,
+ depths=depths,
+ num_heads=num_heads,
+ dilation=dilation,
+ temporal_downsample=temporal_downsample,
+ mlp_ratio=mlp_ratio,
+ dropout=dropout,
+ mapping_depth=mapping_depth,
+ positional_embedding=positional_embedding,
+ temporal_first_frame_special=temporal_first_frame_special,
+ bottleneck_3d=bottleneck_3d,
+ bottleneck_3d_max_t=bottleneck_3d_max_t,
+ bottleneck_3d_causal_time=bottleneck_3d_causal_time,
+ bottleneck_3d_rope=bottleneck_3d_rope,
+ temporal_mixer=temporal_mixer,
+ temporal_conv_kernel=temporal_conv_kernel,
+ circular_padding=circular_padding,
+ )
+
+ self.distribution = ContinuousFormulation[formulation].value()
+
+ self.quant_conv = nn.Conv2d(z_dim * 2, z_dim * 2, 1)
+ self.post_quant_conv = nn.Conv2d(z_dim, z_dim, 1)
+
+ # Resolve the (possibly asymmetric) decoder topology. Defaults mirror
+ # the encoder so every existing config stays bit-identical.
+ _dec_depths = decoder_depths if decoder_depths is not None else depths
+ _dec_num_heads = decoder_num_heads if decoder_num_heads is not None else num_heads
+ _dec_dilation = decoder_dilation if decoder_dilation is not None else dilation
+ if decoder_depths is not None:
+ if decoder_temporal_upsample is None:
+ raise ValueError(
+ "decoder_depths is set (asymmetric decoder) so the encoder's "
+ "temporal schedule no longer has the matching length; pass "
+ "decoder_temporal_upsample with length "
+ f"len(decoder_depths) - 1 = {len(_dec_depths) - 1}."
+ )
+ _dec_temporal_upsample = decoder_temporal_upsample
+ _dec_temporal_downsample = decoder_temporal_upsample
+ else:
+ _dec_temporal_upsample = temporal_upsample
+ _dec_temporal_downsample = temporal_downsample
+
+ self.decoder = Decoder(
+ resolution=resolution,
+ out_channels=self._decoder_out_channels,
+ z_dim=z_dim,
+ base_channels=base_channels,
+ patch_size=patch_size,
+ out_patch_size=out_patch_size,
+ window_size=window_size,
+ depths=_dec_depths,
+ num_heads=_dec_num_heads,
+ dilation=_dec_dilation,
+ temporal_downsample=_dec_temporal_downsample,
+ temporal_upsample=_dec_temporal_upsample,
+ mlp_ratio=mlp_ratio,
+ dropout=dropout,
+ mapping_depth=mapping_depth,
+ positional_embedding=positional_embedding,
+ temporal_expand_wan_style=temporal_expand_wan_style,
+ bottleneck_3d=bottleneck_3d,
+ bottleneck_3d_max_t=bottleneck_3d_max_t,
+ bottleneck_3d_causal_time=bottleneck_3d_causal_time,
+ bottleneck_3d_rope=bottleneck_3d_rope,
+ temporal_mixer=temporal_mixer,
+ temporal_conv_kernel=temporal_conv_kernel,
+ circular_padding=circular_padding,
+ )
+ self.supports_temporal_kv_cache = not any(
+ _dec_temporal_upsample if _dec_temporal_upsample is not None else _dec_temporal_downsample
+ ) and (
+ (bottleneck_3d and bottleneck_3d_causal_time and bottleneck_3d_rope and temporal_mixer == "attention")
+ or (not bottleneck_3d and temporal_mixer == "attention")
+ )
+
+ if self.freeze_encoder:
+ for p in self.encoder.parameters():
+ p.requires_grad = False
+ for p in self.quant_conv.parameters():
+ p.requires_grad = False
+ self.encoder.eval()
+ self.quant_conv.eval()
+
+ def train(self, mode: bool = True):
+ """Keep frozen submodules in eval mode regardless of ``model.train()``."""
+ super().train(mode)
+ if self.freeze_encoder:
+ self.encoder.eval()
+ self.quant_conv.eval()
+ return self
+
+ def encode(self, x: torch.Tensor):
+ """
+ Args:
+ x: (B, C, H, W) [single-frame] or (B, C, T, H, W)
+ Returns:
+ sample: (B, z_dim, T', H_z, W_z)
+ posteriors: (mu, log_var) each (B, z_dim, T', H_z, W_z)
+ """
+ # Single-frame (image) input: add a T=1 axis so the video encoder
+ # path works unchanged. The causal temporal ops degrade gracefully
+ # at T=1 (TemporalMerging pads -> kernel-3 stride-2 -> T stays 1) and
+ # ``decode()`` squeezes the axis back. ``_input_T`` records T so the
+ # decode crop / squeeze stay consistent.
+ if x.ndim == 4:
+ x = x.unsqueeze(2) # (B, C, H, W) -> (B, C, 1, H, W)
+ self._input_T = x.shape[2]
+ coords = self.coords.to(x)
+ if self.gradient_checkpointing and self.training and torch.is_grad_enabled():
+ parameters = checkpoint(self.encoder, x, coords, use_reentrant=False)
+ else:
+ parameters = self.encoder(x, coords)
+
+ # quant_conv is 2D — apply per-frame
+ B, C2, Tz, Hz, Wz = parameters.shape
+ parameters = einops.rearrange(parameters, "B C T H W -> (B T) C H W")
+ parameters = self.quant_conv(parameters)
+ parameters = einops.rearrange(parameters, "(B T) C H W -> B C T H W", B=B, T=Tz)
+
+ sample, (mu, log_var) = self.distribution(parameters, split=True)
+
+ return sample, (mu, log_var)
+
+ def encode_stream(
+ self,
+ x_new: torch.Tensor,
+ temporal_kv_cache: dict[str, tuple[torch.Tensor, torch.Tensor]] | None = None,
+ ):
+ """Encode only newly appended frames using temporal / joint-3D KV caches."""
+ if not self.supports_temporal_kv_cache:
+ raise NotImplementedError("This TransformerVAE config does not support temporal KV-cache streaming")
+ if x_new.ndim != 5:
+ raise ValueError(f"encode_stream expects [B,C,T,H,W], got {tuple(x_new.shape)}")
+ B = x_new.shape[0]
+ T_new = x_new.shape[2]
+ parameters, new_cache = self.encoder.forward_stream(x_new, self.coords.to(x_new), temporal_kv_cache)
+ parameters = einops.rearrange(parameters, "B C T H W -> (B T) C H W")
+ parameters = self.quant_conv(parameters)
+ parameters = einops.rearrange(parameters, "(B T) C H W -> B C T H W", B=B, T=T_new)
+ sample, (mu, log_var) = self.distribution(parameters, split=True)
+ return sample, (mu, log_var), new_cache
+
+ def encode_streaming(
+ self,
+ x: torch.Tensor,
+ *,
+ chunk_frames: int = 9,
+ context_frames: int | None = None,
+ sample_posterior: bool = False,
+ ) -> torch.Tensor:
+ """WAN-style chunked encode for arbitrarily long videos.
+
+ Processes ``x`` in windows of ``chunk_frames`` with a causal KV cache so
+ peak activation stays O(chunk). For total ``T <= bottleneck_3d_max_t``,
+ results match a single :meth:`encode` call (up to SDPA numerics).
+ """
+ if x.ndim != 5:
+ raise ValueError(f"encode_streaming expects [B,C,T,H,W], got {tuple(x.shape)}")
+ if chunk_frames < 1:
+ raise ValueError(f"chunk_frames must be positive, got {chunk_frames}")
+ if context_frames is not None and context_frames < chunk_frames:
+ raise ValueError(f"context_frames must be at least chunk_frames, got {context_frames=} and {chunk_frames=}")
+ T = x.shape[2]
+ latents = []
+ cache = None
+ for start in range(0, T, chunk_frames):
+ chunk = x[:, :, start : start + chunk_frames]
+ cache = self._limit_streaming_cache(cache, context_frames, new_frames=chunk.shape[2])
+ sample, (mu, _), cache = self.encode_stream(chunk, cache)
+ latents.append(sample if sample_posterior else mu)
+ return torch.cat(latents, dim=2)
+
+ @staticmethod
+ def _limit_streaming_cache(
+ cache: dict[str, tuple[torch.Tensor, torch.Tensor]] | None,
+ context_frames: int | None,
+ *,
+ new_frames: int,
+ ) -> dict[str, tuple[torch.Tensor, torch.Tensor]] | None:
+ """Keep enough cached K/V frames for the requested total context."""
+ if cache is None or context_frames is None:
+ return cache
+ keep_past = context_frames - new_frames
+ if keep_past <= 0:
+ return None
+ return {
+ key: (key_cache[:, :, -keep_past:], value_cache[:, :, -keep_past:])
+ for key, (key_cache, value_cache) in cache.items()
+ }
+
+ def decode(
+ self,
+ z: torch.Tensor,
+ return_validity: bool = False,
+ ):
+ """
+ Args:
+ z: (B, z_dim, T', H_z, W_z)
+ Returns:
+ (B, C, T, H, W) — cropped to original T if temporal padding was used
+ When ``return_validity=True`` and ``predict_validity``, also returns
+ validity logits as a second value.
+ """
+ # post_quant_conv is 2D — apply per-frame
+ B, Cz, Tz, Hz, Wz = z.shape
+ z = einops.rearrange(z, "B C T H W -> (B T) C H W")
+ z = self.post_quant_conv(z)
+ z = einops.rearrange(z, "(B T) C H W -> B C T H W", B=B, T=Tz)
+
+ coords = self.coords.to(z)
+ if self.gradient_checkpointing and self.training and torch.is_grad_enabled():
+ out = checkpoint(self.decoder, z, coords, use_reentrant=False)
+ else:
+ out = self.decoder(z, coords)
+
+ # Crop decoder temporal padding back to the requested length.
+ # For 1x temporal compression, key off this latent's T so a stale
+ # encode-side ``_input_T`` cannot silently shorten a longer decode
+ # (e.g. generation after encoding a short clip on the same module).
+ if self.temporal_compression == 1:
+ target_T = Tz
+ else:
+ target_T = getattr(self, "_input_T", None)
+ if target_T is not None and out.shape[2] > target_T:
+ out = out[:, :, -target_T:]
+
+ # Single-frame input: drop the T=1 axis so the output matches the 4D
+ # (B, C, H, W) recon target produced by the image dataloader path.
+ if target_T == 1:
+ out = out.squeeze(2)
+
+ # Split off the validity-logit channel (last) when enabled. By default
+ # decode returns only the data channels so all existing call sites
+ # (inference, callbacks, DiT) are unaffected.
+ if self.predict_validity:
+ recon = out[:, : self.out_channels]
+ validity_logits = out[:, self.out_channels :]
+ if return_validity:
+ return recon, validity_logits
+ return recon
+ if return_validity:
+ return out, None
+ return out
+
+ def decode_stream(
+ self,
+ z_new: torch.Tensor,
+ temporal_kv_cache: dict[str, tuple[torch.Tensor, torch.Tensor]] | None = None,
+ return_validity: bool = False,
+ ):
+ """Decode only newly appended latent frames using temporal KV caches.
+
+ Exact for the 1x-temporal causal decoder (factorized temporal attn or
+ joint-3D with ``causal_time`` + 3D RoPE): spatial ops run only on
+ ``z_new`` while temporal / joint-3D layers attend to the cached prefix.
+ """
+ if not self.supports_temporal_kv_cache:
+ raise NotImplementedError("This TransformerVAE config does not support temporal KV-cache streaming")
+ B, _, T_new, _, _ = z_new.shape
+ z_new = einops.rearrange(z_new, "B C T H W -> (B T) C H W")
+ z_new = self.post_quant_conv(z_new)
+ z_new = einops.rearrange(z_new, "(B T) C H W -> B C T H W", B=B, T=T_new)
+ out, new_cache = self.decoder(
+ z_new,
+ self.coords.to(z_new),
+ temporal_kv_cache=temporal_kv_cache,
+ return_temporal_kv_cache=True,
+ )
+
+ if self.predict_validity:
+ recon = out[:, : self.out_channels]
+ validity_logits = out[:, self.out_channels :]
+ if return_validity:
+ return (recon, validity_logits), new_cache
+ return recon, new_cache
+ if return_validity:
+ return (out, None), new_cache
+ return out, new_cache
+
+ def decode_streaming(
+ self,
+ z: torch.Tensor,
+ *,
+ chunk_frames: int = 9,
+ context_frames: int | None = None,
+ return_validity: bool = False,
+ ):
+ """WAN-style chunked decode for arbitrarily long latent videos."""
+ if z.ndim != 5:
+ raise ValueError(f"decode_streaming expects [B,C,T,H,W], got {tuple(z.shape)}")
+ if chunk_frames < 1:
+ raise ValueError(f"chunk_frames must be positive, got {chunk_frames}")
+ if context_frames is not None and context_frames < chunk_frames:
+ raise ValueError(f"context_frames must be at least chunk_frames, got {context_frames=} and {chunk_frames=}")
+ T = z.shape[2]
+ recons = []
+ validities = []
+ cache = None
+ for start in range(0, T, chunk_frames):
+ chunk = z[:, :, start : start + chunk_frames]
+ cache = self._limit_streaming_cache(cache, context_frames, new_frames=chunk.shape[2])
+ if return_validity:
+ (recon, logits), cache = self.decode_stream(chunk, cache, return_validity=True)
+ recons.append(recon)
+ validities.append(logits)
+ else:
+ recon, cache = self.decode_stream(chunk, cache, return_validity=False)
+ recons.append(recon)
+ recon = torch.cat(recons, dim=2)
+ if return_validity:
+ validity = torch.cat(validities, dim=2) if validities[0] is not None else None
+ return recon, validity
+ return recon
+
+ def forward(
+ self,
+ x: torch.Tensor,
+ valid_mask: torch.Tensor | None = None,
+ ) -> dict[str, Any]:
+ """Training forward: encode → sample → decode (+ optional validity).
+
+ Args:
+ x: ``[B, C, T, H, W]`` tokenizer input. For the 3-channel path ``C=3``
+ (range, intensity, mask) with ``mask_as_input=False``. For the
+ 16x path ``C=1`` range with ``mask_as_input=True`` (mask is
+ concatenated here).
+ valid_mask: optional ``[B, 1, T, H, W]`` validity used when
+ ``mask_as_input=True`` instead of deriving the mask from range.
+
+ Returns:
+ Dict with ``reconstructions``, ``posteriors`` ``(mu, log_var)``,
+ ``latent``, and optionally ``validity_logits``.
+ """
+ if self.mask_as_input:
+ if valid_mask is not None:
+ valid = valid_mask.to(dtype=x.dtype)
+ else:
+ valid = (x[:, :1] > (self.mask_input_min_value + self.mask_input_eps)).to(dtype=x.dtype)
+ x = torch.cat([x, valid], dim=1)
+
+ if self.freeze_encoder:
+ with torch.no_grad():
+ sample, posteriors = self.encode(x)
+ sample = sample.detach()
+ posteriors = tuple(p.detach() for p in posteriors)
+ else:
+ sample, posteriors = self.encode(x)
+
+ decoded = self.decode(sample, return_validity=self.predict_validity)
+ if self.predict_validity:
+ recon, validity_logits = decoded
+ else:
+ recon, validity_logits = decoded, None
+
+ out: dict[str, Any] = {
+ "reconstructions": recon,
+ "posteriors": posteriors,
+ "latent": sample,
+ }
+ if validity_logits is not None:
+ out["validity_logits"] = validity_logits
+ return out
diff --git a/cosmos_framework/model/generator/tokenizers/lidar/normalization.py b/cosmos_framework/model/generator/tokenizers/lidar/normalization.py
new file mode 100644
index 000000000..b4f5c9fe2
--- /dev/null
+++ b/cosmos_framework/model/generator/tokenizers/lidar/normalization.py
@@ -0,0 +1,122 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+
+"""Shared LiDAR input and output normalization.
+
+Both V0 and V1 encode metric range, unit intensity, and a sensor mask through
+:func:`metric_lidar_to_network`. V0 then drops intensity before its 2-channel
+network; V1 keeps all three channels. Both come back out through
+:func:`network_lidar_to_metric_clip`, so a decoded clip carries the same channels
+in the same units whichever version produced it.
+
+This module composes the three-channel layout and decides validity; the per-channel
+scaling it composes lives in ``preprocessing`` so that the encode side, the decode
+side and the visualizers cannot drift onto separate copies of the same affine.
+"""
+
+from __future__ import annotations
+
+import torch
+
+from cosmos_framework.model.generator.tokenizers.lidar.postprocessing import (
+ DEFAULT_VALIDITY_THRESHOLD,
+ validity_mask,
+)
+from cosmos_framework.model.generator.tokenizers.lidar.preprocessing import (
+ INVALID_NORMALIZED_VALUE,
+ network_intensity_to_unit,
+ network_range_to_metric,
+ normalize_range_map,
+ unit_intensity_to_network,
+)
+
+
+def metric_lidar_to_network(
+ video: torch.Tensor,
+ *,
+ min_range: float | torch.Tensor,
+ max_range: float | torch.Tensor,
+) -> tuple[torch.Tensor, torch.Tensor]:
+ """Convert metric range, unit intensity, and sensor mask to network space."""
+ if video.ndim != 5 or video.shape[1] != 3:
+ raise ValueError(f"Expected metric LiDAR [B,3,T,H,W], got {tuple(video.shape)}")
+ range_m = video[:, :1]
+ intensity = video[:, 1:2]
+ sensor_valid = validity_mask(video[:, 2:3])
+ min_value = torch.as_tensor(min_range, dtype=range_m.dtype, device=range_m.device)
+ max_value = torch.as_tensor(max_range, dtype=range_m.dtype, device=range_m.device)
+ if bool((max_value <= min_value).item()):
+ raise ValueError(f"max_range must exceed min_range, got {min_value.item()} and {max_value.item()}")
+ valid = sensor_valid & (range_m >= min_value) & (range_m <= max_value)
+ normalized_range = normalize_range_map(range_m, min_range=min_value, max_range=max_value)
+ normalized_intensity = unit_intensity_to_network(intensity)
+ invalid_fill = torch.full_like(normalized_range, INVALID_NORMALIZED_VALUE)
+ normalized_range = torch.where(valid, normalized_range, invalid_fill)
+ normalized_intensity = torch.where(valid, normalized_intensity, invalid_fill)
+ valid_float = valid.to(dtype=video.dtype)
+ return torch.cat((normalized_range, normalized_intensity, valid_float), dim=1), valid_float
+
+
+def network_lidar_to_metric(
+ reconstruction: torch.Tensor,
+ validity: torch.Tensor,
+ *,
+ min_range: float | torch.Tensor,
+ max_range: float | torch.Tensor,
+ apply_validity_mask: bool = True,
+ validity_threshold: float = DEFAULT_VALIDITY_THRESHOLD,
+) -> torch.Tensor:
+ """Convert normalized range/intensity reconstructions to metric public output.
+
+ ``validity_threshold`` is the probability above which a predicted ray is
+ kept. Raising it trades recall for precision: a sparser sweep holding only
+ the returns the decoder is confident about.
+ """
+ if reconstruction.ndim != 5 or reconstruction.shape[1] < 2:
+ raise ValueError(f"Expected reconstruction [B,2+,T,H,W], got {tuple(reconstruction.shape)}")
+ min_value = torch.as_tensor(min_range, dtype=reconstruction.dtype, device=reconstruction.device)
+ max_value = torch.as_tensor(max_range, dtype=reconstruction.dtype, device=reconstruction.device)
+ range_m = network_range_to_metric(reconstruction[:, :1], min_range=min_value, max_range=max_value)
+ intensity = network_intensity_to_unit(reconstruction[:, 1:2])
+ if apply_validity_mask:
+ valid = validity_mask(validity, threshold=validity_threshold)
+ range_m = torch.where(valid, range_m, torch.zeros_like(range_m))
+ intensity = torch.where(valid, intensity, torch.zeros_like(intensity))
+ return torch.cat((range_m, intensity), dim=1)
+
+
+def network_lidar_to_metric_clip(
+ normalized: torch.Tensor,
+ validity: torch.Tensor,
+ *,
+ min_range: float | torch.Tensor,
+ max_range: float | torch.Tensor,
+ apply_validity_mask: bool = True,
+ validity_threshold: float = DEFAULT_VALIDITY_THRESHOLD,
+) -> torch.Tensor:
+ """The public decode payload, shared by V0 and V1: ``[B,3,T,H,W]`` in metric units.
+
+ Channel 0 is range in metres, channel 1 unit intensity, channel 2 the resolved
+ ``{0, 1}`` mask -- the same layout, in the same units, that the dataloader hands
+ to encode. Both tokenizers return this, so a consumer never has to ask which
+ version produced a clip in order to know what its numbers mean.
+
+ V0's frozen network predicts no intensity. It passes the absent fill through
+ channel 1 rather than special-casing the layout, and the inverse below maps
+ that fill to zero intensity.
+ """
+ metric = network_lidar_to_metric(
+ normalized,
+ validity,
+ min_range=min_range,
+ max_range=max_range,
+ apply_validity_mask=False,
+ validity_threshold=validity_threshold,
+ ) # [B,2,T,H,W]
+ if apply_validity_mask:
+ valid = validity_mask(validity, threshold=validity_threshold)
+ metric = torch.where(valid, metric, torch.zeros_like(metric))
+ channel = valid.to(dtype=validity.dtype)
+ else:
+ channel = validity
+ return torch.cat((metric, channel), dim=1) # [B,3,T,H,W]
diff --git a/cosmos_framework/model/generator/tokenizers/lidar/normalization_test.py b/cosmos_framework/model/generator/tokenizers/lidar/normalization_test.py
new file mode 100644
index 000000000..66d534460
--- /dev/null
+++ b/cosmos_framework/model/generator/tokenizers/lidar/normalization_test.py
@@ -0,0 +1,172 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+
+import pytest
+import torch
+
+import cosmos_framework.model.generator.tokenizers.lidar.normalization as normalization
+from cosmos_framework.model.generator.tokenizers.lidar.checkpoint import parse_lidar_checkpoint_stats, parse_lidar_stats
+from cosmos_framework.model.generator.tokenizers.lidar.normalization import (
+ metric_lidar_to_network,
+ network_lidar_to_metric,
+ network_lidar_to_metric_clip,
+)
+from cosmos_framework.model.generator.tokenizers.lidar.preprocessing import (
+ network_intensity_to_unit,
+ network_range_to_metric,
+ normalize_range_map,
+ unit_intensity_to_network,
+ unnormalize_range_map,
+)
+
+pytestmark = [pytest.mark.L0, pytest.mark.CPU]
+
+
+def test_channel_scaling_round_trips_through_the_shared_primitives() -> None:
+ """The forward and inverse affines invert each other inside the clamped domain."""
+ metric = torch.tensor([[[5.0, 42.5, 100.0]]])
+ unit = torch.tensor([[[0.0, 0.25, 1.0]]])
+
+ normalized_range = normalize_range_map(metric, min_range=5.0, max_range=100.0)
+ normalized_intensity = unit_intensity_to_network(unit)
+
+ torch.testing.assert_close(normalized_range.flatten(), torch.tensor([-1.0, -0.21052632, 1.0]))
+ torch.testing.assert_close(normalized_intensity.flatten(), torch.tensor([-1.0, -0.5, 1.0]))
+ torch.testing.assert_close(
+ network_range_to_metric(normalized_range, min_range=5.0, max_range=100.0),
+ metric,
+ )
+ torch.testing.assert_close(network_intensity_to_unit(normalized_intensity), unit)
+
+
+def test_composed_normalization_matches_the_primitives_it_delegates_to() -> None:
+ """``metric_lidar_to_network`` and its inverse must not re-derive the scaling.
+
+ Pinning the composed three-channel path against the per-channel primitives is what
+ keeps the encode side, the decode side and the visualizers on one affine.
+ """
+ video = torch.tensor([[[[[7.5, 60.0]]], [[[0.25, 0.75]]], [[[1.0, 1.0]]]]])
+
+ normalized, valid = metric_lidar_to_network(video, min_range=5.0, max_range=100.0)
+
+ torch.testing.assert_close(normalized[:, :1], normalize_range_map(video[:, :1], min_range=5.0, max_range=100.0))
+ torch.testing.assert_close(normalized[:, 1:2], unit_intensity_to_network(video[:, 1:2]))
+
+ metric = network_lidar_to_metric(normalized, valid, min_range=5.0, max_range=100.0)
+
+ torch.testing.assert_close(
+ metric[:, :1], network_range_to_metric(normalized[:, :1], min_range=5.0, max_range=100.0)
+ )
+ torch.testing.assert_close(metric[:, 1:2], network_intensity_to_unit(normalized[:, 1:2]))
+
+
+def test_prediction_overshoot_is_clamped_for_output_but_masked_for_display() -> None:
+ """The two inverse-range callers disagree on overshoot, and that is deliberate.
+
+ A decoded public sweep clamps back to the range envelope, while the display
+ helper leaves the value alone and lets the derived mask drop it, so neither
+ path silently adopts the other's handling.
+ """
+ overshoot = torch.tensor([[[1.5]]])
+
+ torch.testing.assert_close(
+ network_range_to_metric(overshoot, min_range=5.0, max_range=100.0),
+ torch.tensor([[[100.0]]]),
+ )
+ unclamped, mask = unnormalize_range_map(overshoot, min_range=5.0, max_range=100.0)
+ torch.testing.assert_close(unclamped, torch.tensor([[[0.0]]]))
+ assert not bool(mask.any())
+
+
+def test_metric_lidar_round_trip_and_range_validity() -> None:
+ video = torch.tensor([[[[[0.5, 1.0, 50.5, 101.0]]], [[[0.2, 0.3, 0.4, 0.5]]], [[[1.0, 1.0, 1.0, 1.0]]]]])
+
+ normalized, valid = metric_lidar_to_network(video, min_range=1.0, max_range=100.0)
+ metric = network_lidar_to_metric(
+ normalized,
+ valid,
+ min_range=1.0,
+ max_range=100.0,
+ apply_validity_mask=True,
+ )
+
+ torch.testing.assert_close(valid.flatten(), torch.tensor([0.0, 1.0, 1.0, 0.0]))
+ torch.testing.assert_close(metric[0, 0].flatten(), torch.tensor([0.0, 1.0, 50.5, 0.0]))
+ torch.testing.assert_close(metric[0, 1].flatten(), torch.tensor([0.0, 0.3, 0.4, 0.0]))
+
+
+def test_network_lidar_to_metric_honors_a_raised_validity_threshold() -> None:
+ """Raising the cut trades recall for precision on the predicted rays."""
+ reconstruction = torch.zeros((1, 3, 1, 1, 3))
+ validity = torch.tensor([[[[[0.4, 0.7, 0.99]]]]])
+
+ default_cut = network_lidar_to_metric(reconstruction, validity, min_range=1.0, max_range=101.0)
+ strict_cut = network_lidar_to_metric(
+ reconstruction,
+ validity,
+ min_range=1.0,
+ max_range=101.0,
+ validity_threshold=0.9,
+ )
+
+ torch.testing.assert_close(default_cut[0, 0].flatten(), torch.tensor([0.0, 51.0, 51.0]))
+ torch.testing.assert_close(strict_cut[0, 0].flatten(), torch.tensor([0.0, 0.0, 51.0]))
+
+
+def test_metric_clip_resolves_validity_once(monkeypatch: pytest.MonkeyPatch) -> None:
+ calls = 0
+ original = normalization.validity_mask
+
+ def counting_validity_mask(validity: torch.Tensor, *, threshold: float) -> torch.Tensor:
+ nonlocal calls
+ calls += 1
+ return original(validity, threshold=threshold)
+
+ monkeypatch.setattr(normalization, "validity_mask", counting_validity_mask)
+ reconstruction = torch.zeros((1, 2, 1, 1, 2))
+ validity = torch.tensor([[[[[0.4, 0.8]]]]])
+
+ decoded = network_lidar_to_metric_clip(reconstruction, validity, min_range=1.0, max_range=101.0)
+
+ assert calls == 1
+ torch.testing.assert_close(decoded[0, 2].flatten(), torch.tensor([0.0, 1.0]))
+
+
+def test_lidar_stats_preserve_bounds_and_accept_legacy_payloads() -> None:
+ mean, std, min_range, max_range = parse_lidar_stats(
+ {"mean": torch.zeros(2), "std": torch.ones(2), "min_range": 1.0, "max_range": 105.0}
+ )
+ assert mean.shape == std.shape == (2,)
+ assert (min_range, max_range) == (1.0, 105.0)
+
+ _, _, legacy_min, legacy_max = parse_lidar_stats((torch.zeros(2), torch.ones(2)))
+ assert legacy_min is legacy_max is None
+
+
+def test_lidar_stats_can_be_read_from_training_checkpoint() -> None:
+ mean, std, min_range, max_range = parse_lidar_checkpoint_stats(
+ {
+ "model": {
+ "latent_mean": torch.tensor([1.0, 2.0]),
+ "latent_std": torch.tensor([3.0, 4.0]),
+ "min_range": torch.tensor(1.0),
+ "max_range": torch.tensor(105.0),
+ }
+ }
+ )
+
+ torch.testing.assert_close(mean, torch.tensor([1.0, 2.0]))
+ torch.testing.assert_close(std, torch.tensor([3.0, 4.0]))
+ assert (min_range, max_range) == (1.0, 105.0)
+
+ _, _, min_range, max_range = parse_lidar_checkpoint_stats(
+ {
+ "model": {
+ "latent_mean": torch.tensor([1.0, 2.0]),
+ "latent_std": torch.tensor([3.0, 4.0]),
+ }
+ }
+ )
+ assert min_range is max_range is None
+
+
diff --git a/cosmos_framework/model/generator/tokenizers/lidar/paths.py b/cosmos_framework/model/generator/tokenizers/lidar/paths.py
new file mode 100644
index 000000000..b665f8c69
--- /dev/null
+++ b/cosmos_framework/model/generator/tokenizers/lidar/paths.py
@@ -0,0 +1,60 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+
+"""Canonical S3 / local asset paths for the LiDAR tokenizer."""
+
+from __future__ import annotations
+
+import os
+from pathlib import Path
+from typing import Any
+
+# =============================================================================
+# Published asset locations
+# =============================================================================
+
+
+DEFAULT_CREDENTIAL_PATH = "credentials/gcp_checkpoint.secret"
+
+# imaginaire4 repo root: .../tokenizers/lidar/paths.py -> parents[5]
+_REPO_ROOT = Path(__file__).resolve().parents[5]
+
+
+# =============================================================================
+# Path / credential helpers
+# =============================================================================
+
+
+def is_remote_uri(path: str | Path) -> bool:
+ """Return True when ``path`` is an object-store URI (e.g. ``s3://...``)."""
+ return "://" in str(path)
+
+
+def resolve_credential_path(credential_path: str | None = None) -> str | None:
+ """Resolve the object-store credential file, honoring env overrides.
+
+ Lookup order:
+ 1. ``LIDAR_TOKENIZER_CREDENTIAL_PATH``
+ 2. explicit ``credential_path`` argument
+ 3. ``DEFAULT_CREDENTIAL_PATH`` (repo-relative or cwd-relative)
+ """
+ path = os.environ.get("LIDAR_TOKENIZER_CREDENTIAL_PATH") or credential_path or DEFAULT_CREDENTIAL_PATH
+ resolved = Path(path).expanduser()
+ if resolved.is_file():
+ return str(resolved)
+
+ repo_relative = _REPO_ROOT / resolved
+ return str(repo_relative) if repo_relative.is_file() else None
+
+
+def s3_backend_args(credential_path: str | None = None) -> dict[str, Any] | None:
+ """Build ``easy_io`` backend args for GCS/S3 LiDAR assets."""
+ resolved = resolve_credential_path(credential_path)
+ if resolved is None:
+ return None
+
+ return {
+ "backend": "s3",
+ "path_mapping": None,
+ "s3_credential_path": resolved,
+ }
diff --git a/cosmos_framework/model/generator/tokenizers/lidar/postprocessing.py b/cosmos_framework/model/generator/tokenizers/lidar/postprocessing.py
new file mode 100644
index 000000000..599f297fe
--- /dev/null
+++ b/cosmos_framework/model/generator/tokenizers/lidar/postprocessing.py
@@ -0,0 +1,440 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+
+"""Decode-side LiDAR postprocessing shared by the V0 and V1 tokenizers.
+
+This module owns the validity cut: sigmoid of mask logits, then a probability
+threshold. Tokenizers apply it in decode; tokenizer-training reconstructions
+are still raw logits, so visualization asks this module for a keep-mask rather
+than reimplementing the cut. Callbacks and renderers do not sigmoid or threshold.
+
+A resolved clip -- dataloader ``{0, 1}``, or decode after the cut -- is read
+with the same half-cut and is never sigmoided again. Range-map preparation on
+the encode side belongs to ``preprocessing``.
+
+This module also owns the range smoothing a generated clip is run through before
+it is unprojected to XYZ, and the single switch that turns it on; see
+:class:`RangeSmoothing` and ``DEFAULT_RANGE_SMOOTHING``.
+"""
+
+from __future__ import annotations
+
+import attrs
+import torch
+
+from cosmos_framework.model.generator.tokenizers.lidar.preprocessing import (
+ INVALID_NORMALIZED_VALUE,
+ MAX_RANGE_METERS,
+ MIN_RANGE_METERS,
+)
+
+DEFAULT_VALIDITY_THRESHOLD = 0.5
+# The sentinel a range-only clip drops rays at is the fill encode writes, so read it
+# from there: a sentinel detector that disagreed with the fill would keep dropped rays.
+DEFAULT_MASK_INPUT_MIN_VALUE = INVALID_NORMALIZED_VALUE
+DEFAULT_MASK_INPUT_EPS = 0.02
+
+
+@attrs.frozen
+class RangeSmoothing:
+ """Validity-aware bilateral filter for a metric range map, run before unprojection.
+
+ A generated range map carries per-ray error that is invisible in the range
+ view -- neighbouring rays differ by a few tenths of a metre either way -- but
+ unprojection turns each ray into a point at its own distance, so that error
+ becomes scatter along the ray direction and reads as a thick, noisy surface
+ in bird's-eye view. Filtering the range before the unprojection is what
+ removes it; nothing downstream of XYZ can tell scatter from structure.
+
+ Each ray is replaced by a weighted mean of its window, where a neighbour's
+ weight is a spatial Gaussian in pixels times a range Gaussian in metres. The
+ range term is what keeps a depth discontinuity intact: rays across the gap
+ fall many sigma away and contribute nothing, so a foreground edge is
+ averaged along itself rather than smeared into the background behind it.
+ Unlike a median, the estimate is continuous in the samples, so it removes
+ sub-bin jitter instead of quantizing every ray onto whichever neighbour
+ happens to sit in the middle.
+
+ That window is one row wide by default, and the two axes are not
+ interchangeable. Along a row the samples are one laser sweeping azimuth, so
+ a smooth surface is nearly iso-range and the neighbours estimate it. Down a
+ column each sample is a different beam at a different elevation, and on the
+ ground plane -- most of a sweep -- that is a steep true range gradient, so
+ those neighbours disagree because the surface really is at another distance.
+ Averaging them in trades noise for distortion. Measured over 15 MADS sweeps,
+ filtering ground truth with the 5x5 window this shipped with moves a
+ surviving ray 0.098 m, against 0.067 m for one row; a 5x1 column-only
+ control accounts for 0.079 m of that, which is where it comes from. On truth
+ plus 0.2 m of per-ray jitter the row leaves 0.125 m against the square's
+ 0.133 m, and on 0.2 m of spatially correlated error -- the shape a decoder's
+ error actually has -- the square leaves 0.165 m where not filtering at all
+ leaves 0.160 m, while the row leaves 0.158 m. Set ``elevation_kernel_size``
+ above 1 to average beams together again.
+
+ The range bandwidth is ``range_sigma_m + range_sigma_ratio * range``,
+ because a fixed bandwidth is the wrong shape for this data: two rays a
+ degree apart that land on the same sloped ground are centimetres apart at
+ 10 m and metres apart at 80 m, so a bandwidth tight enough to clean the near
+ field would treat the whole far field as unrelated surfaces.
+
+ Weighting alone cannot remove a bad ray -- a flying pixel between a
+ foreground edge and the background agrees with itself perfectly and would
+ keep its own value. So a ray is dropped unless the weight its neighbours
+ contribute, ``min_neighbor_support``, says a local surface corroborates it.
+ That one rule covers both the flying pixel, whose neighbours are all too far
+ in range to weigh anything, and isolated speckle, which has no valid
+ neighbours to begin with. A row corroborates nearly as well as a square did:
+ over the same sweeps with 1% of rays displaced, one row catches 97.6% of
+ them against the square's 97.9%. The default asks for 19% of what a full
+ window could contribute, which is the share the 5x5 default asked for, so
+ the bar tracks the window rather than the shape.
+
+ ``strength`` interpolates between the measured range (0.0) and the filtered
+ estimate (1.0), so a run can keep some of the original relief.
+ """
+
+ azimuth_kernel_size: int = 5
+ elevation_kernel_size: int = 1
+ spatial_sigma: float = 1.5
+ range_sigma_m: float = 0.2
+ range_sigma_ratio: float = 0.03
+ min_neighbor_support: float = 0.45
+ strength: float = 1.0
+
+ @property
+ def window_size(self) -> int:
+ """Rays in one window, the ray itself included."""
+ return self.azimuth_kernel_size * self.elevation_kernel_size
+
+ def __attrs_post_init__(self) -> None:
+ if self.azimuth_kernel_size < 3 or self.azimuth_kernel_size % 2 == 0:
+ raise ValueError(f"azimuth_kernel_size must be an odd size of at least 3, got {self.azimuth_kernel_size}")
+ # One is the per-row default rather than a degenerate case; anything
+ # wider still has to be centred on the ray it filters.
+ if self.elevation_kernel_size < 1 or self.elevation_kernel_size % 2 == 0:
+ raise ValueError(f"elevation_kernel_size must be odd and at least 1, got {self.elevation_kernel_size}")
+ if self.spatial_sigma <= 0.0:
+ raise ValueError(f"spatial_sigma must be positive, got {self.spatial_sigma}")
+ if self.range_sigma_m <= 0.0:
+ raise ValueError(f"range_sigma_m must be positive, got {self.range_sigma_m}")
+ if self.range_sigma_ratio < 0.0:
+ raise ValueError(f"range_sigma_ratio must be non-negative, got {self.range_sigma_ratio}")
+ if not 0.0 <= self.min_neighbor_support <= self.window_size - 1:
+ raise ValueError(
+ f"min_neighbor_support must lie in [0, {self.window_size - 1}], got {self.min_neighbor_support}"
+ )
+ if not 0.0 <= self.strength <= 1.0:
+ raise ValueError(f"strength must lie in [0, 1], got {self.strength}")
+
+
+# The one switch for smoothing: the saved generated range map and every BEV unprojection
+# read this constant, so setting it to None goes back to raw generations everywhere. The
+# saved artifact is filtered because its consumers unproject it, which is where the error
+# shows: a generated sweep's per-ray range error is too small to see in the range view,
+# which stays raw so a reconstruction can be judged on what the tokenizer produced.
+# Measured over 48 sweeps of a generated V0 clip, the 5x5 window this shipped with took
+# the disagreement between a ray and its neighbours from 0.73 m to 0.26 m while moving a
+# surviving ray 0.15 m, and dropped the 4.6% of rays no local surface corroborated. Those
+# figures describe the square, not this default; what replaced it is measured against
+# ground truth in :class:`RangeSmoothing`, which is the comparison that showed the square
+# was spending more of that movement on distorting real structure than on noise.
+DEFAULT_RANGE_SMOOTHING: RangeSmoothing | None = RangeSmoothing()
+
+
+def _pad_range_window(value: torch.Tensor, *, pad_elevation: int, pad_azimuth: int) -> torch.Tensor:
+ """Pad ``[N,1,H,W]`` for a neighbourhood read: azimuth wraps, elevation does not."""
+ padded = torch.nn.functional.pad(value, (pad_azimuth, pad_azimuth, 0, 0), mode="circular")
+ if pad_elevation == 0:
+ return padded
+ return torch.nn.functional.pad(padded, (0, 0, pad_elevation, pad_elevation), mode="replicate")
+
+
+def _spatial_weights(
+ elevation: int, azimuth: int, sigma: float, *, dtype: torch.dtype, device: torch.device
+) -> torch.Tensor:
+ """Gaussian falloff over an ``elevation x azimuth`` window, flattened to one axis."""
+ rows = torch.arange(elevation, dtype=dtype, device=device) - elevation // 2
+ columns = torch.arange(azimuth, dtype=dtype, device=device) - azimuth // 2
+ squared = rows.reshape(-1, 1) ** 2 + columns.reshape(1, -1) ** 2
+ return torch.exp(-squared / (2.0 * sigma**2)).reshape(-1)
+
+
+def _smooth_range_sweeps(
+ ranges: torch.Tensor,
+ keep: torch.Tensor,
+ *,
+ smoothing: RangeSmoothing,
+) -> tuple[torch.Tensor, torch.Tensor]:
+ """Filter a ``[N,1,H,W]`` block of sweeps; see :func:`smooth_metric_range`."""
+ elevation = smoothing.elevation_kernel_size
+ azimuth = smoothing.azimuth_kernel_size
+ window = smoothing.window_size
+
+ def neighbourhood(value: torch.Tensor) -> torch.Tensor:
+ """Read each ray's window as a trailing axis."""
+ return (
+ _pad_range_window(value, pad_elevation=elevation // 2, pad_azimuth=azimuth // 2)
+ .unfold(2, elevation, 1)
+ .unfold(3, azimuth, 1)
+ .reshape(*ranges.shape, window)
+ )
+
+ # A dropped ray carries no distance, so it is zeroed here and then zero-weighted
+ # below; the same mechanism excludes a pad column that wrapped in over the seam.
+ windows = neighbourhood(torch.where(keep, ranges, torch.zeros_like(ranges)))
+ window_keep = neighbourhood(keep.to(ranges.dtype))
+
+ sigma = (smoothing.range_sigma_m + smoothing.range_sigma_ratio * ranges.abs()).unsqueeze(-1)
+ similarity = torch.exp(-0.5 * ((windows - ranges.unsqueeze(-1)) / sigma) ** 2)
+ spatial = _spatial_weights(elevation, azimuth, smoothing.spatial_sigma, dtype=ranges.dtype, device=ranges.device)
+ weights = window_keep * spatial * similarity
+
+ total = weights.sum(dim=-1)
+ filtered = (weights * windows).sum(dim=-1) / total.clamp_min(torch.finfo(ranges.dtype).tiny)
+
+ # The ray's own vote is excluded from the support: a flying pixel agrees with
+ # itself perfectly, and it is the neighbours that have to corroborate it.
+ support = total - weights[..., window // 2]
+ kept = keep & (support >= smoothing.min_neighbor_support)
+
+ smoothed = torch.lerp(ranges, filtered, smoothing.strength)
+ return torch.where(kept, smoothed, ranges), kept
+
+
+# A window holds one float per ray in it and the filter keeps a few tensors that
+# shape at once, so a 400-sweep clip at full azimuth would allocate tens of gigabytes
+# in one shot -- on the decode device, where that is fatal rather than merely slow.
+# Sweeps are filtered independently, so slicing them into blocks of roughly this many
+# window elements bounds the peak without changing the result.
+_WINDOW_ELEMENT_BUDGET = 16_000_000
+
+
+def smooth_metric_range(
+ metric_range: torch.Tensor,
+ valid: torch.Tensor,
+ *,
+ smoothing: RangeSmoothing = RangeSmoothing(),
+) -> tuple[torch.Tensor, torch.Tensor]:
+ """Bilateral-filter a metric range map ``[...,H,W]``, returning ranges and validity.
+
+ Only kept rays carry weight, and a dropped ray keeps its input range so no
+ distance is invented behind the mask. Azimuth wraps,
+ because the canvas spans a full turn with ``azimuth_endpoint=False`` and its
+ two edge columns are therefore neighbours on the sensor; elevation does not,
+ so a window that spans beams edge-pads the rows instead. The default window
+ is one row and reads no other beam at all.
+ """
+ if metric_range.shape != valid.shape:
+ raise ValueError(f"Expected matching range and validity shapes, got {metric_range.shape} and {valid.shape}")
+ if metric_range.ndim < 2:
+ raise ValueError(f"Expected a range map with [...,H,W] axes, got {tuple(metric_range.shape)}")
+
+ height, width = metric_range.shape[-2:]
+ if width < smoothing.azimuth_kernel_size:
+ raise ValueError(f"Azimuth width {width} is narrower than the {smoothing.azimuth_kernel_size}-column window")
+
+ leading = metric_range.shape[:-2]
+ ranges = metric_range.reshape(-1, 1, height, width).float()
+ keep = valid.reshape(-1, 1, height, width)
+
+ block = max(1, _WINDOW_ELEMENT_BUDGET // (height * width * smoothing.window_size))
+ smoothed_blocks: list[torch.Tensor] = []
+ kept_blocks: list[torch.Tensor] = []
+ for start in range(0, ranges.shape[0], block):
+ smoothed_block, kept_block = _smooth_range_sweeps(
+ ranges[start : start + block],
+ keep[start : start + block],
+ smoothing=smoothing,
+ )
+ smoothed_blocks.append(smoothed_block)
+ kept_blocks.append(kept_block)
+
+ smoothed = torch.cat(smoothed_blocks) if len(smoothed_blocks) > 1 else smoothed_blocks[0]
+ kept = torch.cat(kept_blocks) if len(kept_blocks) > 1 else kept_blocks[0]
+ return (
+ smoothed.reshape(*leading, height, width).to(dtype=metric_range.dtype),
+ kept.reshape(*leading, height, width),
+ )
+
+
+def validate_validity_threshold(threshold: float) -> float:
+ """Check that a validity probability cut lies strictly inside ``(0, 1)``.
+
+ Zero would keep every ray and one would drop every ray, so both ends erase
+ the prediction the cut is meant to read.
+ """
+ value = float(threshold)
+ if not 0.0 < value < 1.0:
+ raise ValueError(f"validity_threshold must lie in (0, 1), got {threshold}")
+ return value
+
+
+def validity_probability(mask_logits: torch.Tensor) -> torch.Tensor:
+ """Convert mask logits to a per-ray keep probability."""
+ return torch.sigmoid(mask_logits)
+
+
+def validity_mask(
+ probability: torch.Tensor,
+ *,
+ threshold: float = DEFAULT_VALIDITY_THRESHOLD,
+) -> torch.Tensor:
+ """Boolean keep-mask from a probability (or already-resolved ``{0, 1}``) channel."""
+ return probability >= validate_validity_threshold(threshold)
+
+
+def resolved_validity_channel(
+ validity: torch.Tensor,
+ *,
+ should_mask: bool,
+ threshold: float = DEFAULT_VALIDITY_THRESHOLD,
+) -> torch.Tensor:
+ """Channel 2 of a decoded clip, never raw logits.
+
+ A dataloader clip carries a ``{0, 1}`` sensor mask in this slot, so a
+ decoded clip resolves the cut here as well and every consumer reads the
+ two the same way. Leaving logits here is what previously pushed the
+ sigmoid out into the display path, where it drifted out of step with the
+ cut applied to range and intensity.
+
+ With ``should_mask=False`` there is no cut to resolve, so the probability
+ passes through and a reader's own ``> 0.5`` reproduces the default.
+ """
+ if not should_mask:
+ return validity
+ return validity_mask(validity, threshold=threshold).to(dtype=validity.dtype)
+
+
+def validity_from_mask_logits(
+ mask_logits: torch.Tensor,
+ *,
+ threshold: float = DEFAULT_VALIDITY_THRESHOLD,
+) -> torch.Tensor:
+ """Sigmoid then cut. The only path that may see raw mask logits."""
+ return validity_mask(validity_probability(mask_logits), threshold=threshold)
+
+
+def validity_from_range_sentinel(
+ range_channel: torch.Tensor,
+ *,
+ mask_input_min_value: float = DEFAULT_MASK_INPUT_MIN_VALUE,
+ mask_input_eps: float = DEFAULT_MASK_INPUT_EPS,
+) -> torch.Tensor:
+ """Keep rays whose normalized range sits strictly above the invalid fill."""
+ finite = torch.isfinite(range_channel)
+ return finite & (range_channel > (float(mask_input_min_value) + float(mask_input_eps)))
+
+
+def validity_from_video(
+ video: torch.Tensor,
+ *,
+ mask_input_min_value: float = DEFAULT_MASK_INPUT_MIN_VALUE,
+ mask_input_eps: float = DEFAULT_MASK_INPUT_EPS,
+) -> torch.Tensor:
+ """Boolean keep-mask for a dataloader or encoded clip ``[B,C,T,H,W]``.
+
+ Three-channel clips carry an already-resolved ``{0, 1}`` mask; one-channel
+ V0 clips drop rays at the range sentinel.
+ """
+ if video.shape[1] >= 3:
+ return validity_mask(video[:, 2:3])
+ return validity_from_range_sentinel(
+ video[:, :1],
+ mask_input_min_value=mask_input_min_value,
+ mask_input_eps=mask_input_eps,
+ )
+
+
+def validity_from_reconstruction(
+ reconstruction: torch.Tensor,
+ validity: torch.Tensor | None = None,
+ *,
+ mask_input_min_value: float = DEFAULT_MASK_INPUT_MIN_VALUE,
+ mask_input_eps: float = DEFAULT_MASK_INPUT_EPS,
+) -> torch.Tensor:
+ """Boolean keep-mask for a tokenizer training reconstruction ``[B,C,T,H,W]``.
+
+ ``reconstruction`` must be the raw network output, whose channel 2 is a mask
+ logit, because that channel goes through :func:`validity_from_mask_logits`.
+ A clip that came out of ``decode()`` is not such an input: channel 2 there is
+ already resolved to ``{0, 1}`` (or to a probability under ``apply_mask=False``),
+ and sigmoid of either always clears the cut, so every ray would read as kept.
+ Use :func:`validity_from_video` for a decoded clip, or pass the mask through
+ ``validity`` here, which is taken as a probability and only cut.
+
+ One-channel V0 reconstructions carry no mask channel and fall back to the
+ range sentinel.
+ """
+ if validity is not None:
+ return validity_mask(validity)
+ if reconstruction.shape[1] >= 3:
+ return validity_from_mask_logits(reconstruction[:, 2:3])
+ return validity_from_range_sentinel(
+ reconstruction[:, :1],
+ mask_input_min_value=mask_input_min_value,
+ mask_input_eps=mask_input_eps,
+ )
+
+
+def as_lidar_clip(value: torch.Tensor, *, name: str = "LiDAR") -> torch.Tensor:
+ """Peel a packed ``[1,C,T,H,W]`` batch down to ``[C,T,H,W]`` with ``C`` in ``{1, 2, 3}``.
+
+ ``C=1`` is legacy V0 range-only. ``C=2`` is V0 range plus an empty intensity
+ channel. ``C=3`` is V1 metric range, intensity, and mask.
+ """
+ sample = value.detach().float().cpu()
+ if sample.ndim == 5:
+ if sample.shape[0] != 1:
+ raise ValueError(f"Expected one LiDAR sample, got {name} shape {tuple(sample.shape)}")
+ sample = sample[0]
+ if sample.ndim != 4 or sample.shape[0] not in (1, 2, 3):
+ raise ValueError(f"Expected {name} shape [1|2|3,T,H,W] or [1,1|2|3,T,H,W], got {tuple(sample.shape)}")
+ return sample
+
+
+def prepare_range_for_display(
+ lidar: torch.Tensor,
+ *,
+ min_range: float | None = None,
+ max_range: float | None = None,
+ smoothing: RangeSmoothing | None = None,
+) -> tuple[torch.Tensor, torch.Tensor | None, float, float]:
+ """Turn a V0 or V1 clip into normalized range, validity, and colour bounds.
+
+ Both versions colorize over ``[MIN_RANGE_METERS, MAX_RANGE_METERS]`` (5--100 m),
+ the Drive-Dreams span the V0 tokenizer and V1 transfer recipes use.
+
+ A clip with fewer than three channels is still in the network's normalized space
+ with dropped rays at ``-1``, so validity is left to the renderer to read off that
+ sentinel. That covers the legacy single-channel V0 prepare and a tokenizer-training
+ reconstruction; channel 1, when present, is an empty intensity slot, ignored here.
+
+ A three-channel clip is metric range, unit intensity, and a mask -- what the
+ dataloader prepares and what both tokenizers' ``decode()`` returns. Channel 2 is
+ ``{0, 1}`` on a dataloader clip and on a decoded one alike, so a plain half cut
+ reads it;
+ when a caller decodes with ``apply_mask=False`` the channel holds the
+ probability instead and the same half cut reproduces the default. A ray is
+ kept only when its metric range is also positive, because the tokenizer
+ zeroes the range of a dropped ray and renormalizing that zero would paint a
+ centimetre ghost at the near end of the colormap.
+
+ ``smoothing`` filters the metric range against the resolved mask before
+ normalizing, so a caller that unprojects the result gets the filtered
+ geometry. It needs both channels and is therefore ignored on a clip with
+ fewer than three, where the mask is left to the renderer's sentinel read.
+ """
+ clip = as_lidar_clip(lidar, name="LiDAR range")
+ lo = float(MIN_RANGE_METERS if min_range is None else min_range)
+ hi = float(MAX_RANGE_METERS if max_range is None else max_range)
+ if hi <= lo:
+ raise ValueError(f"max_range must exceed min_range, got {lo} and {hi}")
+ if clip.shape[0] < 3:
+ return clip[0], None, lo, hi
+
+ metric = clip[0]
+ valid = (metric > 0.0) & validity_mask(clip[2])
+ if smoothing is not None:
+ metric, valid = smooth_metric_range(metric, valid, smoothing=smoothing)
+ normalized = 2.0 * (metric - lo) / (hi - lo) - 1.0
+ return normalized, valid, lo, hi
diff --git a/cosmos_framework/model/generator/tokenizers/lidar/postprocessing_test.py b/cosmos_framework/model/generator/tokenizers/lidar/postprocessing_test.py
new file mode 100644
index 000000000..f99f19288
--- /dev/null
+++ b/cosmos_framework/model/generator/tokenizers/lidar/postprocessing_test.py
@@ -0,0 +1,308 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+
+import pytest
+import torch
+
+from cosmos_framework.model.generator.tokenizers.lidar.postprocessing import (
+ DEFAULT_VALIDITY_THRESHOLD,
+ RangeSmoothing,
+ as_lidar_clip,
+ prepare_range_for_display,
+ smooth_metric_range,
+ validate_validity_threshold,
+ validity_from_mask_logits,
+ validity_from_reconstruction,
+ validity_from_video,
+ validity_mask,
+)
+
+pytestmark = [pytest.mark.L0, pytest.mark.CPU]
+
+
+@pytest.mark.parametrize("threshold", [0.0, 1.0, -0.1, 1.5])
+def test_validity_threshold_rejects_degenerate_cuts(threshold: float) -> None:
+ """Zero keeps every ray and one drops every ray, so both erase the prediction."""
+ with pytest.raises(ValueError, match="validity_threshold"):
+ validate_validity_threshold(threshold)
+
+
+def test_validity_threshold_accepts_the_sweep_points_and_normalizes_type() -> None:
+ """The cut is read from config, so it arrives as whatever OmegaConf held."""
+ assert validate_validity_threshold(DEFAULT_VALIDITY_THRESHOLD) == 0.5
+ assert validate_validity_threshold(0.7) == pytest.approx(0.7)
+ assert validate_validity_threshold(0.99) == pytest.approx(0.99)
+ assert isinstance(validate_validity_threshold("0.7"), float) # type: ignore[arg-type]
+
+
+def test_as_lidar_clip_peels_a_singleton_batch() -> None:
+ clip = torch.zeros((1, 3, 2, 4, 8))
+ assert as_lidar_clip(clip).shape == (3, 2, 4, 8)
+
+
+def test_prepare_range_for_display_treats_v0_with_empty_intensity_as_v0() -> None:
+ clip = torch.full((1, 1, 1, 2), -1.0)
+ clip[0, 0, 0, 1] = 0.0
+ with_empty_intensity = torch.cat((clip, torch.zeros_like(clip)), dim=0)
+ normalized, valid, lo, hi = prepare_range_for_display(with_empty_intensity)
+ assert valid is None
+ assert (lo, hi) == (5.0, 100.0)
+ torch.testing.assert_close(normalized, clip[0])
+
+
+def test_validity_from_mask_logits_sigmoids_then_cuts() -> None:
+ """Training reconstructions carry logits; this is the only place they are cut."""
+ logits = torch.tensor([0.25, 4.0, -4.0, 0.0])
+ assert validity_from_mask_logits(logits).tolist() == [True, True, False, True]
+
+
+def test_validity_mask_reads_a_resolved_channel_without_sigmoid() -> None:
+ assert validity_mask(torch.tensor([1.0, 0.0, 0.56, 0.44])).tolist() == [True, False, True, False]
+
+
+def test_validity_from_video_reads_a_resolved_mask_channel() -> None:
+ video = torch.zeros((1, 3, 1, 1, 2))
+ video[0, 2, 0, 0, 0] = 1.0
+ assert validity_from_video(video).flatten().tolist() == [True, False]
+
+
+def test_validity_from_reconstruction_sigmoids_training_logits() -> None:
+ recon = torch.zeros((1, 3, 1, 1, 2))
+ recon[0, 2, 0, 0, 0] = 0.25
+ recon[0, 2, 0, 0, 1] = -4.0
+ assert validity_from_reconstruction(recon).flatten().tolist() == [True, False]
+
+
+def test_validity_from_reconstruction_does_not_sigmoid_an_explicit_probability() -> None:
+ recon = torch.zeros((1, 3, 1, 1, 1))
+ recon[0, 2] = 4.0
+ probability = torch.tensor([0.44]).reshape(1, 1, 1, 1, 1)
+ assert bool(validity_from_reconstruction(recon, probability).item()) is False
+
+
+def test_prepare_range_for_display_leaves_v0_on_its_sentinel() -> None:
+ clip = torch.full((1, 1, 1, 2), -1.0)
+ clip[0, 0, 0, 1] = 0.0
+ normalized, valid, lo, hi = prepare_range_for_display(clip)
+ assert valid is None
+ assert (lo, hi) == (5.0, 100.0)
+ torch.testing.assert_close(normalized, clip[0])
+
+
+def test_prepare_range_for_display_reads_a_resolved_v1_mask() -> None:
+ """Channel 2 is ``{0, 1}`` on a dataloader clip and on a decoded one alike.
+
+ The tokenizer resolves its own cut, so display only reads the answer back.
+ """
+ lidar = torch.zeros((3, 1, 1, 3))
+ lidar[0] = 50.0
+ lidar[2, :, :, 0] = 1.0
+ lidar[2, :, :, 1] = 0.0
+ lidar[2, :, :, 2] = 1.0
+
+ _, valid, lo, hi = prepare_range_for_display(lidar)
+
+ assert (lo, hi) == (5.0, 100.0)
+ assert valid is not None
+ assert valid.flatten().tolist() == [True, False, True]
+
+
+def test_prepare_range_for_display_reads_an_unmasked_decode_as_probability() -> None:
+ """``apply_mask=False`` leaves the probability in channel 2, not a logit.
+
+ The half cut then reproduces what the tokenizer's default would have done,
+ so an unmasked decode still renders sensibly without a threshold argument.
+ """
+ lidar = torch.zeros((3, 1, 1, 2))
+ lidar[0] = 50.0
+ lidar[2, :, :, 0] = 0.56
+ lidar[2, :, :, 1] = 0.44
+
+ _, valid, _, _ = prepare_range_for_display(lidar)
+
+ assert valid is not None
+ assert valid.flatten().tolist() == [True, False]
+
+
+def test_prepare_range_for_display_drops_zeroed_range_even_when_the_mask_is_set() -> None:
+ """A raised tokenizer threshold zeroes range; a stale mask must not resurrect it."""
+ lidar = torch.zeros((3, 1, 1, 2))
+ lidar[2] = 1.0
+ lidar[0, :, :, 0] = 50.0
+ lidar[0, :, :, 1] = 0.0
+
+ _, valid, _, _ = prepare_range_for_display(lidar)
+
+ assert valid is not None
+ assert valid.flatten().tolist() == [True, False]
+
+
+@pytest.mark.parametrize(
+ "kwargs",
+ [
+ {"azimuth_kernel_size": 2},
+ {"azimuth_kernel_size": 1},
+ {"elevation_kernel_size": 0},
+ {"elevation_kernel_size": 2},
+ {"spatial_sigma": 0.0},
+ {"range_sigma_m": 0.0},
+ {"range_sigma_ratio": -0.1},
+ # Above what the default window's four neighbours can contribute, so no
+ # ray could clear it and the filter would drop the sweep.
+ {"min_neighbor_support": 5.0},
+ {"strength": 1.5},
+ ],
+)
+def test_range_smoothing_rejects_a_window_that_cannot_estimate_a_surface(kwargs: dict[str, float]) -> None:
+ with pytest.raises(ValueError):
+ RangeSmoothing(**kwargs) # type: ignore[arg-type]
+
+
+def test_smoothing_reads_one_beam_unless_asked_for_more() -> None:
+ """The default window is a row, because the other beam is at another distance.
+
+ Down a column each sample is a different laser at a different elevation, so
+ on ground the beams above and below are legitimately at another distance --
+ close enough to weigh, which is what makes them harmful rather than merely
+ useless. Here they sit 0.3 m out, well inside the 1.1 m bandwidth at this
+ range, and a row-only window leaves the middle beam exactly where it was.
+ Widening the window is the knob that reaches them again.
+ """
+ ranges = torch.full((1, 3, 8), 30.3)
+ ranges[0, 1, :] = 30.0
+ valid = torch.ones_like(ranges, dtype=torch.bool)
+
+ per_row, _ = smooth_metric_range(ranges, valid)
+ across_beams, _ = smooth_metric_range(
+ ranges, valid, smoothing=RangeSmoothing(elevation_kernel_size=3, min_neighbor_support=2.0)
+ )
+
+ torch.testing.assert_close(per_row[0, 1], torch.full((8,), 30.0))
+ assert across_beams[0, 1].min().item() > 30.15
+
+
+def test_smooth_metric_range_pulls_a_jittered_wall_flat() -> None:
+ """Per-ray jitter on a flat surface is what sprays the BEV along each ray."""
+ flat = torch.full((1, 5, 8), 30.0)
+ jittered = flat.clone()
+ jittered[0, 2, 3] = 30.4
+ jittered[0, 1, 5] = 29.6
+ valid = torch.ones_like(flat, dtype=torch.bool)
+
+ smoothed, kept = smooth_metric_range(jittered, valid)
+
+ assert bool(kept.all())
+ # A weighted mean keeps a fraction of the ray's own vote, so it lands near the
+ # surface rather than exactly on it. One row is four neighbours rather than a
+ # square's twenty-four, so that fraction is larger and a lone spike keeps more
+ # of itself: a third of this 0.4 m one survives where a twentieth used to.
+ # Real jitter is on every ray at once, where averaging fewer but better
+ # neighbours still wins; see the measurements on RangeSmoothing.
+ torch.testing.assert_close(smoothed, flat, atol=0.15, rtol=0.0)
+
+
+def test_smooth_metric_range_drops_a_flying_pixel_instead_of_averaging_it_in() -> None:
+ """A ray far off its neighbours spans a depth gap, and weighting alone cannot fix it.
+
+ Its neighbours are many range sigma away, so they weigh nothing and the filter
+ would hand the ray straight back. Only the support rule removes it.
+ """
+ ranges = torch.full((1, 5, 8), 30.0)
+ ranges[0, 2, 3] = 70.0
+ valid = torch.ones_like(ranges, dtype=torch.bool)
+
+ smoothed, kept = smooth_metric_range(ranges, valid)
+
+ assert not bool(kept[0, 2, 3])
+ assert bool(kept.sum() == kept.numel() - 1)
+ # A dropped ray keeps its measured range so nothing invents geometry behind the mask.
+ assert smoothed[0, 2, 3] == pytest.approx(70.0)
+
+
+def test_smooth_metric_range_widens_the_bandwidth_with_distance() -> None:
+ """Adjacent beams on the same ground diverge with range, so a fixed bandwidth over-drops far."""
+ smoothing = RangeSmoothing(range_sigma_m=0.2, range_sigma_ratio=0.03)
+ near = torch.full((1, 5, 8), 10.0)
+ far = torch.full((1, 5, 8), 80.0)
+ near[0, 2, 3] += 2.0
+ far[0, 2, 3] += 2.0
+ valid = torch.ones_like(near, dtype=torch.bool)
+
+ _, near_kept = smooth_metric_range(near, valid, smoothing=smoothing)
+ _, far_kept = smooth_metric_range(far, valid, smoothing=smoothing)
+
+ assert not bool(near_kept[0, 2, 3])
+ assert bool(far_kept[0, 2, 3])
+
+
+def test_smooth_metric_range_drops_isolated_speckle() -> None:
+ """One lit ray in an empty window is noise, not a surface worth keeping."""
+ ranges = torch.zeros((1, 5, 8))
+ valid = torch.zeros_like(ranges, dtype=torch.bool)
+ ranges[0, 2, 3] = 30.0
+ valid[0, 2, 3] = True
+
+ _, kept = smooth_metric_range(ranges, valid)
+
+ assert not bool(kept.any())
+
+
+def test_smooth_metric_range_ignores_dropped_neighbors() -> None:
+ """Masked rays carry no distance, so giving them weight would pull the estimate down."""
+ ranges = torch.full((1, 3, 8), 30.0)
+ ranges[0, :, 4] = 0.0
+ valid = torch.ones_like(ranges, dtype=torch.bool)
+ valid[0, :, 4] = False
+
+ smoothed, _ = smooth_metric_range(ranges, valid)
+
+ torch.testing.assert_close(smoothed[0, 1, 3], torch.tensor(30.0))
+ torch.testing.assert_close(smoothed[0, 1, 5], torch.tensor(30.0))
+
+
+def test_smooth_metric_range_wraps_across_the_azimuth_seam() -> None:
+ """The canvas spans a full turn, so its first and last columns are neighbours."""
+ ranges = torch.full((1, 3, 8), 30.0)
+ ranges[0, 1, 0] = 30.4
+ valid = torch.ones_like(ranges, dtype=torch.bool)
+
+ seam_only = ranges.clone()
+ seam_only[0, :, 1:-1] = 0.0
+ seam_valid = valid.clone()
+ seam_valid[0, :, 1:-1] = False
+
+ smoothed, _ = smooth_metric_range(ranges, valid)
+ wrapped, wrapped_kept = smooth_metric_range(seam_only, seam_valid)
+
+ torch.testing.assert_close(smoothed[0, 1, 0], torch.tensor(30.0), atol=0.15, rtol=0.0)
+ # The column-zero window has no support without the wrap, so only the wrap keeps this
+ # ray. Its pull is weaker than above because the wrapped column is all the support
+ # there is, and the ray's own vote is a larger share of a smaller total.
+ assert bool(wrapped_kept[0, 1, 0])
+ torch.testing.assert_close(wrapped[0, 1, 0], torch.tensor(30.0), atol=0.3, rtol=0.0)
+
+
+def test_smooth_metric_range_strength_interpolates_toward_the_filtered_estimate() -> None:
+ ranges = torch.full((1, 3, 8), 30.0)
+ ranges[0, 1, 4] = 30.4
+ valid = torch.ones_like(ranges, dtype=torch.bool)
+
+ full, _ = smooth_metric_range(ranges, valid)
+ half, _ = smooth_metric_range(ranges, valid, smoothing=RangeSmoothing(strength=0.5))
+
+ assert full[0, 1, 4].item() < 30.4
+ assert half[0, 1, 4].item() == pytest.approx(0.5 * (30.4 + full[0, 1, 4].item()))
+
+
+def test_prepare_range_for_display_smooths_before_it_normalizes() -> None:
+ """Unprojection reads this output, so the filter has to land on the metric range."""
+ lidar = torch.full((3, 1, 3, 8), 30.0)
+ lidar[1] = 0.0
+ lidar[2] = 1.0
+ lidar[0, 0, 1, 4] = 30.4
+
+ normalized, valid, lo, hi = prepare_range_for_display(lidar, smoothing=RangeSmoothing())
+
+ assert valid is not None and bool(valid.all())
+ metric = (normalized + 1.0) * 0.5 * (hi - lo) + lo
+ torch.testing.assert_close(metric, torch.full_like(metric, 30.0), atol=0.15, rtol=0.0)
diff --git a/cosmos_framework/model/generator/tokenizers/lidar/preprocessing.py b/cosmos_framework/model/generator/tokenizers/lidar/preprocessing.py
new file mode 100644
index 000000000..df60ab942
--- /dev/null
+++ b/cosmos_framework/model/generator/tokenizers/lidar/preprocessing.py
@@ -0,0 +1,933 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+
+"""LiDAR range-map loading and tokenizer preprocessing."""
+
+from __future__ import annotations
+
+import io
+import tarfile
+from pathlib import Path
+from typing import IO, Any, TypeAlias
+
+import numpy as np
+import torch
+import torch.nn.functional as F
+
+from cosmos_framework.model.generator.tokenizers.lidar.paths import is_remote_uri, s3_backend_args
+from cosmos_framework.model.generator.tokenizers.lidar.range_projection import (
+ V0_TRANSFER_RANGE_PROJECTION,
+ LidarRangeProjectionConfig,
+)
+
+# Native 3-channel tokenizer layout (range + intensity + mask @ 128x3600), i.e.
+# no azimuth downsampling. Spatial compression there is asymmetric: H/8, W/16
+# (patch 2x4 plus two 2x2 merges).
+DEFAULT_RANGE_PROJECTION = LidarRangeProjectionConfig()
+RANGE_HEIGHT = DEFAULT_RANGE_PROJECTION.native_height
+RANGE_RAW_WIDTH = DEFAULT_RANGE_PROJECTION.native_width
+RANGE_DOWNSAMPLED_WIDTH = V0_TRANSFER_RANGE_PROJECTION.semantic_width
+TOKENIZER_WIDTH = V0_TRANSFER_RANGE_PROJECTION.model_width
+# Shared metric span for V0 and V1: Drive-Dreams [5, 100] m.
+MIN_RANGE_METERS = DEFAULT_RANGE_PROJECTION.min_range_m
+MAX_RANGE_METERS = DEFAULT_RANGE_PROJECTION.max_range_m
+# Network-space fill for a dropped ray after the [-1, 1] affine.
+INVALID_NORMALIZED_VALUE = -1.0
+# The tokenizer trained on 8 sampled sweeps plus the duplicated frame-0 prepend.
+TOKENIZER_SAMPLE_FRAMES = 8
+
+TarSource: TypeAlias = str | Path | bytes | bytearray | memoryview | IO[bytes]
+ExpectedClipKey: TypeAlias = str | tuple[str, ...]
+
+
+def _backend_args_or_default(backend_args: dict[str, Any] | None) -> dict[str, Any] | None:
+ return backend_args if backend_args is not None else s3_backend_args()
+
+
+def _open_tar_archive(
+ tar_source: TarSource,
+ *,
+ backend_args: dict[str, Any] | None = None,
+) -> tarfile.TarFile:
+ """Open a local, remote, or in-memory LidarGEN clip tar."""
+ if isinstance(tar_source, (bytes, bytearray, memoryview)):
+ return tarfile.open(fileobj=io.BytesIO(tar_source), mode="r")
+ if hasattr(tar_source, "read"):
+ return tarfile.open(fileobj=tar_source, mode="r")
+ path = str(tar_source)
+ if is_remote_uri(path):
+ from cosmos_framework.utils.easy_io import easy_io
+
+ payload = easy_io.get(path, backend_args=_backend_args_or_default(backend_args))
+ return tarfile.open(fileobj=io.BytesIO(payload), mode="r")
+ return tarfile.open(path, "r")
+
+
+def normalize_range_map(
+ range_map: np.ndarray | torch.Tensor,
+ *,
+ min_range: float = MIN_RANGE_METERS,
+ max_range: float = MAX_RANGE_METERS,
+) -> np.ndarray | torch.Tensor:
+ """Normalize metric ranges to ``[-1, 1]`` with invalid rays at ``-1``."""
+ if max_range <= min_range:
+ raise ValueError(f"max_range must exceed min_range, got {min_range=} and {max_range=}")
+ if isinstance(range_map, torch.Tensor):
+ clipped = range_map.clamp(min_range, max_range) # [T,H,W]
+ else:
+ clipped = np.clip(range_map, min_range, max_range)
+ unit = (clipped - min_range) / (max_range - min_range) # [T,H,W]
+ return unit * 2.0 - 1.0 # [T,H,W]
+
+
+def normalize_intensity_values(intensities: np.ndarray) -> np.ndarray:
+ """Normalize LiDAR return intensities to finite float32 values in ``[0, 1]``."""
+ values = np.asarray(intensities)
+ if not np.isfinite(values).all():
+ raise ValueError("LiDAR intensity values must be finite")
+ if values.size == 0:
+ return values.astype(np.float32)
+ if np.issubdtype(values.dtype, np.integer):
+ minimum = int(values.min())
+ maximum = int(values.max())
+ if minimum < 0 or maximum > 255:
+ raise ValueError(f"LiDAR integer intensity values must lie in [0,255], got [{minimum},{maximum}]")
+ return values.astype(np.float32) / 255.0
+
+ values = values.astype(np.float32)
+ minimum = float(values.min())
+ maximum = float(values.max())
+ tolerance = 1e-3
+ if minimum < -tolerance or maximum > 1.0 + tolerance:
+ raise ValueError(
+ f"LiDAR floating-point intensity values must lie in [0,1] within tolerance {tolerance}, "
+ f"got [{minimum},{maximum}]"
+ )
+ return np.clip(values, 0.0, 1.0)
+
+
+def network_range_to_metric(
+ range_map: np.ndarray | torch.Tensor,
+ *,
+ min_range: float = MIN_RANGE_METERS,
+ max_range: float = MAX_RANGE_METERS,
+ clamp: bool = True,
+) -> np.ndarray | torch.Tensor:
+ """Invert :func:`normalize_range_map`: ``[-1, 1]`` back to meters, without masking.
+
+ This is the only affine that maps a normalized range to meters; the callers
+ that also drop rays layer their own mask on top, because they disagree about
+ what a dropped ray becomes.
+
+ ``clamp`` bounds the input to ``[-1, 1]`` first, which matters for a raw
+ network prediction: without it an overshoot leaves the sensor's range
+ envelope and reads as a return that the tokenizer never could have encoded.
+ """
+ if clamp:
+ if isinstance(range_map, torch.Tensor):
+ range_map = range_map.clamp(-1.0, 1.0)
+ else:
+ range_map = np.clip(range_map, -1.0, 1.0)
+ return (range_map + 1.0) * 0.5 * (max_range - min_range) + min_range
+
+
+def unit_intensity_to_network(intensity: torch.Tensor) -> torch.Tensor:
+ """Map unit intensity ``[0, 1]`` to the network's ``[-1, 1]``."""
+ return intensity.clamp(0.0, 1.0) * 2.0 - 1.0
+
+
+def network_intensity_to_unit(normalized: torch.Tensor) -> torch.Tensor:
+ """Invert :func:`unit_intensity_to_network`: ``[-1, 1]`` back to ``[0, 1]``."""
+ return (normalized.clamp(-1.0, 1.0) + 1.0) * 0.5
+
+
+def unnormalize_range_map(
+ range_map: np.ndarray | torch.Tensor,
+ *,
+ min_range: float = MIN_RANGE_METERS,
+ max_range: float = MAX_RANGE_METERS,
+ valid_mask: np.ndarray | torch.Tensor | None = None,
+) -> tuple[np.ndarray | torch.Tensor, np.ndarray | torch.Tensor]:
+ """Convert normalized ranges to meters and zero invalid rays.
+
+ The input is not clamped: a caller that passes a raw prediction and no
+ ``valid_mask`` relies on the derived mask below to drop an overshoot, rather
+ than on the affine folding it back to the range envelope.
+ """
+ metric = network_range_to_metric(range_map, min_range=min_range, max_range=max_range, clamp=False) # [T,H,W]
+ if valid_mask is None:
+ valid_mask = (metric > min_range) & (metric < max_range) # [T,H,W]
+ if isinstance(metric, torch.Tensor):
+ metric = torch.where(valid_mask, metric, torch.zeros_like(metric)) # [T,H,W]
+ else:
+ metric = np.where(valid_mask, metric, np.zeros_like(metric))
+ return metric, valid_mask
+
+
+class RangeMapDownsampler:
+ """Downsample a dense range map by retaining the nearest nonzero return."""
+
+ def __init__(self, row_factor: int = 1, col_factor: int = 3) -> None:
+ if row_factor < 1 or col_factor < 1:
+ raise ValueError("Downsampling factors must be positive")
+ self.row_factor = row_factor
+ self.col_factor = col_factor
+
+ @staticmethod
+ def _minimum_nonzero(groups: np.ndarray, axis: int) -> np.ndarray:
+ sentinel = np.finfo(groups.dtype).max
+ nonzero = np.where(groups == 0, sentinel, groups)
+ minimum = nonzero.min(axis=axis)
+ return np.where(minimum == sentinel, 0, minimum)
+
+ @staticmethod
+ def _nearest_nonzero_range_and_value(
+ range_groups: np.ndarray,
+ value_groups: np.ndarray,
+ *,
+ axis: int,
+ ) -> tuple[np.ndarray, np.ndarray]:
+ sentinel = np.finfo(range_groups.dtype).max
+ nonzero_range = np.where(range_groups == 0, sentinel, range_groups)
+ nearest_index = nonzero_range.argmin(axis=axis, keepdims=True)
+ nearest_range = np.take_along_axis(nonzero_range, nearest_index, axis=axis).squeeze(axis)
+ nearest_value = np.take_along_axis(value_groups, nearest_index, axis=axis).squeeze(axis)
+ valid = nearest_range != sentinel
+ return np.where(valid, nearest_range, 0), np.where(valid, nearest_value, 0)
+
+ def __call__(self, range_map: np.ndarray) -> np.ndarray:
+ """Downsample ``[T,H,W]`` range maps."""
+ if range_map.ndim != 3:
+ raise ValueError(f"Expected [T,H,W], got shape {range_map.shape}")
+ result = np.asarray(range_map, dtype=np.float32)
+ frames, height, width = result.shape
+ if height % self.row_factor:
+ raise ValueError(f"Height {height} is not divisible by {self.row_factor}")
+ if self.row_factor != 1:
+ grouped_rows = result.reshape(
+ frames,
+ height // self.row_factor,
+ self.row_factor,
+ width,
+ )
+ result = self._minimum_nonzero(grouped_rows, axis=2)
+
+ frames, height, width = result.shape
+ if width % self.col_factor:
+ raise ValueError(f"Width {width} is not divisible by {self.col_factor}")
+ if self.col_factor != 1:
+ grouped_cols = result.reshape(
+ frames,
+ height,
+ width // self.col_factor,
+ self.col_factor,
+ )
+ result = self._minimum_nonzero(grouped_cols, axis=3)
+ return result
+
+ def downsample_with_values(
+ self,
+ range_map: np.ndarray,
+ values: np.ndarray,
+ ) -> tuple[np.ndarray, np.ndarray]:
+ """Downsample ranges and retain values from each nearest return."""
+ if range_map.shape != values.shape or range_map.ndim != 3:
+ raise ValueError(f"Expected matching [T,H,W] arrays, got {range_map.shape} and {values.shape}")
+ ranges = np.asarray(range_map, dtype=np.float32)
+ result_values = np.asarray(values, dtype=np.float32)
+ frames, height, width = ranges.shape
+ if height % self.row_factor or width % self.col_factor:
+ raise ValueError(
+ f"Shape {(height, width)} is not divisible by factors {(self.row_factor, self.col_factor)}"
+ )
+ if self.row_factor != 1:
+ ranges, result_values = self._nearest_nonzero_range_and_value(
+ ranges.reshape(frames, height // self.row_factor, self.row_factor, width),
+ result_values.reshape(frames, height // self.row_factor, self.row_factor, width),
+ axis=2,
+ )
+ frames, height, width = ranges.shape
+ if self.col_factor != 1:
+ ranges, result_values = self._nearest_nonzero_range_and_value(
+ ranges.reshape(frames, height, width // self.col_factor, self.col_factor),
+ result_values.reshape(frames, height, width // self.col_factor, self.col_factor),
+ axis=3,
+ )
+ return ranges, result_values
+
+
+def validate_lidar_layout(*, semantic_width: int, model_width: int) -> tuple[int, int]:
+ """Validate semantic/model azimuth widths and return their integer values."""
+ semantic_width = int(semantic_width)
+ model_width = int(model_width)
+ if semantic_width < 1 or RANGE_RAW_WIDTH % semantic_width:
+ raise ValueError(
+ f"semantic_width must be a positive divisor of native width {RANGE_RAW_WIDTH}, got {semantic_width}"
+ )
+ if model_width < semantic_width:
+ raise ValueError(f"model_width must be at least semantic_width, got {model_width} < {semantic_width}")
+ if (model_width - semantic_width) % 2:
+ raise ValueError(
+ "model_width - semantic_width must be even for symmetric circular padding, "
+ f"got {model_width} - {semantic_width}"
+ )
+ return semantic_width, model_width
+
+
+def range_video_layout_for_network_width(network_width: int) -> tuple[int, int | None]:
+ """``(prepare_width, resize_width)`` for a network reading ``network_width`` columns.
+
+ Prefer :data:`V0_TRANSFER_RANGE_PROJECTION` / :data:`V1_TRANSFER_RANGE_PROJECTION`
+ at call sites. This helper only reconstructs that V0 pool-then-resize layout from
+ a bare network width, for callers that have not yet loaded a projection object.
+ """
+ width = int(network_width)
+ if width < 1:
+ raise ValueError(f"network_width must be positive, got {network_width}")
+ if width == V0_TRANSFER_RANGE_PROJECTION.model_width:
+ return V0_TRANSFER_RANGE_PROJECTION.semantic_width, width
+ if RANGE_RAW_WIDTH % width == 0:
+ return width, None
+ raise ValueError(
+ f"No pooling layout for network width {width}: it does not divide native "
+ f"{RANGE_RAW_WIDTH} and is not the V0 transfer width "
+ f"{V0_TRANSFER_RANGE_PROJECTION.model_width}"
+ )
+
+
+def circular_pad_range_video_3ch(video: torch.Tensor, *, model_width: int) -> torch.Tensor:
+ """Symmetrically circular-pad ``[..., W]`` from semantic to model width."""
+ if video.ndim < 1:
+ raise ValueError("Expected a tensor with a width dimension")
+ semantic_width, model_width = validate_lidar_layout(
+ semantic_width=int(video.shape[-1]),
+ model_width=model_width,
+ )
+ pad_each = (model_width - semantic_width) // 2
+ if pad_each == 0:
+ return video
+ if pad_each > semantic_width:
+ raise ValueError(f"Circular padding per side cannot exceed semantic width, got {pad_each} > {semantic_width}")
+ return torch.cat((video[..., -pad_each:], video, video[..., :pad_each]), dim=-1)
+
+
+def crop_range_video_width(video: torch.Tensor, *, semantic_width: int) -> torch.Tensor:
+ """Centrally crop a circular-padded ``[..., W_model]`` tensor."""
+ if video.ndim < 1:
+ raise ValueError("Expected a tensor with a width dimension")
+ semantic_width, model_width = validate_lidar_layout(
+ semantic_width=semantic_width,
+ model_width=int(video.shape[-1]),
+ )
+ crop_each = (model_width - semantic_width) // 2
+ if crop_each == 0:
+ return video
+ return video[..., crop_each : crop_each + semantic_width]
+
+
+def resize_range_video_width(video: torch.Tensor, *, width: int) -> torch.Tensor:
+ """Nearest-resize every channel of a ``[C,T,H,W]`` clip along azimuth."""
+ if video.ndim != 4:
+ raise ValueError(f"Expected [C,T,H,W], got {tuple(video.shape)}")
+ if width < 1:
+ raise ValueError(f"width must be positive, got {width}")
+ if video.shape[-1] == width:
+ return video
+ frames = video.permute(1, 0, 2, 3) # [T,C,H,W]
+ frames = F.interpolate(frames, size=(video.shape[-2], width), mode="nearest-exact") # [T,C,H,W_out]
+ return frames.permute(1, 0, 2, 3).contiguous() # [C,T,H,W_out]
+
+
+def prepare_lidar_video(
+ metric_range_maps: np.ndarray,
+ *,
+ prepend_first_frame: bool = True,
+ output_width: int = TOKENIZER_WIDTH,
+) -> torch.Tensor:
+ """Prepare raw range maps as tokenizer input ``[1,T,128,1024]``."""
+ if metric_range_maps.ndim != 3 or metric_range_maps.shape[1:] != (RANGE_HEIGHT, RANGE_RAW_WIDTH):
+ raise ValueError(f"Expected [T,{RANGE_HEIGHT},{RANGE_RAW_WIDTH}], got {metric_range_maps.shape}")
+ downsampled = RangeMapDownsampler()(metric_range_maps)
+ normalized = normalize_range_map(downsampled)
+ video = torch.from_numpy(np.asarray(normalized, dtype=np.float32))[:, None] # [T,1,H,W]
+ if prepend_first_frame:
+ video = torch.cat((video[:1], video), dim=0) # [T+1,1,H,W]
+ if video.shape[-1] != output_width:
+ video = F.interpolate(video, size=(RANGE_HEIGHT, output_width), mode="nearest-exact") # [T,1,H,W_out]
+ return video.permute(1, 0, 2, 3).contiguous() # [1,T,H,W]
+
+
+def prepare_lidar_video_3ch(
+ metric_range_maps: np.ndarray,
+ intensity_maps: np.ndarray | None,
+ *,
+ prepend_first_frame: bool = True,
+ range_projection: LidarRangeProjectionConfig = DEFAULT_RANGE_PROJECTION,
+) -> torch.Tensor:
+ """Prepare metric range, unit intensity, and mask as ``[3,T,128,W_model]``.
+
+ ``metric_range_maps`` and ``intensity_maps`` are dense ``[T,128,3600]`` arrays
+ (intensity already in ``[0,1]`` on returns, zeros elsewhere). Native rays are
+ nearest-return pooled to ``semantic_width`` with matched intensity, then
+ packed to ``model_width`` by circular padding or resize. Range remains in metres,
+ intensity remains in ``[0,1]``, and mask is ``{0,1}``. Configured range
+ filtering and network normalization belong to the tokenizer model.
+
+ ``intensity_maps=None`` marks a modality that carries no return intensity, such
+ as an HD-map rangemap. Its occupied rays take unit intensity rather than zero:
+ the shared normalization maps zero onto the same value it fills invalid rays
+ with, which would leave the control's populated rays indistinguishable from
+ empty sky on the intensity channel.
+ """
+ semantic_width = range_projection.semantic_width
+ expected_shape = (range_projection.native_height, range_projection.native_width)
+ if metric_range_maps.ndim != 3 or metric_range_maps.shape[1:] != expected_shape:
+ raise ValueError(f"Expected range [T,{expected_shape[0]},{expected_shape[1]}], got {metric_range_maps.shape}")
+ if intensity_maps is not None and intensity_maps.shape != metric_range_maps.shape:
+ raise ValueError(f"Intensity maps must match range shape {metric_range_maps.shape}, got {intensity_maps.shape}")
+ col_factor = range_projection.column_pool_factor
+ if col_factor == 1:
+ down_range = np.asarray(metric_range_maps, dtype=np.float32)
+ down_intensity = None if intensity_maps is None else np.asarray(intensity_maps, dtype=np.float32)
+ elif intensity_maps is None:
+ # A constant intensity holds no nearest-return information for the pooler to
+ # carry, so pool range alone and fill the surviving rays below.
+ down_range = RangeMapDownsampler(col_factor=col_factor)(metric_range_maps)
+ down_intensity = None
+ else:
+ downsampler = RangeMapDownsampler(col_factor=col_factor)
+ down_range, down_intensity = downsampler.downsample_with_values(metric_range_maps, intensity_maps)
+ if down_range.shape[-1] != semantic_width:
+ raise ValueError(f"Expected semantic width {semantic_width}, got {down_range.shape[-1]}")
+
+ valid = down_range > 0.0 # [T,H,W]
+ metric_range = np.where(valid, down_range, 0.0).astype(np.float32)
+ if down_intensity is None:
+ intensity_unit = valid.astype(np.float32)
+ else:
+ intensity_unit = np.clip(down_intensity.astype(np.float32), 0.0, 1.0)
+ intensity_unit = np.where(valid, intensity_unit, 0.0).astype(np.float32)
+ mask = valid.astype(np.float32)
+
+ stacked = np.stack([metric_range, intensity_unit, mask], axis=1) # [T,3,H,W]
+ video = torch.from_numpy(stacked)
+ if prepend_first_frame:
+ video = torch.cat((video[:1], video), dim=0) # [T+1,3,H,W]
+ video = video.permute(1, 0, 2, 3).contiguous() # [3,T,H,W_semantic]
+ return apply_model_width_transform(video, range_projection=range_projection).contiguous()
+
+
+def apply_model_width_transform(
+ video: torch.Tensor, # [C,T,H,W_semantic]
+ *,
+ range_projection: LidarRangeProjectionConfig,
+) -> torch.Tensor:
+ """Pack a semantic-width clip onto the tokenizer's model-width canvas."""
+ if video.ndim != 4:
+ raise ValueError(f"Expected [C,T,H,W], got {tuple(video.shape)}")
+ if video.shape[-1] != range_projection.semantic_width:
+ raise ValueError(f"Expected semantic width {range_projection.semantic_width}, got {video.shape[-1]}")
+ if range_projection.model_width_transform == "circular_pad":
+ return circular_pad_range_video_3ch(video, model_width=range_projection.model_width)
+ if range_projection.model_width_transform == "resize":
+ return resize_range_video_width(video, width=range_projection.model_width)
+ raise ValueError(f"Unsupported model width transform {range_projection.model_width_transform!r}")
+
+
+def undo_model_width_transform(
+ video: torch.Tensor, # [C,T,H,W_model]
+ *,
+ range_projection: LidarRangeProjectionConfig,
+) -> torch.Tensor: # [C,T,H,W_semantic]
+ """Undo model-width packing so eval artifacts sit on the physical grid.
+
+ ``circular_pad`` is a center crop back to ``semantic_width`` (a no-op when the
+ two widths already match, as in transfer V1 at 1800). ``resize`` is a nearest
+ resize from ``model_width`` back to ``semantic_width`` (V0: 1024 to 1200).
+ LiDARBench and Hyperion score on that physical grid, not the VAE canvas.
+ """
+ if video.ndim != 4:
+ raise ValueError(f"Expected [C,T,H,W], got {tuple(video.shape)}")
+ if range_projection.model_width_transform == "circular_pad":
+ return crop_range_video_width(video, semantic_width=range_projection.semantic_width)
+ if range_projection.model_width_transform == "resize":
+ if video.shape[-1] != range_projection.model_width:
+ raise ValueError(
+ f"Expected model width {range_projection.model_width} for resize artifacts, got {video.shape[-1]}"
+ )
+ return resize_range_video_width(video, width=range_projection.semantic_width)
+ raise ValueError(f"Unsupported model width transform {range_projection.model_width_transform!r}")
+
+
+def rasterize_sparse_range_frame(
+ rows: np.ndarray,
+ cols: np.ndarray,
+ ranges: np.ndarray,
+) -> np.ndarray:
+ """Rasterize sparse LiDAR returns into a dense ``[128,3600]`` range map."""
+ if rows.shape != cols.shape or rows.shape != ranges.shape:
+ raise ValueError(
+ f"Sparse range-map arrays must share one length: rows={rows.shape}, cols={cols.shape}, ranges={ranges.shape}"
+ )
+ row_index = rows.astype(np.int64)
+ col_index = cols.astype(np.int64)
+ if row_index.size and (row_index.min() < 0 or row_index.max() >= RANGE_HEIGHT):
+ raise ValueError(f"Range-map rows must lie in [0, {RANGE_HEIGHT}), got max {row_index.max()}")
+ if col_index.size and (col_index.min() < 0 or col_index.max() >= RANGE_RAW_WIDTH):
+ raise ValueError(f"Range-map columns must lie in [0, {RANGE_RAW_WIDTH}), got max {col_index.max()}")
+ dense = np.zeros((RANGE_HEIGHT, RANGE_RAW_WIDTH), dtype=np.float32) # [H,W]
+ dense[row_index, col_index] = ranges.astype(np.float32) # [H,W]
+ return dense
+
+
+def _load_npz_array(path: Path) -> np.ndarray:
+ with np.load(path) as payload:
+ if "arr_0" not in payload:
+ raise KeyError(f"Expected key 'arr_0' in {path}")
+ return np.asarray(payload["arr_0"])
+
+
+def _load_tar_npz_array(archive: tarfile.TarFile, member_name: str) -> np.ndarray:
+ member = archive.extractfile(member_name)
+ if member is None:
+ raise ValueError(f"Could not read tar member {member_name!r}")
+ with np.load(io.BytesIO(member.read()), allow_pickle=False) as payload:
+ if "arr_0" not in payload:
+ raise KeyError(f"Expected key 'arr_0' in tar member {member_name!r}")
+ return np.asarray(payload["arr_0"])
+
+
+def _sparse_frame_members(archive: tarfile.TarFile) -> dict[int, tuple[str, dict[str, str]]]:
+ """Index sparse rangemap members by original source-frame ID."""
+ suffixes = {
+ "row": ".lidar_row.npz",
+ "col": ".lidar_col.npz",
+ "range": ".lidar_range.npz",
+ "intensity": ".lidar_intensity.npz",
+ }
+ frames: dict[int, tuple[str, dict[str, str]]] = {}
+ for name in archive.getnames():
+ kind = next((key for key, suffix in suffixes.items() if name.endswith(suffix)), None)
+ if kind is None:
+ continue
+ prefix = name.removesuffix(suffixes[kind])
+ try:
+ clip_key, frame_token = prefix.rsplit(".", 1)
+ source_frame_id = int(frame_token)
+ except (ValueError, TypeError) as error:
+ raise ValueError(f"Malformed LiDAR rangemap member name {name!r}") from error
+ existing_clip_key, members = frames.setdefault(source_frame_id, (clip_key, {}))
+ if existing_clip_key != clip_key:
+ raise ValueError(
+ f"LiDAR source frame {source_frame_id} has inconsistent clip keys: "
+ f"{existing_clip_key!r} and {clip_key!r}"
+ )
+ if kind in members:
+ raise ValueError(f"Duplicate LiDAR {kind} member for source frame {source_frame_id}")
+ members[kind] = name
+ if not frames:
+ raise ValueError("No lidar_row/lidar_col/lidar_range arrays found in LiDAR rangemap tar")
+ return frames
+
+
+def _expected_clip_key_set(expected_clip_key: ExpectedClipKey | None) -> set[str] | None:
+ """Return the accepted rangemap clip keys, or None when validation is disabled."""
+ if expected_clip_key is None:
+ return None
+ if isinstance(expected_clip_key, str):
+ return {expected_clip_key}
+ return set(expected_clip_key)
+
+
+def _rasterize_tar_frame(archive: tarfile.TarFile, source_frame_id: int, members: dict[str, str]) -> np.ndarray:
+ missing = {"row", "col", "range"} - members.keys()
+ if missing:
+ raise ValueError(f"Incomplete sparse LiDAR frame {source_frame_id}: missing {sorted(missing)}")
+ rows = _load_tar_npz_array(archive, members["row"])
+ cols = _load_tar_npz_array(archive, members["col"])
+ ranges = _load_tar_npz_array(archive, members["range"])
+ return rasterize_sparse_range_frame(rows, cols, ranges)
+
+
+def _rasterize_tar_frame_with_intensity(
+ archive: tarfile.TarFile,
+ source_frame_id: int,
+ members: dict[str, str],
+) -> tuple[np.ndarray, np.ndarray]:
+ missing = {"row", "col", "range", "intensity"} - members.keys()
+ if missing:
+ raise ValueError(f"Incomplete sparse LiDAR frame {source_frame_id}: missing {sorted(missing)}")
+ rows = _load_tar_npz_array(archive, members["row"])
+ cols = _load_tar_npz_array(archive, members["col"])
+ ranges = _load_tar_npz_array(archive, members["range"])
+ intensities = normalize_intensity_values(_load_tar_npz_array(archive, members["intensity"]))
+ if intensities.shape != ranges.shape:
+ raise ValueError(
+ f"Sparse LiDAR intensity and range arrays must share one length: "
+ f"intensity={intensities.shape}, range={ranges.shape}"
+ )
+ return (
+ rasterize_sparse_range_frame(rows, cols, ranges),
+ rasterize_sparse_range_frame(rows, cols, intensities),
+ )
+
+
+def load_lidar_frames(
+ tar_path: TarSource,
+ *,
+ frame_indices: list[int] | None = None,
+ backend_args: dict[str, Any] | None = None,
+) -> np.ndarray:
+ """Load selected positional sweeps from one LidarGEN tar as ``[T,128,3600]``."""
+ with _open_tar_archive(tar_path, backend_args=backend_args) as archive:
+ frame_members = _sparse_frame_members(archive)
+ source_frame_ids = sorted(frame_members)
+ selected = frame_indices if frame_indices is not None else list(range(len(source_frame_ids)))
+ frames: list[np.ndarray] = []
+ for frame_index in selected:
+ if frame_index < 0 or frame_index >= len(source_frame_ids):
+ raise IndexError(f"Frame {frame_index} is outside [0, {len(source_frame_ids)})")
+ source_frame_id = source_frame_ids[frame_index]
+ _, members = frame_members[source_frame_id]
+ frames.append(_rasterize_tar_frame(archive, source_frame_id, members))
+ return np.stack(frames, axis=0)
+
+
+def load_prepared_range_video(
+ tar_source: TarSource,
+ *,
+ min_source_frame_id: int,
+ max_source_frame_id: int,
+ expected_clip_key: ExpectedClipKey | None = None,
+ backend_args: dict[str, Any] | None = None,
+ max_frames_per_chunk: int = 8,
+ with_intensity: bool = False,
+) -> tuple[torch.Tensor, torch.Tensor | None, list[int]]:
+ """Load aligned range maps as ``[1,T,128,1024]`` plus optional intensity and source-frame IDs.
+
+ Works for both LiDAR rangemap tars (row/col/range[/intensity]) and HD-map rangemap
+ tars that reuse the same sparse member naming without intensity.
+
+ Members are keyed by source-frame ID, so a 10Hz sweep stream inside a 30FPS
+ clip is stored at IDs 0, 3, 6, ... Every sweep in the closed interval is
+ returned; the bounds are IDs on that clock, not sweep ordinals.
+ """
+ if min_source_frame_id > max_source_frame_id:
+ raise ValueError(
+ f"min_source_frame_id must not exceed max_source_frame_id, got "
+ f"{min_source_frame_id} > {max_source_frame_id}"
+ )
+ if max_frames_per_chunk < 1:
+ raise ValueError(f"max_frames_per_chunk must be positive, got {max_frames_per_chunk}")
+ with _open_tar_archive(tar_source, backend_args=backend_args) as archive:
+ frame_members = _sparse_frame_members(archive)
+ expected_clip_keys = _expected_clip_key_set(expected_clip_key)
+ if expected_clip_keys is not None:
+ mismatched_clip_keys = sorted(
+ {clip_key for clip_key, _ in frame_members.values() if clip_key not in expected_clip_keys}
+ )
+ if mismatched_clip_keys:
+ raise ValueError(
+ f"Rangemap clip keys {mismatched_clip_keys!r} do not match WebDataset key(s) "
+ f"{tuple(sorted(expected_clip_keys))!r}"
+ )
+ selected_ids = sorted(
+ source_frame_id
+ for source_frame_id in frame_members
+ if min_source_frame_id <= source_frame_id <= max_source_frame_id
+ )
+ if not selected_ids:
+ raise ValueError(
+ f"No rangemap sweeps overlap source-frame interval [{min_source_frame_id}, {max_source_frame_id}]"
+ )
+ prepared_range_chunks: list[torch.Tensor] = []
+ prepared_intensity_chunks: list[torch.Tensor] = []
+ for chunk_start in range(0, len(selected_ids), max_frames_per_chunk):
+ chunk_ids = selected_ids[chunk_start : chunk_start + max_frames_per_chunk]
+ metric_frames: list[np.ndarray] = []
+ intensity_frames: list[np.ndarray] = []
+ for source_frame_id in chunk_ids:
+ _, members = frame_members[source_frame_id]
+ if with_intensity:
+ metric, intensity = _rasterize_tar_frame_with_intensity(archive, source_frame_id, members)
+ intensity_frames.append(intensity)
+ else:
+ metric = _rasterize_tar_frame(archive, source_frame_id, members)
+ metric_frames.append(metric)
+ prepared_range_chunks.append(prepare_lidar_video(np.stack(metric_frames), prepend_first_frame=False))
+ if with_intensity:
+ downsampler = RangeMapDownsampler()
+ _, selected_intensity = downsampler.downsample_with_values(
+ np.stack(metric_frames),
+ np.stack(intensity_frames),
+ )
+ intensity = torch.from_numpy(selected_intensity)[:, None] # [T,1,H,W]
+ intensity = F.interpolate(
+ intensity,
+ size=(RANGE_HEIGHT, TOKENIZER_WIDTH),
+ mode="nearest-exact",
+ )
+ prepared_intensity_chunks.append(intensity.permute(1, 0, 2, 3).contiguous()) # [1,T,H,W]
+ prepared_range = torch.cat(prepared_range_chunks, dim=1)
+ prepared_intensity = torch.cat(prepared_intensity_chunks, dim=1) if with_intensity else None
+ return prepared_range, prepared_intensity, selected_ids
+
+
+def enforce_tokenizer_clip_length(
+ video: torch.Tensor,
+ frame_ids: list[int],
+ *,
+ num_source_frames: int,
+) -> tuple[torch.Tensor, list[int]]:
+ """Pad or truncate a prepared clip to ``num_source_frames + 1`` timesteps.
+
+ Tokenizer clips prepend a duplicate of the first source sweep, so the
+ expected temporal length is ``num_source_frames + 1``. Short clips are
+ right-padded by repeating the last frame; long clips are truncated.
+ ``frame_ids`` tracks source sweeps only (length ``num_source_frames``).
+ """
+ if num_source_frames < 1:
+ raise ValueError(f"num_source_frames must be positive, got {num_source_frames}")
+ if video.ndim != 4:
+ raise ValueError(f"Expected video [C,T,H,W], got {tuple(video.shape)}")
+ want_t = int(num_source_frames) + 1
+ if video.shape[1] > want_t:
+ video = video[:, :want_t]
+ frame_ids = list(frame_ids)[:num_source_frames]
+ elif video.shape[1] < want_t:
+ pad_t = want_t - video.shape[1]
+ video = torch.cat([video, video[:, -1:].expand(-1, pad_t, -1, -1)], dim=1)
+ last = frame_ids[-1] if frame_ids else 0
+ frame_ids = list(frame_ids) + [last] * pad_t
+ else:
+ frame_ids = list(frame_ids)
+ return video, frame_ids
+
+
+def align_prepared_range_videos(
+ control: torch.Tensor,
+ control_frame_ids: list[int],
+ target: torch.Tensor,
+ target_frame_ids: list[int],
+) -> tuple[torch.Tensor, torch.Tensor, list[int]]:
+ """Keep only the source-frame IDs present in both prepared range videos."""
+ if control.ndim != 4 or target.ndim != 4:
+ raise ValueError(f"Expected [C,T,H,W] tensors, got {tuple(control.shape)} and {tuple(target.shape)}")
+ if control.shape[1] != len(control_frame_ids) or target.shape[1] != len(target_frame_ids):
+ raise ValueError(
+ "Prepared range temporal length must match frame-id lists: "
+ f"control T={control.shape[1]} ids={len(control_frame_ids)}, "
+ f"target T={target.shape[1]} ids={len(target_frame_ids)}"
+ )
+ control_index = {frame_id: idx for idx, frame_id in enumerate(control_frame_ids)}
+ target_index = {frame_id: idx for idx, frame_id in enumerate(target_frame_ids)}
+ shared_ids = sorted(set(control_index) & set(target_index))
+ if not shared_ids:
+ raise ValueError("No overlapping source-frame IDs between control and target rangemaps")
+ control_positions = torch.tensor([control_index[frame_id] for frame_id in shared_ids], dtype=torch.long)
+ target_positions = torch.tensor([target_index[frame_id] for frame_id in shared_ids], dtype=torch.long)
+ return control.index_select(1, control_positions), target.index_select(1, target_positions), shared_ids
+
+
+def list_rangemap_sweep_ids(
+ tar_source: TarSource,
+ *,
+ backend_args: dict[str, Any] | None = None,
+) -> list[int]:
+ """Return sorted source-frame IDs present in a rangemap tar (10 Hz grid)."""
+ with _open_tar_archive(tar_source, backend_args=backend_args) as archive:
+ return sorted(_sparse_frame_members(archive))
+
+
+def load_prepared_lidar_video_3ch(
+ tar_source: TarSource,
+ *,
+ min_source_frame_id: int,
+ max_source_frame_id: int,
+ expected_clip_key: ExpectedClipKey | None = None,
+ backend_args: dict[str, Any] | None = None,
+ max_frames_per_chunk: int = 8,
+ prepend_first_frame: bool = False,
+ range_projection: LidarRangeProjectionConfig = DEFAULT_RANGE_PROJECTION,
+ with_intensity: bool = True,
+) -> tuple[torch.Tensor, list[int]]:
+ """Load metric range, unit intensity and validity plus source-frame IDs.
+
+ HD-map rangemaps do not carry return intensity. For those controls,
+ ``with_intensity=False`` gives every occupied ray unit intensity while
+ preserving range-derived validity.
+ """
+ if min_source_frame_id > max_source_frame_id:
+ raise ValueError(
+ f"min_source_frame_id must not exceed max_source_frame_id, got "
+ f"{min_source_frame_id} > {max_source_frame_id}"
+ )
+ if max_frames_per_chunk < 1:
+ raise ValueError(f"max_frames_per_chunk must be positive, got {max_frames_per_chunk}")
+ with _open_tar_archive(tar_source, backend_args=backend_args) as archive:
+ frame_members = _sparse_frame_members(archive)
+ expected_clip_keys = _expected_clip_key_set(expected_clip_key)
+ if expected_clip_keys is not None:
+ mismatched_clip_keys = sorted(
+ {clip_key for clip_key, _ in frame_members.values() if clip_key not in expected_clip_keys}
+ )
+ if mismatched_clip_keys:
+ raise ValueError(
+ f"LiDAR rangemap clip keys {mismatched_clip_keys!r} do not match "
+ f"WebDataset key(s) {tuple(sorted(expected_clip_keys))!r}"
+ )
+ selected_ids = sorted(
+ source_frame_id
+ for source_frame_id in frame_members
+ if min_source_frame_id <= source_frame_id <= max_source_frame_id
+ )
+ if not selected_ids:
+ raise ValueError(
+ f"No LiDAR rangemap sweeps overlap source-frame interval [{min_source_frame_id}, {max_source_frame_id}]"
+ )
+ prepared_chunks: list[torch.Tensor] = []
+ for chunk_start in range(0, len(selected_ids), max_frames_per_chunk):
+ chunk_ids = selected_ids[chunk_start : chunk_start + max_frames_per_chunk]
+ metric_frames: list[np.ndarray] = []
+ intensity_frames: list[np.ndarray] = []
+ for source_frame_id in chunk_ids:
+ _, members = frame_members[source_frame_id]
+ if with_intensity:
+ metric, intensity = _rasterize_tar_frame_with_intensity(archive, source_frame_id, members)
+ intensity_frames.append(intensity)
+ else:
+ metric = _rasterize_tar_frame(archive, source_frame_id, members)
+ metric_frames.append(metric)
+ # Only prepend on the first chunk when requested, to avoid duplicating
+ # frame 0 at every chunk boundary.
+ chunk_prepend = prepend_first_frame and chunk_start == 0
+ prepared_chunks.append(
+ prepare_lidar_video_3ch(
+ np.stack(metric_frames),
+ np.stack(intensity_frames) if with_intensity else None,
+ prepend_first_frame=chunk_prepend,
+ range_projection=range_projection,
+ )
+ )
+ return torch.cat(prepared_chunks, dim=1), selected_ids
+
+
+def load_lidar_rangeview_directory(
+ rangeview_dir: str | Path,
+ *,
+ frame_prefixes: list[str] | None = None,
+ max_frames: int | None = None,
+) -> np.ndarray:
+ """Load loose ``*.lidar_{row,col,range}.npz`` frames as ``[T,128,3600]``."""
+ root = Path(rangeview_dir)
+ if not root.is_dir():
+ raise FileNotFoundError(f"Rangeview directory not found: {root}")
+ row_files = sorted(root.glob("*.lidar_row.npz"))
+ if not row_files:
+ raise ValueError(f"No lidar_row arrays found in {root}")
+ prefixes = [path.name.removesuffix(".lidar_row.npz") for path in row_files]
+ if frame_prefixes is not None:
+ missing = [prefix for prefix in frame_prefixes if prefix not in prefixes]
+ if missing:
+ raise FileNotFoundError(f"Missing rangeview prefixes under {root}: {missing}")
+ selected_prefixes = frame_prefixes
+ else:
+ selected_prefixes = prefixes if max_frames is None else prefixes[:max_frames]
+ frames: list[np.ndarray] = []
+ for prefix in selected_prefixes:
+ rows = _load_npz_array(root / f"{prefix}.lidar_row.npz")
+ cols = _load_npz_array(root / f"{prefix}.lidar_col.npz")
+ ranges = _load_npz_array(root / f"{prefix}.lidar_range.npz")
+ frames.append(rasterize_sparse_range_frame(rows, cols, ranges))
+ return np.stack(frames, axis=0)
+
+
+def load_lidar_sample(
+ sample_path: str | Path,
+ *,
+ num_frames: int | None = TOKENIZER_SAMPLE_FRAMES,
+ backend_args: dict[str, Any] | None = None,
+) -> np.ndarray:
+ """Load selected sweeps of one sample as ``[T,128,3600]``.
+
+ Accepts a local or remote LidarGEN clip tar, or a directory of loose sparse
+ frames. ``num_frames=None`` loads every available sweep.
+ """
+ if num_frames is not None and num_frames < 1:
+ raise ValueError(f"num_frames must be positive, got {num_frames}")
+ path_str = str(sample_path)
+ if not is_remote_uri(path_str) and Path(path_str).is_dir():
+ frames = load_lidar_rangeview_directory(path_str, max_frames=num_frames) # [T,128,3600]
+ else:
+ frame_indices = None if num_frames is None else list(range(num_frames))
+ frames = load_lidar_frames(path_str, frame_indices=frame_indices, backend_args=backend_args) # [T,128,3600]
+ if num_frames is not None and frames.shape[0] != num_frames:
+ raise ValueError(f"Requested {num_frames} frames but {path_str} yielded {frames.shape[0]}")
+ return frames
+
+
+def _find_local_lidar_samples(data_root: Path) -> list[str]:
+ for tar_dir in (data_root / "lidar", data_root):
+ tars = sorted(tar_dir.glob("*.tar"))
+ if tars:
+ return [str(path) for path in tars]
+ for frame_dir in (data_root / "rangeview", data_root):
+ if any(frame_dir.glob("*.lidar_row.npz")):
+ return [str(frame_dir)]
+ return []
+
+
+def _find_remote_lidar_samples(
+ data_root: str,
+ *,
+ backend_args: dict[str, Any] | None,
+) -> list[str]:
+ from cosmos_framework.utils.easy_io import easy_io
+
+ root = data_root.rstrip("/")
+ candidates = (f"{root}/lidar", root)
+ args = _backend_args_or_default(backend_args)
+ samples: list[str] = []
+ last_error: Exception | None = None
+ for prefix in candidates:
+ try:
+ names = list(
+ easy_io.list_dir_or_file(
+ prefix if prefix.endswith("/") else f"{prefix}/",
+ list_dir=False,
+ list_file=True,
+ suffix=".tar",
+ recursive=False,
+ backend_args=args,
+ )
+ )
+ except Exception as error:
+ last_error = error
+ continue
+ for name in sorted(names):
+ samples.append(name if name.startswith("s3://") else f"{prefix.rstrip('/')}/{name.lstrip('/')}")
+ if samples:
+ return samples
+ if last_error is not None:
+ raise RuntimeError(f"Could not list remote LiDAR samples under {data_root}") from last_error
+ return []
+
+
+def find_lidar_samples(
+ data_root: str | Path,
+ *,
+ backend_args: dict[str, Any] | None = None,
+) -> list[str]:
+ """List clip tars or loose-frame directories under a LiDAR data root.
+
+ Handles the LidarGEN layout (``/lidar/*.tar`` beside ``metadata/``),
+ a flat directory of tars, a directory of loose ``*.lidar_row.npz``, and the
+ matching S3 prefixes under ``s3://bucket0/lidar/data``.
+ """
+ root = str(data_root)
+ if is_remote_uri(root):
+ return _find_remote_lidar_samples(root, backend_args=backend_args)
+ path = Path(root)
+ if not path.is_dir():
+ return []
+ return _find_local_lidar_samples(path)
diff --git a/cosmos_framework/model/generator/tokenizers/lidar/range_projection.py b/cosmos_framework/model/generator/tokenizers/lidar/range_projection.py
new file mode 100644
index 000000000..cfeddf508
--- /dev/null
+++ b/cosmos_framework/model/generator/tokenizers/lidar/range_projection.py
@@ -0,0 +1,142 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+
+"""Shared definition of the LiDAR tokenizer's physical range-view contract."""
+
+from __future__ import annotations
+
+from collections.abc import Mapping
+from typing import Any
+
+import attrs
+
+
+@attrs.frozen
+class LidarRangeProjectionConfig:
+ """Settings that determine the contents and interpretation of a range view.
+
+ This object is deliberately independent of Hydra, model construction, and
+ storage. Training uses it directly and inference records the resolved object
+ so evaluation can reconstruct the exact deterministic projection contract.
+ """
+
+ # Geometry currently uses the Pandar128 beam-elevation calibration. Keep
+ # the sensor identity explicit so incompatible 128-row data fails early.
+ sensor: str = "pandar128"
+ native_height: int = 128
+ native_width: int = 3600
+ semantic_height: int = 128
+ semantic_width: int = 3600
+ model_width: int = 3600
+ model_width_transform: str = "circular_pad"
+ azimuth_start_degrees: float = 180.0
+ azimuth_end_degrees: float = -180.0
+ azimuth_endpoint: bool = False
+ return_selection: str = "nearest"
+ # Drive-Dreams metric span shared by every LiDAR tokenizer checkpoint. The
+ # span fixes the [-1, 1] affine, so it must match the span a checkpoint was
+ # trained with or metric decode is silently wrong.
+ min_range_m: float = 5.0
+ max_range_m: float = 100.0
+ intensity_encoding: str = "unit"
+ invalid_range_m: float = 0.0
+ validity_threshold: float = 0.5
+ coordinate_system: str = "x_forward_y_left_z_up"
+
+ def __attrs_post_init__(self) -> None:
+ if self.sensor != "pandar128":
+ raise ValueError(f"Unsupported LiDAR sensor {self.sensor!r}")
+ if self.native_height < 1 or self.native_width < 1:
+ raise ValueError(
+ f"Native range-view dimensions must be positive, got {(self.native_height, self.native_width)}"
+ )
+ if self.semantic_height != self.native_height:
+ raise ValueError(
+ "Vertical resampling is not supported; "
+ f"semantic_height must equal native_height, got {self.semantic_height} != {self.native_height}"
+ )
+ if self.semantic_width < 1 or self.native_width % self.semantic_width:
+ raise ValueError(
+ f"semantic_width must be a positive divisor of native_width {self.native_width}, "
+ f"got {self.semantic_width}"
+ )
+ if self.model_width_transform == "circular_pad":
+ if self.model_width < self.semantic_width:
+ raise ValueError(
+ f"model_width must be at least semantic_width, got {self.model_width} < {self.semantic_width}"
+ )
+ if (self.model_width - self.semantic_width) % 2:
+ raise ValueError(
+ "model_width - semantic_width must be even for symmetric circular padding, "
+ f"got {self.model_width} - {self.semantic_width}"
+ )
+ elif self.model_width_transform == "resize":
+ if self.model_width < 1:
+ raise ValueError(f"model_width must be positive, got {self.model_width}")
+ if self.semantic_width == self.native_width and self.model_width != self.semantic_width:
+ raise ValueError(
+ "Resizing the native azimuth canvas skips nearest-return pooling. "
+ "Set semantic_width to a divisor of native_width first "
+ f"(V0 transfer uses semantic_width=1200, model_width=1024); "
+ f"got semantic_width={self.semantic_width} and model_width={self.model_width}"
+ )
+ else:
+ raise ValueError(f"Unsupported model width transform {self.model_width_transform!r}")
+ if self.return_selection != "nearest":
+ raise ValueError(f"Unsupported return selection {self.return_selection!r}")
+ if self.max_range_m <= self.min_range_m:
+ raise ValueError(f"max_range_m must exceed min_range_m, got {self.min_range_m=} and {self.max_range_m=}")
+ if self.intensity_encoding != "unit":
+ raise ValueError(f"Unsupported intensity encoding {self.intensity_encoding!r}")
+ if self.invalid_range_m != 0.0:
+ raise ValueError(f"Only zero invalid range is supported, got {self.invalid_range_m}")
+ if not 0.0 < self.validity_threshold < 1.0:
+ raise ValueError(f"validity_threshold must lie in (0,1), got {self.validity_threshold}")
+ if self.coordinate_system != "x_forward_y_left_z_up":
+ raise ValueError(f"Unsupported coordinate system {self.coordinate_system!r}")
+
+ @property
+ def column_pool_factor(self) -> int:
+ """Number of adjacent native columns considered for each semantic ray."""
+ return self.native_width // self.semantic_width
+
+ def to_dict(self) -> dict[str, Any]:
+ """Serialize the resolved projection for a generated artifact manifest."""
+ return attrs.asdict(self)
+
+ @classmethod
+ def from_dict(cls, value: Mapping[str, Any]) -> LidarRangeProjectionConfig:
+ """Reconstruct and validate a serialized projection configuration."""
+ return cls(**dict(value))
+
+
+# Azimuth canvases used by transfer recipes. All inherit the class metric span;
+# they differ only in how many columns the tokenizer's network consumes.
+V1_TRANSFER_RANGE_PROJECTION = LidarRangeProjectionConfig(
+ semantic_width=1800,
+ model_width=1800,
+)
+V0_TRANSFER_RANGE_PROJECTION = LidarRangeProjectionConfig(
+ semantic_width=1200,
+ model_width=1024,
+ model_width_transform="resize",
+)
+# V1.2 keeps V1.1's 1800 semantic rays and widens the canvas to the next
+# multiple of its 16x compression. Padding rather than resizing is what keeps
+# the 1800 rays on their native angular grid, and wrapping the azimuth means the
+# eight added columns hold real neighbors instead of edge replicas.
+V1P2_TRANSFER_RANGE_PROJECTION = LidarRangeProjectionConfig(
+ semantic_width=1800,
+ model_width=1808,
+)
+
+
+def transfer_range_projection(*, tokenizer_v1: bool) -> LidarRangeProjectionConfig:
+ """Range-view contract for a transfer recipe's LiDAR tokenizer version.
+
+ The boolean reaches only the two canvases that also differ in channel
+ layout. V1.2 shares V1's three channels and differs from it in canvas width
+ alone, so a recipe on that tokenizer passes
+ ``V1P2_TRANSFER_RANGE_PROJECTION`` explicitly and still counts as V1 here.
+ """
+ return V1_TRANSFER_RANGE_PROJECTION if tokenizer_v1 else V0_TRANSFER_RANGE_PROJECTION
diff --git a/cosmos_framework/model/generator/tokenizers/lidar/range_projection_test.py b/cosmos_framework/model/generator/tokenizers/lidar/range_projection_test.py
new file mode 100644
index 000000000..91b9d1ef5
--- /dev/null
+++ b/cosmos_framework/model/generator/tokenizers/lidar/range_projection_test.py
@@ -0,0 +1,79 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+
+from __future__ import annotations
+
+import pytest
+
+from cosmos_framework.model.generator.tokenizers.lidar.range_projection import (
+ V0_TRANSFER_RANGE_PROJECTION,
+ V1_TRANSFER_RANGE_PROJECTION,
+ LidarRangeProjectionConfig,
+)
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_every_projection_shares_the_drive_dreams_metric_span() -> None:
+ """The span fixes the [-1, 1] affine, so a stray default silently breaks metric decode."""
+ assert (LidarRangeProjectionConfig().min_range_m, LidarRangeProjectionConfig().max_range_m) == (5.0, 100.0)
+ for projection in (V0_TRANSFER_RANGE_PROJECTION, V1_TRANSFER_RANGE_PROJECTION):
+ assert (projection.min_range_m, projection.max_range_m) == (5.0, 100.0)
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_range_projection_round_trip() -> None:
+ projection = LidarRangeProjectionConfig(
+ semantic_width=1800,
+ model_width=1808,
+ max_range_m=105.0,
+ validity_threshold=0.6,
+ )
+
+ restored = LidarRangeProjectionConfig.from_dict(projection.to_dict())
+
+ assert restored == projection
+ assert restored.column_pool_factor == 2
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_range_projection_rejects_inconsistent_circular_layout() -> None:
+ with pytest.raises(ValueError, match="must be even"):
+ LidarRangeProjectionConfig(semantic_width=1800, model_width=1801)
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+@pytest.mark.parametrize("validity_threshold", [0.0, 1.0])
+def test_range_projection_rejects_degenerate_validity_threshold(validity_threshold: float) -> None:
+ with pytest.raises(ValueError, match="validity_threshold"):
+ LidarRangeProjectionConfig(validity_threshold=validity_threshold)
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_range_projection_records_legacy_resize() -> None:
+ projection = LidarRangeProjectionConfig(
+ semantic_width=1200,
+ model_width=1024,
+ model_width_transform="resize",
+ min_range_m=5.0,
+ max_range_m=100.0,
+ )
+
+ assert projection.model_width == 1024
+ assert projection.column_pool_factor == 3
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_range_projection_rejects_resize_from_native_width() -> None:
+ with pytest.raises(ValueError, match="nearest-return pooling"):
+ LidarRangeProjectionConfig(
+ model_width=1024,
+ model_width_transform="resize",
+ min_range_m=5.0,
+ max_range_m=100.0,
+ )
diff --git a/cosmos_framework/model/generator/utils/data_and_condition.py b/cosmos_framework/model/generator/utils/data_and_condition.py
index 636b58436..ce9b47ae7 100644
--- a/cosmos_framework/model/generator/utils/data_and_condition.py
+++ b/cosmos_framework/model/generator/utils/data_and_condition.py
@@ -201,3 +201,41 @@ def build_dense_sound_schedule(
idx_sound = torch.tensor(sound_sample_indices, dtype=torch.long, device=timesteps.device) # [n_sound]
return timesteps[idx_sound], sigmas[idx_sound] # [n_sound,...], [n_sound,...]
+
+
+def select_target_image_sizes(
+ image_sizes: list[torch.Tensor],
+ num_vision_items_per_sample: list[int] | None,
+ batch_size: int,
+) -> list[torch.Tensor]:
+ """Pick one ``image_size`` per sample, the size of the generated (last) vision item.
+
+ Single-item batches carry one ``image_size`` per sample. Multi-item samples (transfer, SR) carry
+ one entry per vision item, flattened by the joint dataloader in item order. Resolution-dependent
+ settings such as the rectified-flow shift must follow the target item, not the conditioning
+ item, so this selects the last item of each sample.
+
+ Args:
+ image_sizes: flattened list of ``[4]`` or ``[1,4]`` tensors ``[target_h, target_w, orig_h, orig_w]``.
+ num_vision_items_per_sample: items per sample, or None for one item per sample.
+ batch_size: number of samples.
+
+ Returns:
+ list of ``batch_size`` tensors.
+ """
+ if num_vision_items_per_sample is None or len(image_sizes) == batch_size:
+ return list(image_sizes[:batch_size])
+ if len(num_vision_items_per_sample) != batch_size:
+ raise ValueError(
+ f"num_vision_items_per_sample has {len(num_vision_items_per_sample)} entries for batch_size {batch_size}"
+ )
+ if sum(num_vision_items_per_sample) != len(image_sizes):
+ raise ValueError(
+ f"image_size has {len(image_sizes)} entries but samples declare {sum(num_vision_items_per_sample)} items"
+ )
+ selected: list[torch.Tensor] = []
+ offset = 0
+ for num_items in num_vision_items_per_sample:
+ offset += num_items
+ selected.append(image_sizes[offset - 1])
+ return selected
diff --git a/cosmos_framework/model/generator/utils/data_and_condition_test.py b/cosmos_framework/model/generator/utils/data_and_condition_test.py
new file mode 100644
index 000000000..1aaa8cc08
--- /dev/null
+++ b/cosmos_framework/model/generator/utils/data_and_condition_test.py
@@ -0,0 +1,40 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+
+import pytest
+import torch
+
+from cosmos_framework.model.generator.utils.data_and_condition import select_target_image_sizes
+
+pytestmark = [pytest.mark.L0, pytest.mark.CPU]
+
+
+def _size(h: int, w: int) -> torch.Tensor: # returns [4]
+ return torch.tensor([float(h), float(w), float(h), float(w)])
+
+
+def test_single_item_batches_pass_through() -> None:
+ sizes = [_size(480, 832), _size(720, 1280)]
+ assert select_target_image_sizes(sizes, None, 2) == sizes
+ assert select_target_image_sizes(sizes, [1, 1], 2) == sizes
+
+
+def test_multi_item_samples_return_last_item_per_sample() -> None:
+ # Sample 0: SR pair (LR 240p, HR 480p). Sample 1: SR pair at 720p. Flattened in item order.
+ sizes = [_size(240, 416), _size(480, 832), _size(360, 640), _size(720, 1280)]
+ selected = select_target_image_sizes(sizes, [2, 2], 2)
+ assert [s.tolist() for s in selected] == [sizes[1].tolist(), sizes[3].tolist()]
+
+
+def test_mixed_item_counts_are_indexed_correctly() -> None:
+ sizes = [_size(256, 256), _size(240, 416), _size(480, 832)]
+ selected = select_target_image_sizes(sizes, [1, 2], 2)
+ assert [s.tolist() for s in selected] == [sizes[0].tolist(), sizes[2].tolist()]
+
+
+def test_inconsistent_counts_raise() -> None:
+ sizes = [_size(240, 416), _size(480, 832), _size(720, 1280)]
+ with pytest.raises(ValueError, match="declare"):
+ select_target_image_sizes(sizes, [2, 2], 2)
+ with pytest.raises(ValueError, match="batch_size"):
+ select_target_image_sizes(sizes, [2, 1, 1], 2)
diff --git a/cosmos_framework/model/generator/utils/kv_cache.py b/cosmos_framework/model/generator/utils/kv_cache.py
index 81a5aae44..3eac1af84 100644
--- a/cosmos_framework/model/generator/utils/kv_cache.py
+++ b/cosmos_framework/model/generator/utils/kv_cache.py
@@ -12,6 +12,8 @@
from __future__ import annotations
+import weakref
+from collections.abc import Sequence
from dataclasses import dataclass, field
from typing import Any
@@ -20,7 +22,7 @@
# Re-exported from memory.py for backward compatibility.
from cosmos_framework.model.generator.utils.memory import KVToStore, MemoryState, MemoryValue
-from cosmos_framework.data.generator.sequence_packing.runtime import get_num_real_samples
+from cosmos_framework.data.generator.sequence_packing.runtime import get_num_real_samples, to_device_nonblocking
from cosmos_framework.configs.base.defaults.replay_attention import TeacherForcingReplayPolicyConfig
from cosmos_framework.model.generator.utils.kv_storage_backend import (
BF16StorageBackend,
@@ -60,6 +62,86 @@ def zero_null_action_values(
return gen_v # [B,S,H,D]
+class KVBufferPool:
+ """Process-lifetime pool of reusable KV workspaces keyed by ``(name, slot)``.
+
+ AR inference creates fresh caches for every generation. Handing their large
+ workspaces (per-layer K/V rings and the static history buffers read by the
+ CUDA-graph path) out of one pool keeps the tensors, and therefore their
+ addresses, stable across generations. A stable address is what lets
+ ``torch._dynamo.mark_static_address`` tell CUDA-graph trees to read a buffer
+ in place instead of copying it into graph-owned inputs on every replay.
+
+ The marking uses ``guard=False``, the same treatment Dynamo gives module
+ parameters: a ``guard=True`` mark adds an object-identity guard, and since
+ all decoder layers share one compiled code object while each layer owns a
+ distinct buffer, that would force one recompile per layer and trip the
+ recompile limit. Without the guard, CUDA-graph trees still select the
+ recorded graph by matching static-input addresses (per layer) and the
+ tensor shape/dtype guards still cover reallocation. A buffer is only ever
+ reallocated when the requested shape, dtype or device changes, so a live
+ address never silently moves. Contents are never trusted across
+ acquisitions: callers rewrite what they read.
+ """
+
+ def __init__(self) -> None:
+ self._buffers: dict[tuple[str, int], torch.Tensor] = {}
+ # Live owner per key so two concurrently active caches can never share a slot silently.
+ self._owners: dict[tuple[str, int], weakref.ReferenceType[object]] = {}
+
+ def acquire(
+ self,
+ name: str,
+ slot: int,
+ shape: tuple[int, ...],
+ *,
+ dtype: torch.dtype,
+ device: torch.device,
+ static_address: bool = False,
+ owner: object | None = None,
+ ) -> torch.Tensor:
+ key = (name, int(slot))
+ current_owner_ref = self._owners.get(key)
+ current_owner = current_owner_ref() if current_owner_ref is not None else None
+ if owner is not None and current_owner is not None and current_owner is not owner:
+ raise RuntimeError(
+ f"KV buffer {name!r} slot {slot} is still owned by a live cache; concurrent AR generations on one "
+ "model need distinct pool slots (release the previous generation's caches first)"
+ )
+ buffer = self._buffers.get(key)
+ if buffer is None or tuple(buffer.shape) != tuple(shape) or buffer.dtype != dtype or buffer.device != device:
+ # A plain (non-inference) tensor stays writable in place from both
+ # ``torch.inference_mode`` and ``torch.no_grad`` generations.
+ with torch.inference_mode(mode=False):
+ # Zero-filled on purpose: the varlen FMHA reads each sequence's last 128-token KV tile in
+ # full and masks the overrun, so padding past ``cu_seqlens`` must hold finite values
+ # (0 x finite = 0, 0 x NaN poisons the row). One memset per allocation.
+ buffer = torch.zeros(tuple(shape), dtype=dtype, device=device)
+ if static_address:
+ torch._dynamo.mark_static_address(buffer, guard=False)
+ self._buffers[key] = buffer
+ if owner is not None:
+ self._owners[key] = weakref.ref(owner)
+ return buffer
+
+ def release_owner(self, owner: object) -> None:
+ """Drop ``owner``'s claim on its slots (buffers stay pooled for the next generation)."""
+ for key, ref in list(self._owners.items()):
+ holder = ref()
+ if holder is None or holder is owner:
+ del self._owners[key]
+
+ def discard(self, name: str, slot: int) -> None:
+ """Forget a buffer so its memory is freed once no cache references it."""
+ key = (name, int(slot))
+ self._buffers.pop(key, None)
+ self._owners.pop(key, None)
+
+ def clear(self) -> None:
+ self._buffers.clear()
+ self._owners.clear()
+
+
class KVCache:
"""Manages key/value caching for efficient sequential attention using a circular buffer.
@@ -428,7 +510,26 @@ def __init__(
cache_size: int | None = None,
backend: KVStorageBackend | None = None,
attention_sink_size: int = 0,
+ *,
+ buffer_pool: KVBufferPool | None = None,
+ pool_slot: int = 0,
+ preallocate_ring: bool = False,
) -> None:
+ """Args beyond ``KVCache``:
+ buffer_pool: Optional pool that owns this cache's large workspaces so
+ they are reused across generations (stable addresses).
+ pool_slot: Pool key for this cache (one per layer and CFG branch).
+ preallocate_ring: Store entries in place into one preallocated ring
+ tensor per K and V (no per-entry clone, no allocator churn).
+ Only used with a finite cache, the BF16 backend and uniform entry
+ shapes; otherwise the per-entry clone path is kept.
+ """
+ self._buffer_pool = buffer_pool
+ self._pool_slot = int(pool_slot)
+ self._preallocate_ring = preallocate_ring
+ self._ring_k: torch.Tensor | None = None
+ self._ring_v: torch.Tensor | None = None
+ self._ring_disabled = False
# initialize static KV cache used or AR inference use case.
# This buffer hosts materialization of the rolling window in chronological order.
self._static_k_buf: torch.Tensor | None = None
@@ -436,26 +537,188 @@ def __init__(
self._static_valid_frame_idx: int | None = None
self._static_real_len: int = 0
self._static_tokens_per_frame: int = 0
+ # The ``history_frame_indices`` request behind the current buffer contents (None = ring-native).
+ self._static_history_indices: tuple[int, ...] | None = None
+ # Composite layout: ``prefix_tokens`` = und + current-frame tokens ahead of the history region,
+ # and the identity of the und K/V last copied into the und region (None = not primed).
+ self._static_prefix_tokens: int = 0
+ self._static_batch_rows: int = 1
+ self._static_und_key: tuple | None = None
+ self._static_real_lens: tuple[int, ...] = ()
+ # Per-row episode restarts (batched Transfer): the global cache index at which row r last
+ # restarted, or None when no row ever did. A restarted row keeps storing into the shared
+ # global slots but reads a fresh-cache history plan expressed in its local frame count, with
+ # its first ``attention_sink_size`` local entries mirrored into the pinned sink slots of its
+ # own ring row (see ``row_history_slots``).
+ self._row_reset_frame: list[int] | None = None
super().__init__(cache_size=cache_size, backend=backend, attention_sink_size=attention_sink_size)
def reset(self) -> None:
- """Reset cache state and discard static inference workspaces."""
+ """Reset cache state, discard static inference workspaces and release pooled slots."""
super().reset()
self._static_k_buf = None
self._static_v_buf = None
self._static_valid_frame_idx = None
self._static_real_len = 0
self._static_tokens_per_frame = 0
+ self._static_history_indices = None
+ self._static_prefix_tokens = 0
+ self._static_batch_rows = 1
+ self._static_und_key = None
+ self._static_real_lens = ()
+ self._row_reset_frame = None
+ self._ring_k = None
+ self._ring_v = None
+ self._ring_disabled = False
+ if self._buffer_pool is not None:
+ self._buffer_pool.release_owner(self)
def store_kv(self, k: torch.Tensor, v: torch.Tensor, frame_idx: int) -> None:
"""Store K/V and invalidate static read workspaces."""
- super().store_kv(k, v, frame_idx)
+ if not self._store_kv_in_ring(k, v, frame_idx):
+ super().store_kv(k, v, frame_idx)
# A store can change the chronological history for the next AR frame,
# and can also overwrite a physical ring-buffer slot that the static
# buffer previously copied from. Mark the workspace stale.
# refresh happens once on the next static read.
self._static_valid_frame_idx = None
+ def _store_kv_in_ring(self, k: torch.Tensor, v: torch.Tensor, frame_idx: int) -> bool:
+ """Write ``k``/``v`` in place into the preallocated ring; ``False`` = use the clone path."""
+ if not self._preallocate_ring or self._ring_disabled:
+ return False
+ if self.cache_size >= MAX_CACHE_SIZE or not isinstance(self.backend, BF16StorageBackend):
+ return False
+ if k.shape != v.shape or k.dtype != v.dtype or k.device != v.device:
+ return False
+ ring_shape = (self.cache_size, *k.shape)
+ if self._ring_k is None:
+ if self._buffer_pool is not None:
+ self._ring_k = self._buffer_pool.acquire(
+ "ring_k", self._pool_slot, ring_shape, dtype=k.dtype, device=k.device, owner=self
+ )
+ self._ring_v = self._buffer_pool.acquire(
+ "ring_v", self._pool_slot, ring_shape, dtype=k.dtype, device=k.device, owner=self
+ )
+ else:
+ with torch.inference_mode(mode=False):
+ self._ring_k = torch.empty(ring_shape, dtype=k.dtype, device=k.device) # [slots,B,S,H,D]
+ self._ring_v = torch.empty(ring_shape, dtype=k.dtype, device=k.device) # [slots,B,S,H,D]
+ assert self._ring_k is not None and self._ring_v is not None
+ if tuple(self._ring_k.shape) != ring_shape or self._ring_k.dtype != k.dtype or self._ring_k.device != k.device:
+ # Variable-size entries: keep the already-stored views (they hold the
+ # ring alive) and fall back to per-entry clones for the rest of this
+ # cache's lifetime. Drop the ring from the pool so it is not pinned
+ # for later generations that will never use it either.
+ self._ring_disabled = True
+ if self._buffer_pool is not None:
+ self._buffer_pool.discard("ring_k", self._pool_slot)
+ self._buffer_pool.discard("ring_v", self._pool_slot)
+ return False
+ index = self._cache_index(int(frame_idx))
+ # In-place copy: detached from autograd and independent of the source
+ # storage (which may live in a CUDA-graph pool), same as the clone path.
+ self._ring_k[index].copy_(k.detach())
+ self._ring_v[index].copy_(v.detach())
+ self._mirror_reset_row_sinks(index, int(frame_idx))
+ k_entry = self._ring_k[index] # [B,S,H,D] view into the ring
+ v_entry = self._ring_v[index] # [B,S,H,D] view into the ring
+ self.k_cache[index] = k_entry
+ self.v_cache[index] = v_entry
+ self.backend.update_cached_kv_metadata(index, k_entry, v_entry)
+ return True
+
+ def transfer_history_frame_indices(
+ self,
+ frame_idx: int,
+ sink_entries: int,
+ max_recent_entries: int,
+ ) -> list[int]:
+ """Return the logical cache entries a Transfer forward at ``frame_idx`` may attend to.
+
+ Entry-level mirror of the token slicing in ``ARMemoryState.read_for_layer``:
+ the first ``sink_entries`` entries of the ring history stay pinned (complete
+ logical Transfer sink frames, one control + one RGB entry each) and only the
+ newest ``max_recent_entries`` of the remaining history are kept. Control seeds
+ pass ``2 * recent`` and target forwards ``2 * recent + 1`` so that ``R_t`` also
+ sees its aligned ``C_t``. Requires every cached entry to hold the same token
+ count (framewise Transfer), which the static AR path asserts when it copies.
+ """
+ if sink_entries < 0:
+ raise ValueError(f"sink_entries must be >= 0, got {sink_entries}")
+ if max_recent_entries < 0:
+ raise ValueError(f"max_recent_entries must be >= 0, got {max_recent_entries}")
+ history = self._history_frame_indices(int(frame_idx))
+ sink = history[:sink_entries]
+ rest = history[sink_entries:]
+ recent = rest[-max_recent_entries:] if max_recent_entries > 0 else []
+ return sink + recent
+
+ def reset_rows(self, rows: Sequence[int], at_frame_idx: int) -> None:
+ """Restart the episodes of ``rows`` at global cache index ``at_frame_idx`` (batched Transfer).
+
+ The row keeps storing into the shared global ring slots (all rows step in lockstep) but its
+ visible history becomes that of a fresh cache whose logical frame ``l`` is global
+ ``at_frame_idx + l``; its first ``attention_sink_size`` entries are additionally mirrored
+ into the pinned sink slots of its own ring row so the fresh-cache sink survives eviction.
+ Only tensor values change afterwards, so captured CUDA graphs keep replaying.
+ """
+ if not self._preallocate_ring or self._ring_disabled or self.cache_size is None:
+ raise RuntimeError("per-row resets require the in-place K/V ring (finite cache, BF16 storage)")
+ if self._ring_k is not None:
+ batch_rows = int(self._ring_k.shape[1])
+ for row in rows:
+ if row < 0 or row >= batch_rows:
+ raise IndexError(f"reset row {row} outside [0, {batch_rows})")
+ if self._row_reset_frame is None:
+ self._row_reset_frame = []
+ for row in rows:
+ while len(self._row_reset_frame) <= row:
+ self._row_reset_frame.append(0)
+ self._row_reset_frame[row] = int(at_frame_idx)
+ self._static_valid_frame_idx = None
+
+ def row_reset_frame(self, row: int) -> int:
+ """Global cache index at which ``row`` last restarted (0 = never)."""
+ if self._row_reset_frame is None or row >= len(self._row_reset_frame):
+ return 0
+ return self._row_reset_frame[row]
+
+ @property
+ def has_row_resets(self) -> bool:
+ return self._row_reset_frame is not None and any(self._row_reset_frame)
+
+ def _mirror_reset_row_sinks(self, slot: int, frame_idx: int) -> None:
+ """After storing global entry ``frame_idx`` at ``slot``, copy restarted rows' early entries into their sink slots."""
+ if self._row_reset_frame is None or self._ring_k is None or self._ring_v is None:
+ return
+ for row, reset_frame in enumerate(self._row_reset_frame):
+ if reset_frame <= 0:
+ continue
+ local = frame_idx - reset_frame
+ if 0 <= local < self.attention_sink_size and local != slot:
+ self._ring_k[local, row].copy_(self._ring_k[slot, row])
+ self._ring_v[local, row].copy_(self._ring_v[slot, row])
+
+ def row_history_slots(self, row: int, frame_idx: int, sink_entries: int, max_recent_entries: int) -> list[int]:
+ """Physical ring slots row ``row`` may attend to at global ``frame_idx`` (Transfer visibility).
+
+ Without a restart this is ``transfer_history_frame_indices`` mapped through ``_cache_index``.
+ After a restart the plan is that of a fresh cache at the row's local frame: its sink entries
+ live in the pinned sink slots (mirrored at store time), the rest at the global slots.
+ """
+ reset_frame = self.row_reset_frame(row)
+ local_frame = int(frame_idx) - reset_frame
+ if local_frame < 0:
+ raise ValueError(f"row {row} restarted at {reset_frame}, after frame {frame_idx}")
+ slots: list[int] = []
+ for logical_local in self.transfer_history_frame_indices(local_frame, sink_entries, max_recent_entries):
+ if reset_frame > 0 and logical_local < self.attention_sink_size:
+ slots.append(logical_local) # mirrored sink slot of this row
+ else:
+ slots.append(self._cache_index(reset_frame + logical_local))
+ return slots
+
def _ensure_static_history_buffer_allocated(
self,
max_tokens: int,
@@ -464,17 +727,29 @@ def _ensure_static_history_buffer_allocated(
head_dim: int,
device: torch.device | None,
dtype: torch.dtype | None,
+ prefix_tokens: int = 0,
+ batch_rows: int = 1,
) -> None:
- """Allocate the fixed-size history buffer used by Cuda Graph AR inference."""
+ """Allocate the fixed-size K/V buffer used by CUDA Graph AR inference.
+
+ ``prefix_tokens`` reserves ``[und | curr]`` room ahead of the ``max_tokens`` history
+ region so the compiled block reads one contiguous ``[und | curr | hist | pad]`` buffer.
+ ``batch_rows`` rows of that layout are laid end to end with a fixed row stride
+ ``prefix_tokens + max_tokens`` (one row per batched sample; ``B=1`` is today's buffer).
+ """
if device is None:
raise ValueError("device is required when allocating static gen KV buffers")
if dtype is None:
raise ValueError("dtype is required when allocating static gen KV buffers")
+ if batch_rows < 1:
+ raise ValueError(f"batch_rows must be >= 1, got {batch_rows}")
- expected_shape = (1, max_tokens, num_heads, head_dim)
+ expected_shape = (1, batch_rows * (prefix_tokens + max_tokens), num_heads, head_dim)
needs_alloc = (
self._static_k_buf is None
or self._static_v_buf is None
+ or self._static_prefix_tokens != prefix_tokens
+ or self._static_batch_rows != batch_rows
or tuple(self._static_k_buf.shape) != expected_shape
or tuple(self._static_v_buf.shape) != expected_shape
or self._static_k_buf.device != device
@@ -485,11 +760,39 @@ def _ensure_static_history_buffer_allocated(
if not needs_alloc:
return
- self._static_k_buf = torch.empty(expected_shape, device=device, dtype=dtype) # [1,S_max,H,D]
- self._static_v_buf = torch.empty(expected_shape, device=device, dtype=dtype) # [1,S_max,H,D]
+ if self._buffer_pool is not None:
+ # Pooled + static address: the compiled AR block reads these in place
+ # on every CUDA-graph replay instead of copying them into graph inputs.
+ self._static_k_buf = self._buffer_pool.acquire(
+ "static_k",
+ self._pool_slot,
+ expected_shape,
+ dtype=dtype,
+ device=device,
+ static_address=True,
+ owner=self,
+ ) # [1,S_max,H,D]
+ self._static_v_buf = self._buffer_pool.acquire(
+ "static_v",
+ self._pool_slot,
+ expected_shape,
+ dtype=dtype,
+ device=device,
+ static_address=True,
+ owner=self,
+ ) # [1,S_max,H,D]
+ else:
+ with torch.inference_mode(mode=False):
+ # Zero-filled: the attention kernel reads past the real prefix within the last KV tile.
+ self._static_k_buf = torch.zeros(expected_shape, device=device, dtype=dtype) # [1,S_max,H,D]
+ self._static_v_buf = torch.zeros(expected_shape, device=device, dtype=dtype) # [1,S_max,H,D]
self._static_valid_frame_idx = None
self._static_real_len = 0
self._static_tokens_per_frame = 0
+ self._static_history_indices = None
+ self._static_prefix_tokens = prefix_tokens
+ self._static_batch_rows = batch_rows
+ self._static_und_key = None
def _first_cached_history_k(self, frame_idx: int) -> torch.Tensor | None:
"""Return the first cached K tensor that will contribute to ``frame_idx``."""
@@ -506,28 +809,55 @@ def _first_cached_history_k(self, frame_idx: int) -> torch.Tensor | None:
return self.backend.decode(k_entry) # [B,S_frame,H,D]
return None
- def _rebuild_static_history(self, frame_idx: int, max_tokens: int, tokens_per_frame: int) -> int:
- """Rebuild the fixed history buffer from the list-backed circular cache."""
+ def _rebuild_static_history(
+ self,
+ frame_idx: int,
+ max_tokens: int,
+ tokens_per_frame: int,
+ history_frame_indices: list[int] | None = None,
+ row_slot_lists: list[list[int]] | None = None,
+ ) -> int:
+ """Rebuild the fixed history buffer from the list-backed circular cache.
+
+ ``history_frame_indices`` overrides the ring-native chronological history
+ with an explicit list of logical entries (Transfer sink + recent selection).
+ ``row_slot_lists`` gives every buffer row its own physical slot list instead
+ (batched Transfer with per-row episode restarts); rows may then differ in length.
+ """
assert self._static_k_buf is not None
assert self._static_v_buf is not None
+ if self._static_k_buf.is_cuda and torch.cuda.is_current_stream_capturing():
+ # The copies below read the ring slots of *this* frame; recorded into a graph they
+ # would replay the same slots for every later frame. Callers refresh the buffer
+ # before capture/replay (``prepare_for_coarse_cuda_graph_replay``).
+ raise RuntimeError("static AR history rebuild requested inside a CUDA graph capture")
current_idx = int(frame_idx)
- if current_idx <= 0:
- # Frame 0 has no gen history. The returned buffer may contain
+ if row_slot_lists is not None:
+ return self._rebuild_static_history_rows(current_idx, max_tokens, tokens_per_frame, row_slot_lists)
+ requested_indices = None if history_frame_indices is None else tuple(history_frame_indices)
+ logical_indices = (
+ list(history_frame_indices)
+ if history_frame_indices is not None
+ else (self._history_frame_indices(current_idx) if current_idx > 0 else [])
+ )
+ if not logical_indices:
+ # No visible gen history. The returned buffer may contain
# uninitialized tail data, but real_len=0 and the static AR path
# will not expose any history tokens to attention.
self._static_valid_frame_idx = current_idx
self._static_real_len = 0
self._static_tokens_per_frame = tokens_per_frame
+ self._static_history_indices = requested_indices
return 0
- # Chronological logical frames in history. This mirrors ``fetch_kv``
- # exactly, including the cache_size-1 history limit that leaves one
- # slot available for the current frame after it is stored.
+ # Chronological logical frames in history. Without an override this
+ # mirrors ``fetch_kv`` exactly, including the cache_size-1 history limit
+ # that leaves one slot available for the current frame after it is stored.
history_indices: list[int] = []
history_k_entries: list[object] = []
history_v_entries: list[object] = []
- for logical_idx in self._history_frame_indices(current_idx):
+ for logical_idx in logical_indices:
# Map logical frame indices to ring-buffer slots. At wraparound,
# this copies frames in logical order even though physical storage
# is no longer contiguous.
@@ -551,32 +881,142 @@ def _rebuild_static_history(self, frame_idx: int, max_tokens: int, tokens_per_fr
history_k_entries.append(k_entry)
history_v_entries.append(v_entry)
- # Decode/concatenate the whole chronological window at once. This
- # avoids launching two copies per cached frame while staging inputs for
- # a coarse graph replay.
- k_history, v_history = self.backend.decode_many(
- history_k_entries,
- history_v_entries,
- slots=history_indices,
- ) # [B,S_hist,H,D] each
dst_end = len(history_indices) * tokens_per_frame
- if k_history.shape[1] != dst_end or v_history.shape[1] != dst_end:
- raise AssertionError(
- f"Static AR cache requires {tokens_per_frame} tokens per frame: "
- f"expected total={dst_end}, got k={k_history.shape[1]}, v={v_history.shape[1]}"
- )
if dst_end > max_tokens:
raise AssertionError(f"Static AR cache overflow: trying to write {dst_end} tokens into {max_tokens}")
- # Copy only the real chronological prefix. The padded suffix is
- # deliberately left untouched and excluded by ``cu_seqlens_kv_t``.
- self._static_k_buf[:, :dst_end].copy_(k_history.detach()) # [1,S_hist,H,D]
- self._static_v_buf[:, :dst_end].copy_(v_history.detach()) # [1,S_hist,H,D]
+ # Row r's history region starts at ``r*R + prefix`` (R = row stride); only the real
+ # chronological prefix of it is written. The padded suffix is deliberately left
+ # untouched (zero-filled at allocation) and excluded by ``cu_seqlens_kv_t``.
+ rows = self._static_batch_rows
+ row_stride = self._static_prefix_tokens + max_tokens
+ hist_start = self._static_prefix_tokens
+ if self._ring_slot_runs_copyable(history_indices, history_k_entries, history_v_entries):
+ # Single pass straight from the ring ``[slots,B,S,H,D]``: each run of consecutive
+ # physical slots is one copy per row. Same bytes in the same order as concatenating
+ # the entries; at B=1 the source run is contiguous.
+ assert self._ring_k is not None and self._ring_v is not None
+ _, _, _, num_heads, head_dim = self._ring_k.shape
+ for row in range(rows):
+ offset = row * row_stride + hist_start
+ for slot_start, slot_end in self._coalesce_slot_runs(history_indices):
+ run_entries = slot_end - slot_start
+ run_tokens = run_entries * tokens_per_frame
+ self._static_k_buf[0, offset : offset + run_tokens].view(
+ run_entries, tokens_per_frame, num_heads, head_dim
+ ).copy_(self._ring_k[slot_start:slot_end, row]) # [n,S,H,D]
+ self._static_v_buf[0, offset : offset + run_tokens].view(
+ run_entries, tokens_per_frame, num_heads, head_dim
+ ).copy_(self._ring_v[slot_start:slot_end, row]) # [n,S,H,D]
+ offset += run_tokens
+ else:
+ # Generic path (FP8 entries, clone storage, ragged rings): decode the whole
+ # chronological window at once, then copy each row into its history region.
+ k_history, v_history = self.backend.decode_many(
+ history_k_entries,
+ history_v_entries,
+ slots=history_indices,
+ ) # [B,S_hist,H,D] each
+ if k_history.shape[1] != dst_end or v_history.shape[1] != dst_end:
+ raise AssertionError(
+ f"Static AR cache requires {tokens_per_frame} tokens per frame: "
+ f"expected total={dst_end}, got k={k_history.shape[1]}, v={v_history.shape[1]}"
+ )
+ if k_history.shape[0] != rows:
+ raise AssertionError(f"Static AR cache holds {rows} rows but the history has {k_history.shape[0]}")
+ for row in range(rows):
+ offset = row * row_stride + hist_start
+ self._static_k_buf[0, offset : offset + dst_end].copy_(k_history[row].detach()) # [S_hist,H,D]
+ self._static_v_buf[0, offset : offset + dst_end].copy_(v_history[row].detach()) # [S_hist,H,D]
self._static_valid_frame_idx = current_idx
self._static_real_len = dst_end
+ self._static_real_lens = (dst_end,) * rows
self._static_tokens_per_frame = tokens_per_frame
+ # Remember the *request* (None = ring-native) so a later call with a different
+ # request at the same frame rebuilds, including switching back to native.
+ self._static_history_indices = requested_indices
return dst_end
+ def _rebuild_static_history_rows(
+ self,
+ current_idx: int,
+ max_tokens: int,
+ tokens_per_frame: int,
+ row_slot_lists: list[list[int]],
+ ) -> int:
+ """Per-row rebuild: row r copies its own physical slots (single pass from the ring)."""
+ assert self._static_k_buf is not None and self._static_v_buf is not None
+ rows = self._static_batch_rows
+ if len(row_slot_lists) != rows:
+ raise AssertionError(f"{len(row_slot_lists)} slot lists for {rows} rows")
+ if self._ring_k is None or self._ring_v is None or self._ring_disabled or self._ring_k.shape[1] != rows:
+ raise RuntimeError("per-row static history requires the in-place K/V ring with one ring row per buffer row")
+ _, _, ring_tokens, num_heads, head_dim = self._ring_k.shape
+ if ring_tokens != tokens_per_frame:
+ raise AssertionError(f"ring entries hold {ring_tokens} tokens, expected {tokens_per_frame}")
+ row_stride = self._static_prefix_tokens + max_tokens
+ real_lens: list[int] = []
+ for row, slots in enumerate(row_slot_lists):
+ for slot in slots:
+ if self.k_cache[slot] is None or self.v_cache[slot] is None:
+ raise AssertionError(f"K/V cache slot {slot} is empty (row {row}, frame_idx={current_idx})")
+ dst_end = len(slots) * tokens_per_frame
+ if dst_end > max_tokens:
+ raise AssertionError(f"Static AR cache overflow: trying to write {dst_end} tokens into {max_tokens}")
+ offset = row * row_stride + self._static_prefix_tokens
+ for slot_start, slot_end in self._coalesce_slot_runs(list(slots)):
+ run_entries = slot_end - slot_start
+ run_tokens = run_entries * tokens_per_frame
+ self._static_k_buf[0, offset : offset + run_tokens].view(
+ run_entries, tokens_per_frame, num_heads, head_dim
+ ).copy_(self._ring_k[slot_start:slot_end, row]) # [n,S,H,D]
+ self._static_v_buf[0, offset : offset + run_tokens].view(
+ run_entries, tokens_per_frame, num_heads, head_dim
+ ).copy_(self._ring_v[slot_start:slot_end, row]) # [n,S,H,D]
+ offset += run_tokens
+ real_lens.append(dst_end)
+ self._static_valid_frame_idx = current_idx
+ self._static_real_len = real_lens[0] if real_lens else 0
+ self._static_real_lens = tuple(real_lens)
+ self._static_tokens_per_frame = tokens_per_frame
+ self._static_history_indices = tuple(tuple(slots) for slots in row_slot_lists)
+ return self._static_real_len
+
+ def _ring_slot_runs_copyable(
+ self,
+ slots: list[int],
+ k_entries: list[object],
+ v_entries: list[object],
+ ) -> bool:
+ """True when every entry is the ring view of its slot (BF16, in-place ring, one ring row per buffer row)."""
+ if (
+ self._ring_k is None
+ or self._ring_v is None
+ or self._ring_disabled
+ or self._ring_k.shape[1] != self._static_batch_rows
+ ):
+ return False
+ for slot, k_entry, v_entry in zip(slots, k_entries, v_entries, strict=True):
+ if not isinstance(k_entry, torch.Tensor) or not isinstance(v_entry, torch.Tensor):
+ return False
+ if (
+ k_entry.data_ptr() != self._ring_k[slot].data_ptr()
+ or v_entry.data_ptr() != self._ring_v[slot].data_ptr()
+ ):
+ return False
+ return True
+
+ @staticmethod
+ def _coalesce_slot_runs(slots: list[int]) -> list[tuple[int, int]]:
+ """Group consecutive physical slots into ``[start, end)`` runs, preserving order."""
+ runs: list[tuple[int, int]] = []
+ for slot in slots:
+ if runs and runs[-1][1] == slot:
+ runs[-1] = (runs[-1][0], slot + 1)
+ else:
+ runs.append((slot, slot + 1))
+ return runs
+
def fetch_kv_static(
self,
frame_idx: int,
@@ -587,22 +1027,44 @@ def fetch_kv_static(
head_dim: int,
device: torch.device | None,
dtype: torch.dtype | None,
+ history_frame_indices: list[int] | None = None,
+ prefix_tokens: int = 0,
+ und_kv: tuple[torch.Tensor, torch.Tensor] | None = None,
+ batch_rows: int = 1,
+ und_lens: tuple[int, ...] | None = None,
+ row_slot_lists: list[list[int]] | None = None,
) -> tuple[torch.Tensor, torch.Tensor, int]:
"""Fetch cached K/V history from a persistent fixed-size buffer.
Unlike ``fetch_kv_padded``, this does not ``cat`` and ``pad`` on every
- read. It rebuilds the persistent buffer only when the requested frame
- changes or the cache is invalidated by a store/reset.
+ read. It rebuilds the persistent buffer only when the requested frame or
+ the explicit ``history_frame_indices`` change, or the cache is invalidated
+ by a store/reset.
+
+ With ``prefix_tokens`` every row is ``[und | curr | hist | pad]`` with stride
+ ``R = prefix_tokens + max_tokens`` and the buffer is ``[1, batch_rows * R, H, D]``; the
+ history of row ``r`` lives at ``r*R + prefix_tokens``. ``und_kv`` (``[B,U,H,D]`` each,
+ ``U`` = padded und length) is copied once per distinct und tensor pair into the und
+ region of each row, right-aligned so that it ends at ``r*R + U`` (``und_lens`` gives the
+ real per-row lengths; default: every row full). The returned ``real_len`` counts history
+ tokens per row.
"""
- first_k = self._first_cached_history_k(frame_idx) # [B,S_frame,H,D] or None
- if first_k is not None:
- # Prefer the cached tensor metadata over caller-provided defaults.
- # This preserves CP/head-sharded cache shapes and also picks up the
- # actual device/dtype after the first frame is cached.
- num_heads = first_k.shape[2]
- head_dim = first_k.shape[3]
- device = first_k.device
- dtype = first_k.dtype
+ if self._ring_k is not None and not self._ring_disabled:
+ # The ring carries the cached shape/device/dtype without walking the history.
+ num_heads = self._ring_k.shape[3]
+ head_dim = self._ring_k.shape[4]
+ device = self._ring_k.device
+ dtype = self._ring_k.dtype
+ else:
+ first_k = self._first_cached_history_k(frame_idx) # [B,S_frame,H,D] or None
+ if first_k is not None:
+ # Prefer the cached tensor metadata over caller-provided defaults.
+ # This preserves CP/head-sharded cache shapes and also picks up the
+ # actual device/dtype after the first frame is cached.
+ num_heads = first_k.shape[2]
+ head_dim = first_k.shape[3]
+ device = first_k.device
+ dtype = first_k.dtype
self._ensure_static_history_buffer_allocated(
max_tokens,
@@ -610,9 +1072,52 @@ def fetch_kv_static(
head_dim=head_dim,
device=device,
dtype=dtype,
+ prefix_tokens=prefix_tokens,
+ batch_rows=batch_rows,
+ )
+ if und_kv is not None:
+ k_und, v_und = und_kv
+ und_capacity = int(k_und.shape[1])
+ lens = tuple(und_lens) if und_lens is not None else (und_capacity,) * batch_rows
+ und_key = (k_und.data_ptr(), v_und.data_ptr(), und_capacity, lens)
+ if self._static_und_key != und_key:
+ assert self._static_k_buf is not None and self._static_v_buf is not None
+ if self._static_k_buf.is_cuda and torch.cuda.is_current_stream_capturing():
+ raise RuntimeError("static AR und-region priming requested inside a CUDA graph capture")
+ if und_capacity > prefix_tokens:
+ raise AssertionError(
+ f"und length {und_capacity} exceeds the static prefix of {prefix_tokens} tokens"
+ )
+ if k_und.shape[0] != batch_rows or len(lens) != batch_rows:
+ raise AssertionError(
+ f"und K/V carry {k_und.shape[0]} rows / {len(lens)} lengths for {batch_rows} rows"
+ )
+ row_stride = prefix_tokens + max_tokens
+ for row, real_len in enumerate(lens):
+ if real_len < 0 or real_len > und_capacity:
+ raise AssertionError(f"und length {real_len} outside [0, {und_capacity}] for row {row}")
+ end = row * row_stride + und_capacity
+ # Right-aligned: every row's current frame starts at the same offset ``U``.
+ self._static_k_buf[0, end - real_len : end].copy_(k_und[row, :real_len]) # [u_r,H,D]
+ self._static_v_buf[0, end - real_len : end].copy_(v_und[row, :real_len]) # [u_r,H,D]
+ self._static_und_key = und_key
+ if row_slot_lists is not None:
+ requested_indices: tuple | None = tuple(tuple(slots) for slots in row_slot_lists)
+ else:
+ requested_indices = None if history_frame_indices is None else tuple(history_frame_indices)
+ needs_rebuild = (
+ self._static_valid_frame_idx != int(frame_idx)
+ or self._static_tokens_per_frame != tokens_per_frame
+ or self._static_history_indices != requested_indices
)
- if self._static_valid_frame_idx != int(frame_idx) or self._static_tokens_per_frame != tokens_per_frame:
- self._rebuild_static_history(frame_idx, max_tokens, tokens_per_frame)
+ if needs_rebuild:
+ self._rebuild_static_history(
+ frame_idx,
+ max_tokens,
+ tokens_per_frame,
+ history_frame_indices=history_frame_indices,
+ row_slot_lists=row_slot_lists,
+ )
assert self._static_k_buf is not None
assert self._static_v_buf is not None
@@ -698,6 +1203,10 @@ def __init__(
kv_cache_dtype: str | None = None,
kv_cache_kernel_impl: str = "triton",
attention_sink_size: int = 0,
+ *,
+ buffer_pool: KVBufferPool | None = None,
+ pool_slot: int = 0,
+ preallocate_ring: bool = False,
) -> None:
"""Initialize dual cache.
@@ -710,6 +1219,7 @@ def __init__(
"triton" is the default fused decode path; "torch" uses the
reference path. FP8 encode always uses the torch path.
attention_sink_size: Number of initial gen frames to pin in cache.
+ buffer_pool / pool_slot / preallocate_ring: See ``GenKVCache``.
"""
self.und_cache = UndKVCache()
if kv_cache_dtype is None:
@@ -722,6 +1232,9 @@ def __init__(
cache_size=gen_cache_size,
backend=backend,
attention_sink_size=attention_sink_size,
+ buffer_pool=buffer_pool,
+ pool_slot=pool_slot,
+ preallocate_ring=preallocate_ring,
)
def reset(self) -> None:
@@ -1611,6 +2124,26 @@ class ARMemoryValue(MemoryValue):
``cumulative_seqlen_KV`` kwarg, giving CUDA Graphs a single
capture that replays for every frame.
+ The static flavor hands the block one pooled buffer per K and V,
+ ``kv_k_static`` / ``kv_v_static`` ``[1, S_und + gen_len + S_hist_max, H_kv, D]``,
+ laid out ``[und | curr | hist | pad]``: the und K/V are copied in once per
+ generation, the block writes the current frame in place at
+ ``static_curr_offset`` (= ``S_und``) and the history is rebuilt outside the
+ compiled region at ``static_hist_offset`` (= ``S_und + gen_len``). This is
+ the same token order the block used to materialise with ``cat`` every
+ forward, so attention reads identical bytes without the 81 MB copy. For a
+ single sample ``und_k_cached`` / ``gen_k_buf_full`` stay populated as *views* of
+ that buffer (the und prefix and the ``[hist | pad]`` window of ``max_gen_tokens``),
+ so readers of the pre-composite layout see the same bytes; a hand-built value with
+ ``kv_k_static=None`` and those two set still takes the legacy ``cat`` path in
+ ``attention_AR_gen_only``. With ``batch_size > 1``
+ the buffer holds one such row per sample at stride ``static_row_stride`` (``[1, B*R, H, D]``),
+ the und of row ``r`` right-aligned to end at ``r*R + static_curr_offset``, and
+ ``cu_seqlens_q_t`` / ``cu_seqlens_kv_t`` are ``[2B]``: entry ``2r`` is row ``r`` and entry
+ ``2r+1`` a zero-query dummy sequence that owns the gap up to the next row (the varlen
+ kernel wants adjacent sequences and skips zero-query entries). ``B=1`` collapses to the
+ ``[2]`` tensors above.
+
All tensor sequence dimensions are fixed per-step. For Context
Parallelism the head dimension is ``H/cp`` (head-sharded); otherwise ``H``.
@@ -1671,6 +2204,11 @@ class ARMemoryValue(MemoryValue):
und_lens: tuple[int, ...] = ()
gen_k_buf_full: torch.Tensor | None = None
gen_v_buf_full: torch.Tensor | None = None
+ kv_k_static: torch.Tensor | None = None
+ kv_v_static: torch.Tensor | None = None
+ static_curr_offset: int = 0
+ static_hist_offset: int = 0
+ static_row_stride: int = 0
real_gen_cache_len_t: torch.Tensor | None = None
real_und_cache_len_t: torch.Tensor | None = None
cu_seqlens_q_t: torch.Tensor | None = None
@@ -1750,6 +2288,10 @@ class ARMemoryState(MemoryState):
positive ``transfer_history_sink_tokens`` this limits only the
recent suffix after the pinned prefix. With no sink tokens it keeps
the legacy behavior of limiting the complete history to a suffix.
+ Supported on the dynamic-shape and the CUDA-graph static-shape
+ flavors (the static flavor selects whole cache entries via
+ ``GenKVCache.transfer_history_frame_indices``); rejected with
+ post-saturation static compile.
"""
def requires_natten_metadata(self) -> bool:
@@ -1808,16 +2350,22 @@ def __init__(
"post_saturation_static_compile=True requires static_und_cache_max_len"
)
if coarse_cuda_graph:
- assert post_saturation_static_compile, "coarse_cuda_graph=True requires post-saturation static compile"
+ assert post_saturation_static_compile or for_cuda_graphs, (
+ "coarse_cuda_graph=True requires post-saturation static compile or the static-shape "
+ "(for_cuda_graphs=True) flavor"
+ )
if transfer_history_sink_tokens < 0:
raise ValueError(f"transfer_history_sink_tokens must be >= 0, got {transfer_history_sink_tokens}")
if transfer_history_max_tokens is not None:
if transfer_history_max_tokens < 0:
raise ValueError(f"transfer_history_max_tokens must be >= 0, got {transfer_history_max_tokens}")
- if for_cuda_graphs or post_saturation_static_compile:
- raise ValueError("transfer history limiting supports only dynamic-shape AR inference")
- if batched and (for_cuda_graphs or post_saturation_static_compile or coarse_cuda_graph):
- raise ValueError("Batched AR memory supports only eager dynamic-shape inference")
+ if post_saturation_static_compile:
+ raise ValueError(
+ "transfer history limiting supports dynamic-shape and CUDA-graph static-shape AR inference, "
+ "not post-saturation static compile"
+ )
+ if batched and post_saturation_static_compile:
+ raise ValueError("Batched AR memory does not support post-saturation static compile")
if kv_head_shard_size > 1:
assert not for_cuda_graphs, "local KV-head cache storage does not support CUDA graph static-cache mode"
assert num_kv_heads is not None, "local KV-head cache storage requires num_kv_heads"
@@ -1834,18 +2382,43 @@ def __init__(
self._cu_seqlens_kv_t: torch.Tensor | None = None
self._real_und_cache_len_t: torch.Tensor | None = None
self._max_seqlen_KV: int = 0
+ # Explicit logical history entries for the static Transfer path (None = ring-native history).
+ self._static_history_indices: tuple[int, ...] | None = None
+ # ``[und | curr]`` tokens ahead of the history region in the composite static buffer.
+ self._static_prefix_tokens: int = 0
+ # Per-row physical slot plans (batched Transfer / per-row restarts); None = shared plan.
+ self._static_row_slots: list[list[int]] | None = None
self._device: torch.device = torch.device("cpu")
self._dtype: torch.dtype = torch.float32
def init(self, hidden_states: dict, device: torch.device) -> None:
if self.batched:
- self._batch_size = int(get_num_real_samples(hidden_states))
- full_sample_ids = hidden_states["_full_only_sample_ids"][: hidden_states["_num_full_tokens"]] # [N_gen]
- causal_sample_ids = hidden_states["_causal_sample_ids"][: hidden_states["_num_causal_tokens"]] # [N_und]
- gen_counts = torch.bincount(full_sample_ids, minlength=self._batch_size) # [B]
- und_counts = torch.bincount(causal_sample_ids, minlength=self._batch_size) # [B]
- self._gen_lens = tuple(int(length) for length in gen_counts.tolist())
- self._current_und_lens = tuple(int(length) for length in und_counts.tolist())
+ # ``get_num_real_samples`` excludes the trailing CUDA-graph pad segment. AR packs give
+ # every sample one generation split of the same size, so the gen lengths follow from the
+ # total without touching device memory. Prompt lengths are read from the split offsets
+ # only while the pack carries text (frame 0, the dynamic path); gen-only frames see no
+ # sync, which keeps init() legal inside a capture.
+ num_samples = int(get_num_real_samples(hidden_states))
+ self._batch_size = num_samples
+ num_full = int(hidden_states["_num_full_tokens"])
+ if num_samples <= 0 or num_full <= 0 or num_full % num_samples:
+ raise ValueError(
+ f"Batched AR requires equal generation lengths: {num_full} tokens over {num_samples} rows"
+ )
+ self._gen_lens = (num_full // num_samples,) * num_samples
+ num_causal = int(hidden_states["_num_causal_tokens"])
+ if num_causal == 0:
+ self._current_und_lens = (0,) * num_samples
+ else:
+ causal_offsets = hidden_states["_causal_seq_offsets"][: num_samples + 1] # [B+1]
+ full_offsets = hidden_states["_full_only_seq_offsets"][: num_samples + 1] # [B+1]
+ if causal_offsets.is_cuda and torch.cuda.is_current_stream_capturing():
+ raise RuntimeError("batched AR prompt lengths cannot be read inside a CUDA graph capture")
+ self._current_und_lens = tuple(int(x) for x in torch.diff(causal_offsets).tolist())
+ if tuple(int(x) for x in torch.diff(full_offsets).tolist()) != self._gen_lens:
+ raise ValueError(
+ f"Batched AR requires equal generation lengths, got {torch.diff(full_offsets).tolist()}"
+ )
if any(length <= 0 for length in self._gen_lens):
raise ValueError(f"Every batched AR sample must contain generation tokens, got {self._gen_lens}")
if len(set(self._gen_lens)) != 1:
@@ -1899,34 +2472,126 @@ def init(self, hidden_states: dict, device: torch.device) -> None:
cache_size = self.dual_kv_cache[0].gen_cache.cache_size
self._max_gen_cache_tokens = (cache_size - 1) * T * S_super
self._tokens_per_frame = T * S_super
-
- # Real (non-padding) length of the cached gen history at this
- # frame. Same rolling-buffer math as KVCacheTrainMemoryState.
- # Held as an int32 tensor of shape ``[1]`` so its *value* can vary
- # across CUDA-graph replays while the *shape* stays static.
- real_len = min(self.frame_idx, cache_size - 1) * T * S_super
- self._real_gen_cache_len_t = torch.tensor([real_len], device=device, dtype=torch.int32)
-
self._device = device
# Match the dtype of the populated und cache (set by the eager
# frame-0 prefill). At AR loop entry the und cache is always
# initialized; for safety fall back to float32 if not.
und_k = self.dual_kv_cache[0].und_cache.k_und
self._dtype = und_k.dtype if und_k is not None else torch.float32
+ self._static_prefix_tokens = self.dual_kv_cache[0].und_cache.cached_len + self._gen_len
+ if self.coarse_cuda_graph and self._real_gen_cache_len_t is not None:
+ # Whole-forward capture: ``init()`` runs inside the recorded region. The
+ # per-frame offset values are owned by ``prepare_for_coarse_cuda_graph_replay``;
+ # a ``fill_`` recorded here would replay the capture-time constants over them.
+ return
+ self._refresh_static_offsets(self.frame_idx, device)
- # Pre-build varlen offsets *outside* the compiled region. Doing the
- # construction inside the captured graph forces Dynamo to specialize
- # on the contained Python ints (gen_len, S_und, ...) and emit
- # value-specific guards (e.g. ``memory_value.frame_idx == N``) — every
- # frame retraces and blows past the recompile limit. Building them
- # here keeps the captured region's view as plain ``[2]`` tensor inputs
- # whose values change per frame but whose shapes/addresses Dynamo
- # never inspects.
- s_und = self.dual_kv_cache[0].und_cache.cached_len
- real_total_kv_len = s_und + self._gen_len + real_len
- self._cu_seqlens_q_t = torch.tensor([0, self._gen_len], device=device, dtype=torch.int32)
- self._cu_seqlens_kv_t = torch.tensor([0, real_total_kv_len], device=device, dtype=torch.int32)
- self._max_seqlen_KV = s_und + self._gen_len + self._max_gen_cache_tokens
+ def _static_history_plan(self, frame_idx: int) -> tuple[tuple[int, ...] | None, int]:
+ """Return ``(explicit history entries or None, real history tokens)`` for ``frame_idx``.
+
+ Transfer limits select whole cache entries (pinned sink pairs plus the newest
+ ``max_tokens`` worth) so the static buffer copies exactly what the dynamic
+ branch slices; otherwise the ring-native chronological window applies.
+ """
+ cache_size = self.dual_kv_cache[0].gen_cache.cache_size
+ tokens_per_frame = self._tokens_per_frame
+ if self.transfer_history_max_tokens is None:
+ return None, min(int(frame_idx), cache_size - 1) * tokens_per_frame
+ if (
+ self.transfer_history_sink_tokens % tokens_per_frame != 0
+ or self.transfer_history_max_tokens % tokens_per_frame != 0
+ ):
+ raise ValueError(
+ "Transfer history limits must be whole cache entries on the static AR path: "
+ f"sink_tokens={self.transfer_history_sink_tokens}, "
+ f"max_tokens={self.transfer_history_max_tokens}, tokens_per_frame={tokens_per_frame}"
+ )
+ history_indices = self.dual_kv_cache[0].gen_cache.transfer_history_frame_indices(
+ frame_idx,
+ sink_entries=self.transfer_history_sink_tokens // tokens_per_frame,
+ max_recent_entries=self.transfer_history_max_tokens // tokens_per_frame,
+ )
+ return tuple(history_indices), len(history_indices) * tokens_per_frame
+
+ def _static_history_plan_rows(self, frame_idx: int) -> tuple[list[list[int]], list[int]] | None:
+ """Per-row physical slot plans for batched Transfer (``None`` when the shared plan applies).
+
+ Rows get their own plan when there is more than one row or any row restarted its episode
+ (``GenKVCache.reset_rows``); each row's plan is a fresh-cache plan in its local frame count.
+ """
+ gen_cache = self.dual_kv_cache[0].gen_cache
+ if self.transfer_history_max_tokens is None or (self._batch_size == 1 and not gen_cache.has_row_resets):
+ return None
+ tokens_per_frame = self._tokens_per_frame
+ sink_entries = self.transfer_history_sink_tokens // tokens_per_frame
+ max_recent_entries = self.transfer_history_max_tokens // tokens_per_frame
+ row_slots = [
+ gen_cache.row_history_slots(row, frame_idx, sink_entries, max_recent_entries)
+ for row in range(self._batch_size)
+ ]
+ return row_slots, [len(slots) * tokens_per_frame for slots in row_slots]
+
+ def _refresh_static_offsets(self, frame_idx: int, device: torch.device) -> None:
+ """(Re)compute the static-flavor history plan and the varlen offset tensors for ``frame_idx``.
+
+ The ``[1]`` / ``[2]`` int32 tensors are created once and afterwards updated in
+ place: their *addresses* are baked into CUDA graphs (per-block trees or a
+ forward-scope capture) and only their *values* change per frame. The update
+ must run *outside* any whole-forward capture: a recorded ``fill_`` replays its
+ capture-time constants, so the coarse flavor refreshes only from
+ ``prepare_for_coarse_cuda_graph_replay``.
+ """
+ self._static_history_indices, real_len = self._static_history_plan(frame_idx)
+ self._static_row_slots = None
+ real_lens = [real_len] * self._batch_size
+ row_plan = self._static_history_plan_rows(frame_idx)
+ if row_plan is not None:
+ self._static_row_slots, real_lens = row_plan
+ real_len = real_lens[0]
+ und_cache = self.dual_kv_cache[0].und_cache
+ s_und = und_cache.cached_len # padded und length U (= the single prompt length at B=1)
+ rows = self._batch_size
+ row_stride = s_und + self._gen_len + self._max_gen_cache_tokens
+ self._max_seqlen_KV = row_stride
+ shared_single_row = rows == 1 and self._static_row_slots is None
+ if shared_single_row:
+ real_total_kv_len = s_und + self._gen_len + real_len
+ cu_q_values = [0, self._gen_len]
+ cu_kv_values = [0, real_total_kv_len]
+ else:
+ # Row r: [pad | und_r right-aligned to r*R+U | curr at r*R+U | hist_r | pad]; entry 2r is the
+ # row, entry 2r+1 the zero-query dummy that owns the gap to the next row.
+ und_lens = und_cache.cached_lens if len(und_cache.cached_lens) == rows else (s_und,) * rows
+ cu_q_values, cu_kv_values = [0], []
+ for row in range(rows):
+ cu_kv_values += [
+ row * row_stride + s_und - und_lens[row],
+ row * row_stride + s_und + self._gen_len + real_lens[row],
+ ]
+ cu_q_values += [(row + 1) * self._gen_len] if row == rows - 1 else [(row + 1) * self._gen_len] * 2
+ if self._real_gen_cache_len_t is None or self._cu_seqlens_q_t is None or self._cu_seqlens_kv_t is None:
+ # Pre-build varlen offsets *outside* the compiled region. Doing the
+ # construction inside the captured graph forces Dynamo to specialize
+ # on the contained Python ints (gen_len, S_und, ...) and emit
+ # value-specific guards -- every frame retraces and blows past the
+ # recompile limit. Building them here keeps the captured region's
+ # view as plain tensor inputs whose values change per frame but whose
+ # shapes/addresses Dynamo never inspects.
+ self._real_gen_cache_len_t = to_device_nonblocking(
+ torch.tensor([real_len], dtype=torch.int32), device
+ ) # [1]
+ self._cu_seqlens_q_t = to_device_nonblocking(torch.tensor(cu_q_values, dtype=torch.int32), device) # [2B]
+ self._cu_seqlens_kv_t = to_device_nonblocking(torch.tensor(cu_kv_values, dtype=torch.int32), device) # [2B]
+ else:
+ self._real_gen_cache_len_t.fill_(real_len)
+ if shared_single_row:
+ self._cu_seqlens_q_t[1:].fill_(self._gen_len)
+ self._cu_seqlens_kv_t[1:].fill_(real_total_kv_len)
+ else:
+ # Same addresses, new values: an asynchronous copy from pinned host memory.
+ self._cu_seqlens_kv_t.copy_(
+ to_device_nonblocking(torch.tensor(cu_kv_values, dtype=torch.int32), device)
+ )
def read_for_layer(self, layer_idx: int) -> ARMemoryValue:
cache = self.dual_kv_cache[layer_idx]
@@ -2051,17 +2716,25 @@ def read_for_layer(self, layer_idx: int) -> ARMemoryValue:
post_saturation_static_compile=self.post_saturation_static_compile,
)
- # Static-shape branch: hand the layer the full preallocated gen
- # buffer + a scalar real-length tensor. Shapes are constant
- # across frames so a single CUDA-graph capture replays.
- assert und_k_cached is not None, (
+ # Static-shape branch: hand the layer one composite ``[und | curr | hist | pad]``
+ # buffer per K and V plus the varlen offset tensors. Shapes and addresses are
+ # constant across frames so a single CUDA-graph capture replays; the und region
+ # is primed once per generation and the history region rebuilt per frame here
+ # (per-block graphs) or in ``prepare_for_coarse_cuda_graph_replay`` (forward scope).
+ assert und_k_cached is not None and und_v_cached is not None, (
"ARMemoryState(for_cuda_graphs=True) requires the und cache to be "
"populated by frame-0 prefill before entering the AR loop"
)
assert self._real_gen_cache_len_t is not None
assert self._num_kv_heads is not None
assert self._head_dim is not None
- gen_k_buf, gen_v_buf, _ = cache.gen_cache.fetch_kv_static(
+ s_und = cache.und_cache.cached_len
+ assert s_und + self._gen_len == self._static_prefix_tokens, (
+ f"static prefix mismatch: und={s_und} gen={self._gen_len} prefix={self._static_prefix_tokens}"
+ )
+ rows = self._batch_size
+ und_lens = cache.und_cache.cached_lens if len(cache.und_cache.cached_lens) == rows else None
+ kv_k_static, kv_v_static, _ = cache.gen_cache.fetch_kv_static(
self.frame_idx,
self._max_gen_cache_tokens,
self._tokens_per_frame,
@@ -2069,18 +2742,43 @@ def read_for_layer(self, layer_idx: int) -> ARMemoryValue:
head_dim=self._head_dim,
device=self._device,
dtype=self._dtype,
+ history_frame_indices=(
+ None if self._static_history_indices is None else list(self._static_history_indices)
+ ),
+ prefix_tokens=self._static_prefix_tokens,
+ und_kv=(und_k_cached, und_v_cached),
+ batch_rows=rows,
+ und_lens=und_lens,
+ row_slot_lists=self._static_row_slots,
)
- torch._dynamo.mark_static(gen_k_buf, 1)
- torch._dynamo.mark_static(gen_v_buf, 1)
+ torch._dynamo.mark_static(kv_k_static, 1)
+ torch._dynamo.mark_static(kv_v_static, 1)
+ # Legacy single-sample views of the composite buffer: the und prefix and the
+ # ``[hist | pad]`` window. No copies; the block itself reads ``kv_k_static``.
+ # Batched rows keep them ``None`` (per-row right-aligned und, strided rows).
+ hist_start = self._static_prefix_tokens
+ hist_end = hist_start + self._max_gen_cache_tokens
+ legacy_und_k = kv_k_static[:, :s_und] if rows == 1 and s_und > 0 else None # [1,S_und,H_kv,D]
+ legacy_und_v = kv_v_static[:, :s_und] if rows == 1 and s_und > 0 else None # [1,S_und,H_kv,D]
+ legacy_gen_k = kv_k_static[:, hist_start:hist_end] if rows == 1 else None # [1,max_gen_tokens,H_kv,D]
+ legacy_gen_v = kv_v_static[:, hist_start:hist_end] if rows == 1 else None # [1,max_gen_tokens,H_kv,D]
return ARMemoryValue(
- und_k_cached=und_k_cached,
- und_v_cached=und_v_cached,
+ und_k_cached=legacy_und_k,
+ und_v_cached=legacy_und_v,
gen_k_hist=None,
gen_v_hist=None,
frame_idx=self.frame_idx,
gen_len=self._gen_len,
- gen_k_buf_full=gen_k_buf,
- gen_v_buf_full=gen_v_buf,
+ batch_size=rows,
+ gen_lens=self._gen_lens,
+ und_lens=(cache.und_cache.cached_lens if self.batched else ()),
+ gen_k_buf_full=legacy_gen_k,
+ gen_v_buf_full=legacy_gen_v,
+ kv_k_static=kv_k_static,
+ kv_v_static=kv_v_static,
+ static_curr_offset=s_und,
+ static_hist_offset=self._static_prefix_tokens,
+ static_row_stride=(self._max_seqlen_KV if rows > 1 else 0),
real_gen_cache_len_t=self._real_gen_cache_len_t,
cu_seqlens_q_t=self._cu_seqlens_q_t,
cu_seqlens_kv_t=self._cu_seqlens_kv_t,
@@ -2138,12 +2836,14 @@ def write_for_layer(self, layer_idx: int, kv_to_store: KVToStore) -> None:
)
def prepare_for_coarse_cuda_graph_replay(self, frame_idx: int) -> None:
- """Refresh fixed-address history buffers before replaying a coarse graph."""
+ """Refresh fixed-address history buffers (and, on the static flavor, the offset tensors) before replay."""
if not self.coarse_cuda_graph:
raise RuntimeError("prepare_for_coarse_cuda_graph_replay requires coarse_cuda_graph=True")
assert self._num_kv_heads is not None
assert self._head_dim is not None
self.frame_idx = frame_idx
+ if self.for_cuda_graphs:
+ self._refresh_static_offsets(frame_idx, self._device)
for cache in self.dual_kv_cache:
cache.gen_cache.fetch_kv_static(
frame_idx,
@@ -2153,6 +2853,18 @@ def prepare_for_coarse_cuda_graph_replay(self, frame_idx: int) -> None:
head_dim=self._head_dim,
device=self._device,
dtype=self._dtype,
+ history_frame_indices=(
+ None if self._static_history_indices is None else list(self._static_history_indices)
+ ),
+ prefix_tokens=self._static_prefix_tokens if self.for_cuda_graphs else 0,
+ und_kv=cache.und_cache.get() if self.for_cuda_graphs else None,
+ batch_rows=self._batch_size if self.for_cuda_graphs else 1,
+ und_lens=(
+ cache.und_cache.cached_lens
+ if self.for_cuda_graphs and len(cache.und_cache.cached_lens) == self._batch_size
+ else None
+ ),
+ row_slot_lists=self._static_row_slots if self.for_cuda_graphs else None,
)
def commit_staged_gen_cache(self, frame_idx: int) -> None:
@@ -2178,6 +2890,7 @@ def is_gen_only(self) -> bool:
"KVCache",
"UndKVCache",
"GenKVCache",
+ "KVBufferPool",
"DualKVCache",
"KVCacheTrainMemoryState",
"KVTrainMemoryValue",
diff --git a/cosmos_framework/model/generator/utils/kv_cache_test.py b/cosmos_framework/model/generator/utils/kv_cache_test.py
index dc8ba5eb9..1ed5b9705 100644
--- a/cosmos_framework/model/generator/utils/kv_cache_test.py
+++ b/cosmos_framework/model/generator/utils/kv_cache_test.py
@@ -3,6 +3,8 @@
"""Tests for KV cache and non-CP AR inference logic."""
+from unittest.mock import patch
+
import pytest
import torch
@@ -27,6 +29,7 @@
FlexARMemoryState,
FlexARMemoryValue,
GenKVCache,
+ KVBufferPool,
KVCache,
KVCacheTrainMemoryState,
TeacherForcingMemoryState,
@@ -1168,10 +1171,13 @@ def test_batched_ar_memory_state_tracks_unequal_prompt_lengths_per_sample() -> N
batch_size, gen_len, max_und_len, num_heads, head_dim = 2, 3, 4, 1, 2
cache = DualKVCache(gen_cache_size=4)
hidden_states = {
- "sample_offsets": torch.tensor([0, 5, 12]), # [B+1]
+ "sample_offsets": torch.tensor([0, 5, 12, 15]), # [B+1] plus a trailing CUDA-graph pad segment
+ "_has_pad_segment": True,
"_full_only_sample_ids": torch.tensor([0, 0, 0, 1, 1, 1]), # [B*S_gen]
+ "_full_only_seq_offsets": torch.tensor([0, 3, 6, 8]), # [B+1] plus the pad segment
"_num_full_tokens": batch_size * gen_len,
"_causal_sample_ids": torch.tensor([0, 0, 1, 1, 1, 1]), # [S_und_total]
+ "_causal_seq_offsets": torch.tensor([0, 2, 6, 7]), # [B+1] plus the pad segment
"_num_causal_tokens": 6,
}
state = ARMemoryState([cache], frame_idx=0, batched=True)
@@ -1224,8 +1230,10 @@ def test_batched_ar_memory_state_rejects_unequal_generation_lengths() -> None:
hidden_states = {
"sample_offsets": torch.tensor([0, 4, 9]), # [B+1]
"_full_only_sample_ids": torch.tensor([0, 0, 1, 1, 1]), # [S_gen_total]
+ "_full_only_seq_offsets": torch.tensor([0, 2, 5]), # generation splits of 2 and 3 tokens
"_num_full_tokens": 5,
"_causal_sample_ids": torch.tensor([0, 1]), # [S_und_total]
+ "_causal_seq_offsets": torch.tensor([0, 1, 2]),
"_num_causal_tokens": 2,
}
state = ARMemoryState([cache], frame_idx=0, batched=True)
@@ -2013,8 +2021,8 @@ def test_ar_memory_state_static_shape_init_and_read():
tensors after ``init`` (constant shape across frames is the prerequisite
for a single CUDA-graph capture).
3. Return an ``ARMemoryValue(for_cuda_graphs=True)`` with constant-shape
- ``gen_k_buf_full`` / ``gen_v_buf_full`` (``[1, max_gen_cache_tokens, H, D]``)
- across every frame.
+ ``kv_k_static`` / ``kv_v_static`` (``[1, S_und + gen_len + max_gen_cache_tokens, H, D]``,
+ laid out ``[und | curr | hist | pad]``) across every frame.
4. Carry the und K/V cached at frame 0 unchanged into frames 1+.
5. Match the dynamic-shape flavor's ``frame_idx``, ``gen_len``, and
``und_k/v_cached`` fields exactly (they're shared by both branches).
@@ -2088,17 +2096,24 @@ def test_ar_memory_state_static_shape_init_and_read():
assert isinstance(mv_static, ARMemoryValue)
assert mv_static.for_cuda_graphs is True
assert mv_static.gen_k_hist is None and mv_static.gen_v_hist is None
+ assert mv_static.kv_k_static is not None and mv_static.kv_v_static is not None
assert mv_static.gen_k_buf_full is not None and mv_static.gen_v_buf_full is not None
assert mv_static.gen_k_buf_full.shape == (B, expected_max_gen_tokens, H, D)
- assert mv_static.gen_v_buf_full.shape == (B, expected_max_gen_tokens, H, D)
+ s_und = und_k.shape[1]
+ assert mv_static.static_curr_offset == s_und and mv_static.static_hist_offset == s_und + S_gen
+ assert mv_static.kv_k_static.shape == (B, s_und + S_gen + expected_max_gen_tokens, H, D)
+ assert mv_static.kv_v_static.shape == (B, s_und + S_gen + expected_max_gen_tokens, H, D)
+ assert mv_static.max_seqlen_KV == mv_static.kv_k_static.shape[1]
assert mv_static.real_gen_cache_len_t is not None
assert mv_static.real_gen_cache_len_t.shape == (1,)
assert mv_static.real_gen_cache_len_t.item() == 1 * S_gen # 1 prior frame
- # und K/V cached at frame 0 carried unchanged into frame 1.
+ # und K/V cached at frame 0 carried unchanged into frame 1 (primed into the und region;
+ # ``und_k_cached`` is a view of that region).
assert mv_static.und_k_cached is not None and mv_static.und_v_cached is not None
assert torch.equal(mv_static.und_k_cached, und_k)
- assert torch.equal(mv_static.und_v_cached, und_v)
+ assert torch.equal(mv_static.kv_k_static[:, :s_und], und_k)
+ assert torch.equal(mv_static.kv_v_static[:, :s_und], und_v)
# Shared fields match dynamic flavor exactly.
mv_dyn = state_dyn.read_for_layer(0)
@@ -2107,13 +2122,16 @@ def test_ar_memory_state_static_shape_init_and_read():
assert mv_static.gen_len == mv_dyn.gen_len == S_gen
assert mv_dyn.und_k_cached is not None
torch.testing.assert_close(mv_static.und_k_cached, mv_dyn.und_k_cached)
- torch.testing.assert_close(mv_static.und_v_cached, mv_dyn.und_v_cached)
+ torch.testing.assert_close(mv_static.kv_k_static[:, :s_und], mv_dyn.und_k_cached)
+ torch.testing.assert_close(mv_static.kv_v_static[:, :s_und], mv_dyn.und_v_cached)
- # Static buffer's real prefix must match dynamic gen_k_hist exactly.
+ # Static buffer's history region must match dynamic gen_k_hist exactly.
real_len = mv_static.real_gen_cache_len_t.item()
+ hist = mv_static.static_hist_offset
assert mv_dyn.gen_k_hist is not None and mv_dyn.gen_v_hist is not None
torch.testing.assert_close(mv_static.gen_k_buf_full[:, :real_len], mv_dyn.gen_k_hist)
- torch.testing.assert_close(mv_static.gen_v_buf_full[:, :real_len], mv_dyn.gen_v_hist)
+ torch.testing.assert_close(mv_static.kv_k_static[:, hist : hist + real_len], mv_dyn.gen_k_hist)
+ torch.testing.assert_close(mv_static.kv_v_static[:, hist : hist + real_len], mv_dyn.gen_v_hist)
# The padded tail is intentionally unspecified: static AR attention uses
# ``cu_seqlens_kv_t`` to exclude it from the varlen kernel.
@@ -2121,7 +2139,7 @@ def test_ar_memory_state_static_shape_init_and_read():
@pytest.mark.L0
def test_ar_memory_state_static_shape_constant_across_frames():
"""The whole point of the static-shape flavor is shape constancy across
- frames: ``cu_seqlens_*`` and ``gen_k_buf_full`` shapes must not change as
+ frames: ``cu_seqlens_*``, ``kv_k_static`` and ``gen_k_buf_full`` shapes must not change as
``frame_idx`` advances; only the *values* in ``cu_seqlens_kv_t`` and
``real_gen_cache_len_t`` should update.
@@ -2155,8 +2173,10 @@ def test_ar_memory_state_static_shape_constant_across_frames():
)
state.init({"_num_full_tokens": S_gen}, torch.device("cpu"))
mv = state.read_for_layer(0)
+ assert mv.kv_k_static is not None
assert mv.gen_k_buf_full is not None
shapes = {
+ "kv_k_static": tuple(mv.kv_k_static.shape),
"gen_k_buf_full": tuple(mv.gen_k_buf_full.shape),
"cu_seqlens_q_t": tuple(state._cu_seqlens_q_t.shape), # type: ignore[union-attr]
"cu_seqlens_kv_t": tuple(state._cu_seqlens_kv_t.shape), # type: ignore[union-attr]
@@ -2288,3 +2308,955 @@ def test_batched_ar_counts_real_samples_in_runtime_pack(batch_size: int, with_te
assert state._batch_size == batch_size
assert state._gen_lens == (4,) * batch_size
assert state._current_und_lens == tuple(text_lengths)
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_gen_kv_cache_transfer_history_frame_indices_pins_sink_and_keeps_recent_entries() -> None:
+ """Entry-level Transfer history: logical sink pairs stay, then the newest ``max_recent`` entries."""
+ tokens_per_entry = 2
+ cache = GenKVCache(cache_size=6, attention_sink_size=2) # W=3, S=1 -> 2W slots, 2S pinned
+ for entry_idx in range(9):
+ entry = torch.zeros(1, tokens_per_entry, 1, 1) # [B,S_entry,H,D]
+ cache.store_kv(entry, entry, frame_idx=entry_idx)
+
+ # Ring-native history at idx 8: pinned [0, 1] + rolling tail [5, 6, 7].
+ assert cache._history_frame_indices(8) == [0, 1, 5, 6, 7]
+ # Control seed (even idx) sees 2*recent = 2 recent entries: drop the orphan RGB entry 5.
+ assert cache.transfer_history_frame_indices(8, sink_entries=2, max_recent_entries=2) == [0, 1, 6, 7]
+ # Target denoise / RGB seed (odd idx) sees 2*recent+1 = 3 recent entries.
+ assert cache.transfer_history_frame_indices(9, sink_entries=2, max_recent_entries=3) == [0, 1, 6, 7, 8]
+ # max_recent_entries == 0 keeps only the sink.
+ assert cache.transfer_history_frame_indices(9, sink_entries=2, max_recent_entries=0) == [0, 1]
+ # Before saturation nothing is dropped.
+ assert cache.transfer_history_frame_indices(2, sink_entries=2, max_recent_entries=2) == [0, 1]
+ assert cache.transfer_history_frame_indices(3, sink_entries=2, max_recent_entries=3) == [0, 1, 2]
+ # Frame 0 has no history.
+ assert cache.transfer_history_frame_indices(0, sink_entries=2, max_recent_entries=2) == []
+
+ no_sink = GenKVCache(cache_size=4)
+ for entry_idx in range(7):
+ entry = torch.zeros(1, tokens_per_entry, 1, 1) # [B,S_entry,H,D]
+ no_sink.store_kv(entry, entry, frame_idx=entry_idx)
+ # Legacy total-suffix behaviour when there is no sink.
+ assert no_sink.transfer_history_frame_indices(7, sink_entries=0, max_recent_entries=2) == [5, 6]
+ assert no_sink.transfer_history_frame_indices(7, sink_entries=0, max_recent_entries=0) == []
+
+ with pytest.raises(ValueError, match="sink_entries must be >= 0"):
+ cache.transfer_history_frame_indices(8, sink_entries=-1, max_recent_entries=2)
+ with pytest.raises(ValueError, match="max_recent_entries must be >= 0"):
+ cache.transfer_history_frame_indices(8, sink_entries=2, max_recent_entries=-1)
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_gen_kv_cache_static_fetch_honours_explicit_history_indices_and_rebuilds_on_change() -> None:
+ """Explicit logical indices select the copied entries; identical indices reuse the buffer."""
+ batch_size, tokens_per_entry, num_heads, head_dim = 1, 2, 1, 1
+ cache = GenKVCache(cache_size=8)
+ entries: list[tuple[torch.Tensor, torch.Tensor]] = []
+ for entry_idx in range(6):
+ key = torch.full((batch_size, tokens_per_entry, num_heads, head_dim), float(entry_idx)) # [B,S_entry,H,D]
+ value = key + 100.0 # [B,S_entry,H,D]
+ entries.append((key, value))
+ cache.store_kv(key, value, frame_idx=entry_idx)
+ max_tokens = 7 * tokens_per_entry
+ kwargs = dict(num_heads=num_heads, head_dim=head_dim, device=torch.device("cpu"), dtype=torch.float32)
+
+ k_buf, v_buf, real_len = cache.fetch_kv_static(
+ 6, max_tokens, tokens_per_entry, history_frame_indices=[0, 1, 4, 5], **kwargs
+ )
+ assert real_len == 4 * tokens_per_entry
+ assert k_buf.shape == (1, max_tokens, num_heads, head_dim)
+ torch.testing.assert_close(k_buf[:, :real_len], torch.cat([entries[i][0] for i in (0, 1, 4, 5)], dim=1))
+ torch.testing.assert_close(v_buf[:, :real_len], torch.cat([entries[i][1] for i in (0, 1, 4, 5)], dim=1))
+ assert cache._static_history_indices == (0, 1, 4, 5)
+
+ with patch.object(GenKVCache, "_rebuild_static_history", wraps=cache._rebuild_static_history) as rebuild:
+ cache.fetch_kv_static(6, max_tokens, tokens_per_entry, history_frame_indices=[0, 1, 4, 5], **kwargs)
+ assert rebuild.call_count == 0 # same frame, same indices: no rebuild
+ k_buf2, _, real_len2 = cache.fetch_kv_static(
+ 6, max_tokens, tokens_per_entry, history_frame_indices=[0, 1, 3, 4, 5], **kwargs
+ )
+ assert rebuild.call_count == 1 # same frame, different indices: rebuild
+ assert real_len2 == 5 * tokens_per_entry
+ assert k_buf2.data_ptr() == k_buf.data_ptr() # persistent buffer, same address
+ torch.testing.assert_close(k_buf2[:, :real_len2], torch.cat([entries[i][0] for i in (0, 1, 3, 4, 5)], dim=1))
+
+ _, _, empty_len = cache.fetch_kv_static(6, max_tokens, tokens_per_entry, history_frame_indices=[], **kwargs)
+ assert empty_len == 0
+ assert cache._static_history_indices == ()
+
+ # Switching back to the ring-native history (no explicit indices) must rebuild too.
+ _, _, native_len = cache.fetch_kv_static(6, max_tokens, tokens_per_entry, **kwargs)
+ assert native_len == 6 * tokens_per_entry
+ assert cache._static_history_indices is None
+
+ cache.reset()
+ assert cache._static_history_indices is None
+
+
+def _assert_transfer_static_matches_dynamic(
+ window: int, sink: int, tokens_per_entry: int = 2, **cache_kwargs: object
+) -> None:
+ """Static Transfer buffer real prefix == dynamic Transfer history for both forward kinds, every frame."""
+ batch_size, num_heads, head_dim, s_und = 1, 1, 1, 3
+ cache = DualKVCache(gen_cache_size=2 * window, attention_sink_size=2 * sink, **cache_kwargs) # type: ignore[arg-type]
+ cache.und_cache.store(
+ torch.randn(batch_size, s_und, num_heads, head_dim), # [B,S_und,H,D]
+ torch.randn(batch_size, s_und, num_heads, head_dim), # [B,S_und,H,D]
+ )
+ recent = window - sink - 1
+ sink_tokens = 2 * sink * tokens_per_entry
+ kinds = {
+ "control": 2 * recent * tokens_per_entry,
+ "target": (2 * recent + 1) * tokens_per_entry,
+ "sink_only": 0,
+ }
+ expected_buf_tokens = (2 * window - 1) * tokens_per_entry
+ for frame_idx in range(1, 4 * window): # runs well past saturation and ring wraparound
+ key = torch.full((batch_size, tokens_per_entry, num_heads, head_dim), float(frame_idx - 1)) # [B,S_entry,H,D]
+ cache.gen_cache.store_kv(key, key + 100.0, frame_idx=frame_idx - 1)
+ for kind, max_tokens in kinds.items():
+ dynamic = ARMemoryState(
+ [cache],
+ frame_idx=frame_idx,
+ transfer_history_sink_tokens=sink_tokens,
+ transfer_history_max_tokens=max_tokens,
+ )
+ dynamic.init({"_num_full_tokens": tokens_per_entry}, torch.device("cpu"))
+ dynamic_value = dynamic.read_for_layer(0)
+ static = ARMemoryState(
+ [cache],
+ frame_idx=frame_idx,
+ vision_token_shapes=[(1, 1, tokens_per_entry)],
+ for_cuda_graphs=True,
+ num_kv_heads=num_heads,
+ head_dim=head_dim,
+ transfer_history_sink_tokens=sink_tokens,
+ transfer_history_max_tokens=max_tokens,
+ )
+ static.init({"_num_full_tokens": tokens_per_entry}, torch.device("cpu"))
+ static_value = static.read_for_layer(0)
+ expected_len = 0 if dynamic_value.gen_k_hist is None else dynamic_value.gen_k_hist.shape[1]
+ context = (window, sink, frame_idx, kind)
+ assert static_value.for_cuda_graphs is True, context
+ assert static_value.kv_k_static is not None and static_value.kv_v_static is not None, context
+ hist = static_value.static_hist_offset
+ assert hist == s_und + tokens_per_entry, context
+ assert static_value.kv_k_static.shape == (1, hist + expected_buf_tokens, num_heads, head_dim), context
+ assert static_value.real_gen_cache_len_t is not None and static_value.cu_seqlens_kv_t is not None, context
+ real_len = int(static_value.real_gen_cache_len_t.item())
+ assert real_len == expected_len, context
+ assert int(static_value.cu_seqlens_kv_t[1].item()) == s_und + tokens_per_entry + expected_len, context
+ assert static_value.max_seqlen_KV == s_und + tokens_per_entry + expected_buf_tokens, context
+ if expected_len:
+ assert dynamic_value.gen_v_hist is not None
+ torch.testing.assert_close(
+ static_value.kv_k_static[:, hist : hist + real_len], dynamic_value.gen_k_hist
+ )
+ torch.testing.assert_close(
+ static_value.kv_v_static[:, hist : hist + real_len], dynamic_value.gen_v_hist
+ )
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_ar_memory_state_transfer_static_matches_dynamic_with_sink() -> None:
+ """W=3, S=1: pinned control/RGB pair plus pair-aligned recent history, static == dynamic."""
+ _assert_transfer_static_matches_dynamic(window=3, sink=1)
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_ar_memory_state_transfer_static_matches_dynamic_without_sink() -> None:
+ """W=4, S=0: legacy total-suffix limiting, static == dynamic."""
+ _assert_transfer_static_matches_dynamic(window=4, sink=0)
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_ar_memory_state_transfer_static_flavor_construction_rules() -> None:
+ """Transfer limiting is allowed with CUDA-graph static shapes but not with post-saturation compile."""
+ tokens_per_entry = 2
+ cache = DualKVCache(gen_cache_size=6, attention_sink_size=2)
+ cache.und_cache.store(torch.randn(1, 3, 1, 1), torch.randn(1, 3, 1, 1)) # [B,S_und,H,D] each
+ cache.gen_cache.store_kv(
+ torch.zeros(1, tokens_per_entry, 1, 1), torch.zeros(1, tokens_per_entry, 1, 1), frame_idx=0
+ ) # [B,S_entry,H,D] each
+
+ state = ARMemoryState(
+ [cache],
+ frame_idx=1,
+ vision_token_shapes=[(1, 1, tokens_per_entry)],
+ for_cuda_graphs=True,
+ num_kv_heads=1,
+ head_dim=1,
+ transfer_history_sink_tokens=2 * tokens_per_entry,
+ transfer_history_max_tokens=2 * tokens_per_entry,
+ )
+ state.init({"_num_full_tokens": tokens_per_entry}, torch.device("cpu"))
+ assert state._static_history_indices == (0,)
+
+ with pytest.raises(ValueError, match="post-saturation"):
+ ARMemoryState(
+ [cache],
+ frame_idx=1,
+ post_saturation_static_compile=True,
+ static_und_cache_max_len=8,
+ transfer_history_sink_tokens=2 * tokens_per_entry,
+ transfer_history_max_tokens=2 * tokens_per_entry,
+ )
+
+ ragged = ARMemoryState(
+ [cache],
+ frame_idx=1,
+ vision_token_shapes=[(1, 1, tokens_per_entry)],
+ for_cuda_graphs=True,
+ num_kv_heads=1,
+ head_dim=1,
+ transfer_history_sink_tokens=2 * tokens_per_entry,
+ transfer_history_max_tokens=3, # not a whole number of entries
+ )
+ with pytest.raises(ValueError, match="whole cache entries"):
+ ragged.init({"_num_full_tokens": tokens_per_entry}, torch.device("cpu"))
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_kv_buffer_pool_reuses_matching_buffers_and_marks_static_address_once(monkeypatch: pytest.MonkeyPatch) -> None:
+ """Same key + shape returns the same tensor (stable address); marking happens only on allocation."""
+ marked: list[tuple[int, bool]] = []
+ monkeypatch.setattr(
+ torch._dynamo, "mark_static_address", lambda t, guard=False: marked.append((t.data_ptr(), guard))
+ )
+ pool = KVBufferPool()
+ first = pool.acquire(
+ "static_k", 3, (1, 8, 2, 4), dtype=torch.float32, device=torch.device("cpu"), static_address=True
+ )
+ again = pool.acquire(
+ "static_k", 3, (1, 8, 2, 4), dtype=torch.float32, device=torch.device("cpu"), static_address=True
+ )
+ assert again is first
+ # guard=False like module parameters: an identity guard would recompile once per layer (shared code object).
+ assert marked == [(first.data_ptr(), False)]
+
+ other_slot = pool.acquire(
+ "static_k", 4, (1, 8, 2, 4), dtype=torch.float32, device=torch.device("cpu"), static_address=True
+ )
+ assert other_slot is not first
+ resized = pool.acquire(
+ "static_k", 3, (1, 16, 2, 4), dtype=torch.float32, device=torch.device("cpu"), static_address=True
+ )
+ assert resized is not first and resized.shape == (1, 16, 2, 4)
+ assert len(marked) == 3
+
+ ring = pool.acquire("ring_k", 3, (6, 1, 2, 2, 4), dtype=torch.float32, device=torch.device("cpu"))
+ assert ring.shape == (6, 1, 2, 2, 4)
+ assert len(marked) == 3 # rings are never CUDA-graph inputs: no marking
+ pool.clear()
+ assert pool.acquire("ring_k", 3, (6, 1, 2, 2, 4), dtype=torch.float32, device=torch.device("cpu")) is not ring
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_gen_kv_cache_preallocated_ring_matches_clone_storage_across_wraparound() -> None:
+ """Ring slots are written in place and exposed as views; history reads equal the clone-based cache."""
+ batch_size, tokens_per_entry, num_heads, head_dim = 1, 2, 1, 1
+ clone_cache = GenKVCache(cache_size=6, attention_sink_size=2)
+ ring_cache = GenKVCache(cache_size=6, attention_sink_size=2, preallocate_ring=True)
+ for entry_idx in range(9): # wraps the 4 rolling slots twice
+ key = torch.full((batch_size, tokens_per_entry, num_heads, head_dim), float(entry_idx)) # [B,S_entry,H,D]
+ value = key + 100.0 # [B,S_entry,H,D]
+ clone_cache.store_kv(key, value, frame_idx=entry_idx)
+ ring_cache.store_kv(key.clone(), value.clone(), frame_idx=entry_idx)
+ for frame_idx in range(1, entry_idx + 2):
+ expected_k, expected_v = clone_cache.fetch_kv(frame_idx)
+ got_k, got_v = ring_cache.fetch_kv(frame_idx)
+ assert expected_k is not None and got_k is not None
+ torch.testing.assert_close(got_k, expected_k)
+ torch.testing.assert_close(got_v, expected_v)
+
+ assert ring_cache._ring_k is not None and ring_cache._ring_v is not None
+ assert ring_cache._ring_k.shape == (6, batch_size, tokens_per_entry, num_heads, head_dim)
+ for slot in range(6):
+ entry = ring_cache.k_cache[slot]
+ assert isinstance(entry, torch.Tensor)
+ assert entry._base is ring_cache._ring_k # a view into the ring, not a clone
+ # Overwriting a source tensor after the store must not leak into the cache (copy, not alias).
+ source = torch.zeros(batch_size, tokens_per_entry, num_heads, head_dim) # [B,S_entry,H,D]
+ ring_cache.store_kv(source, source, frame_idx=9)
+ source.fill_(-1.0)
+ stored = ring_cache.k_cache[ring_cache._cache_index(9)]
+ assert isinstance(stored, torch.Tensor) and torch.all(stored == 0.0)
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_gen_kv_cache_ring_is_skipped_for_unbounded_fp8_or_variable_entries() -> None:
+ """The ring needs a finite BF16 cache with uniform entries; everything else keeps per-entry clones."""
+ entry = torch.zeros(1, 2, 1, 1) # [B,S_entry,H,D]
+ unbounded = GenKVCache(cache_size=None, preallocate_ring=True)
+ unbounded.store_kv(entry, entry, frame_idx=0)
+ assert unbounded._ring_k is None
+ assert unbounded.k_cache[0] is not entry # still cloned
+
+ fp8 = GenKVCache(cache_size=6, backend=FP8StorageBackend(kernel_impl="torch"), preallocate_ring=True)
+ fp8.store_kv(entry, entry, frame_idx=0)
+ assert fp8._ring_k is None
+ assert isinstance(fp8.k_cache[0], tuple)
+
+ ragged = GenKVCache(cache_size=6, preallocate_ring=True)
+ ragged.store_kv(entry, entry, frame_idx=0)
+ assert ragged._ring_k is not None
+ wider = torch.ones(1, 3, 1, 1) # [B,S_entry+1,H,D]: different token count
+ ragged.store_kv(wider, wider, frame_idx=1)
+ assert ragged._ring_disabled is True
+ assert isinstance(ragged.k_cache[1], torch.Tensor) and ragged.k_cache[1]._base is not ragged._ring_k
+ k_hist, _ = ragged.fetch_kv(2)
+ assert k_hist is not None and k_hist.shape[1] == 5 # 2 + 3 tokens, mixed storage still concatenates
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_gen_kv_cache_pool_shares_ring_and_static_workspaces_across_generations() -> None:
+ """A second generation's cache reuses the pooled tensors (same addresses) yet reads only its own entries."""
+ batch_size, tokens_per_entry, num_heads, head_dim = 1, 2, 1, 1
+ pool = KVBufferPool()
+ kwargs = dict(num_heads=num_heads, head_dim=head_dim, device=torch.device("cpu"), dtype=torch.float32)
+ max_tokens = 5 * tokens_per_entry
+
+ first = GenKVCache(cache_size=6, buffer_pool=pool, pool_slot=7, preallocate_ring=True)
+ for entry_idx in range(4):
+ key = torch.full((batch_size, tokens_per_entry, num_heads, head_dim), float(entry_idx)) # [B,S_entry,H,D]
+ first.store_kv(key, key + 100.0, frame_idx=entry_idx)
+ first_k_buf, _, first_len = first.fetch_kv_static(4, max_tokens, tokens_per_entry, **kwargs)
+ assert first_len == 4 * tokens_per_entry
+
+ first_ring_k = first._ring_k
+ first.reset() # the previous generation's cache releases its slots (as when its generator finishes)
+ second = GenKVCache(cache_size=6, buffer_pool=pool, pool_slot=7, preallocate_ring=True)
+ key = torch.full((batch_size, tokens_per_entry, num_heads, head_dim), 42.0) # [B,S_entry,H,D]
+ second.store_kv(key, key + 100.0, frame_idx=0)
+ assert second._ring_k is first_ring_k # pooled ring
+ second_k_buf, _, second_len = second.fetch_kv_static(1, max_tokens, tokens_per_entry, **kwargs)
+ assert second_k_buf is first_k_buf # pooled static buffer: address stable across generations
+ assert second_len == tokens_per_entry
+ torch.testing.assert_close(second_k_buf[:, :second_len], key) # rebuilt from the new generation, not stale
+
+ unrelated = GenKVCache(cache_size=6, buffer_pool=pool, pool_slot=8, preallocate_ring=True)
+ unrelated.store_kv(key, key, frame_idx=0)
+ assert unrelated._ring_k is not first._ring_k
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_ar_memory_state_transfer_static_matches_dynamic_with_pooled_ring_cache() -> None:
+ """The static/dynamic equivalence holds with in-place ring storage and pooled workspaces."""
+ _assert_transfer_static_matches_dynamic(
+ window=3, sink=1, buffer_pool=KVBufferPool(), pool_slot=0, preallocate_ring=True
+ )
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_kv_buffer_pool_refuses_to_share_a_slot_between_live_caches() -> None:
+ """Two live caches on one slot would silently corrupt each other; the pool fails loudly instead."""
+ tokens_per_entry = 2
+ pool = KVBufferPool()
+ entry = torch.zeros(1, tokens_per_entry, 1, 1) # [B,S_entry,H,D]
+ first = GenKVCache(cache_size=6, buffer_pool=pool, pool_slot=3, preallocate_ring=True)
+ first.store_kv(entry, entry, frame_idx=0)
+ second = GenKVCache(cache_size=6, buffer_pool=pool, pool_slot=3, preallocate_ring=True)
+ with pytest.raises(RuntimeError, match="still owned"):
+ second.store_kv(entry, entry, frame_idx=0)
+ # Releasing the first cache (reset or garbage collection) frees the slot.
+ first.reset()
+ second.store_kv(entry, entry, frame_idx=0)
+ assert second._ring_k is not None
+ third = GenKVCache(cache_size=6, buffer_pool=pool, pool_slot=3, preallocate_ring=True)
+ del second
+ third.store_kv(entry, entry, frame_idx=0) # previous owner is gone
+ assert third._ring_k is not None
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_gen_kv_cache_releases_pooled_ring_when_entries_turn_ragged() -> None:
+ """Falling back to clone storage must not leave a dead ring pinned in the pool."""
+ pool = KVBufferPool()
+ cache = GenKVCache(cache_size=6, buffer_pool=pool, pool_slot=0, preallocate_ring=True)
+ cache.store_kv(torch.zeros(1, 2, 1, 1), torch.zeros(1, 2, 1, 1), frame_idx=0) # [B,S_entry,H,D]
+ assert ("ring_k", 0) in pool._buffers
+ cache.store_kv(torch.ones(1, 3, 1, 1), torch.ones(1, 3, 1, 1), frame_idx=1) # different token count
+ assert cache._ring_disabled is True
+ assert ("ring_k", 0) not in pool._buffers and ("ring_v", 0) not in pool._buffers
+ k_hist, _ = cache.fetch_kv(2)
+ assert k_hist is not None and k_hist.shape[1] == 5 # entry 0 stays readable through its own storage
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_kv_buffer_pool_tensors_are_not_inference_tensors() -> None:
+ """Buffers allocated under inference_mode must stay usable (in-place) from a later no_grad generation."""
+ pool = KVBufferPool()
+ with torch.inference_mode():
+ buf = pool.acquire("static_k", 0, (1, 4, 1, 1), dtype=torch.float32, device=torch.device("cpu"))
+ assert not buf.is_inference()
+ with torch.no_grad():
+ buf[:, :2].copy_(torch.ones(1, 2, 1, 1)) # would raise on an inference tensor
+ assert float(buf[0, 0, 0, 0]) == 1.0
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_ar_memory_state_static_flavor_coarse_replay_updates_offsets_in_place_and_matches_dynamic() -> None:
+ """Forward-scope CUDA graphs: offset tensors keep their identity across init calls and replays,
+ KV writes are staged until commit, and the replay-time rebuild reproduces the dynamic transfer history."""
+ window, sink, tokens_per_entry, num_heads, head_dim, s_und = 3, 1, 2, 1, 1, 3
+ pool = KVBufferPool()
+ cache = DualKVCache(
+ gen_cache_size=2 * window, attention_sink_size=2 * sink, buffer_pool=pool, pool_slot=0, preallocate_ring=True
+ )
+ und_k = torch.randn(1, s_und, num_heads, head_dim) # [B,S_und,H,D]
+ und_v = torch.randn(1, s_und, num_heads, head_dim) # [B,S_und,H,D]
+ cache.und_cache.store(und_k, und_v)
+ recent = window - sink - 1
+ sink_tokens = 2 * sink * tokens_per_entry
+ target_tokens = (2 * recent + 1) * tokens_per_entry
+ for entry_idx in range(3):
+ key = torch.full((1, tokens_per_entry, num_heads, head_dim), float(entry_idx)) # [B,S_entry,H,D]
+ cache.gen_cache.store_kv(key, key + 100.0, frame_idx=entry_idx)
+ hidden = {"_num_full_tokens": tokens_per_entry}
+ device = torch.device("cpu")
+
+ state = ARMemoryState(
+ [cache],
+ frame_idx=3,
+ vision_token_shapes=[(1, 1, tokens_per_entry)],
+ for_cuda_graphs=True,
+ num_kv_heads=num_heads,
+ head_dim=head_dim,
+ write_gen_cache=True,
+ coarse_cuda_graph=True,
+ stage_gen_cache_writes=True,
+ transfer_history_sink_tokens=sink_tokens,
+ transfer_history_max_tokens=target_tokens,
+ )
+ state.init(hidden, device) # warmup forward #1
+ cu_kv, cu_q, real_len_t = state._cu_seqlens_kv_t, state._cu_seqlens_q_t, state._real_gen_cache_len_t
+ assert cu_kv is not None and cu_q is not None and real_len_t is not None
+ first = state.read_for_layer(0)
+ assert first.kv_k_static is not None
+ buffer_ptr = first.kv_k_static.data_ptr()
+ hist = first.static_hist_offset
+ assert hist == s_und + tokens_per_entry
+ assert torch.equal(first.kv_k_static[:, :s_und], und_k) # und region primed once
+
+ state.init(hidden, device) # warmup forward #2 / capture: same tensors, values left to prepare()
+ assert (
+ state._cu_seqlens_kv_t is cu_kv and state._cu_seqlens_q_t is cu_q and state._real_gen_cache_len_t is real_len_t
+ )
+ assert int(real_len_t.item()) == 3 * tokens_per_entry
+ assert int(cu_kv[1].item()) == s_und + tokens_per_entry + 3 * tokens_per_entry
+
+ key3 = torch.full((1, tokens_per_entry, num_heads, head_dim), 3.0) # [B,S_entry,H,D]
+ state.write_for_layer(0, (key3, key3 + 100.0, und_k, und_v))
+ assert cache.gen_cache.k_cache[cache.gen_cache._cache_index(3)] is None # staged, not stored during capture
+ state.commit_staged_gen_cache(frame_idx=3)
+ stored = cache.gen_cache.k_cache[cache.gen_cache._cache_index(3)]
+ assert isinstance(stored, torch.Tensor) and torch.equal(stored, key3)
+
+ # Replay for the next target frame: same tensors and buffer address, values and contents for frame 4.
+ state.prepare_for_coarse_cuda_graph_replay(frame_idx=4)
+ dynamic = ARMemoryState(
+ [cache], frame_idx=4, transfer_history_sink_tokens=sink_tokens, transfer_history_max_tokens=target_tokens
+ )
+ dynamic.init(hidden, device)
+ dynamic_value = dynamic.read_for_layer(0)
+ assert dynamic_value.gen_k_hist is not None and dynamic_value.gen_v_hist is not None
+ expected_len = dynamic_value.gen_k_hist.shape[1]
+ assert int(real_len_t.item()) == expected_len
+ assert int(cu_kv[1].item()) == s_und + tokens_per_entry + expected_len
+ replayed = state.read_for_layer(0)
+ assert replayed.kv_k_static is not None and replayed.kv_v_static is not None
+ assert replayed.kv_k_static.data_ptr() == buffer_ptr
+ assert replayed.real_gen_cache_len_t is real_len_t and replayed.cu_seqlens_kv_t is cu_kv
+ torch.testing.assert_close(replayed.kv_k_static[:, hist : hist + expected_len], dynamic_value.gen_k_hist)
+ torch.testing.assert_close(replayed.kv_v_static[:, hist : hist + expected_len], dynamic_value.gen_v_hist)
+ assert torch.equal(replayed.kv_k_static[:, :s_und], und_k) # und region untouched by rebuilds
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_ar_memory_state_coarse_static_init_leaves_offset_values_to_prepare() -> None:
+ """A whole-forward capture records ``init()``; it must not write offset values that every replay
+ would then re-apply over the per-frame refresh done by ``prepare_for_coarse_cuda_graph_replay``."""
+ tokens_per_entry, num_heads, head_dim, s_und = 2, 1, 1, 3
+ cache = DualKVCache(gen_cache_size=4)
+ cache.und_cache.store(torch.randn(1, s_und, num_heads, head_dim), torch.randn(1, s_und, num_heads, head_dim))
+ for entry_idx in range(3):
+ key = torch.full((1, tokens_per_entry, num_heads, head_dim), float(entry_idx)) # [B,S_entry,H,D]
+ cache.gen_cache.store_kv(key, key, frame_idx=entry_idx)
+ hidden = {"_num_full_tokens": tokens_per_entry}
+ device = torch.device("cpu")
+ state = ARMemoryState(
+ [cache],
+ frame_idx=1,
+ vision_token_shapes=[(1, 1, tokens_per_entry)],
+ for_cuda_graphs=True,
+ num_kv_heads=num_heads,
+ head_dim=head_dim,
+ coarse_cuda_graph=True,
+ stage_gen_cache_writes=False,
+ )
+ state.init(hidden, device) # warmup #1 (outside capture) creates the tensors with frame-1 values
+ real_len_t, cu_kv = state._real_gen_cache_len_t, state._cu_seqlens_kv_t
+ assert real_len_t is not None and cu_kv is not None
+ assert int(real_len_t.item()) == tokens_per_entry
+
+ state.prepare_for_coarse_cuda_graph_replay(frame_idx=3) # replay-time refresh for frame 3
+ assert int(real_len_t.item()) == 3 * tokens_per_entry
+
+ # The recorded forward ran ``init()`` with the capture frame; replaying it must not clobber frame 3.
+ state.frame_idx = 1
+ state.init(hidden, device)
+ assert state._real_gen_cache_len_t is real_len_t and state._cu_seqlens_kv_t is cu_kv
+ assert int(real_len_t.item()) == 3 * tokens_per_entry
+ assert int(cu_kv[1].item()) == s_und + tokens_per_entry + 3 * tokens_per_entry
+
+ # Per-block CUDA graphs (non-coarse) rebuild the state every forward and keep refreshing in init().
+ per_block = ARMemoryState(
+ [cache],
+ frame_idx=1,
+ vision_token_shapes=[(1, 1, tokens_per_entry)],
+ for_cuda_graphs=True,
+ num_kv_heads=num_heads,
+ head_dim=head_dim,
+ )
+ per_block.init(hidden, device)
+ per_block.frame_idx = 3
+ per_block.init(hidden, device)
+ assert per_block._real_gen_cache_len_t is not None
+ assert int(per_block._real_gen_cache_len_t.item()) == 3 * tokens_per_entry
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_static_kv_workspaces_are_zero_filled_on_allocation() -> None:
+ """The varlen FMHA reads each sequence's last KV tile in full and masks the overrun; padding past
+ the real prefix must therefore be finite. Pooled and plain workspaces start zeroed."""
+ pooled = KVBufferPool().acquire("static_k", 0, (1, 8, 2, 4), dtype=torch.float32, device=torch.device("cpu"))
+ assert torch.equal(pooled, torch.zeros_like(pooled))
+ cache = GenKVCache(cache_size=3)
+ cache.store_kv(torch.ones(1, 2, 2, 4), torch.ones(1, 2, 2, 4), frame_idx=0)
+ k_buf, v_buf, real_len = cache.fetch_kv_static(
+ 1, 4, 2, num_heads=2, head_dim=4, device=torch.device("cpu"), dtype=torch.float32
+ )
+ assert real_len == 2
+ assert torch.equal(k_buf[:, real_len:], torch.zeros_like(k_buf[:, real_len:]))
+ assert torch.equal(v_buf[:, real_len:], torch.zeros_like(v_buf[:, real_len:]))
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_gen_kv_cache_composite_static_buffer_single_pass_rebuild_with_wraparound() -> None:
+ """The ``[und | curr | hist | pad]`` buffer: und primed once, history rebuilt straight from the
+ ring in physical-slot runs (sink run + wrapped recent runs), byte-identical to the dynamic history."""
+ cache_size, sink, tokens, num_heads, head_dim, s_und, gen_len = 6, 2, 2, 1, 1, 3, 2
+ pool = KVBufferPool()
+ cache = GenKVCache(
+ cache_size=cache_size, attention_sink_size=sink, buffer_pool=pool, pool_slot=0, preallocate_ring=True
+ )
+ for entry_idx in range(9): # wraps the 4 rolling slots twice
+ key = torch.full((1, tokens, num_heads, head_dim), float(entry_idx)) # [B,S_entry,H,D]
+ cache.store_kv(key, key + 100.0, frame_idx=entry_idx)
+ und_k = torch.randn(1, s_und, num_heads, head_dim) # [B,S_und,H,D]
+ und_v = torch.randn(1, s_und, num_heads, head_dim) # [B,S_und,H,D]
+ max_tokens = (cache_size - 1) * tokens
+ prefix = s_und + gen_len
+
+ k_buf, v_buf, real_len = cache.fetch_kv_static(
+ 8,
+ max_tokens,
+ tokens,
+ num_heads=num_heads,
+ head_dim=head_dim,
+ device=torch.device("cpu"),
+ dtype=torch.float32,
+ prefix_tokens=prefix,
+ und_kv=(und_k, und_v),
+ )
+ assert k_buf.shape == (1, prefix + max_tokens, num_heads, head_dim)
+ slots = [cache._cache_index(i) for i in cache._history_frame_indices(8)]
+ assert len(GenKVCache._coalesce_slot_runs(slots)) == 3 # sink run + rolling run split by the wrap
+ k_dyn, v_dyn = cache.fetch_kv(8)
+ assert k_dyn is not None and v_dyn is not None and real_len == k_dyn.shape[1]
+ assert torch.equal(k_buf[:, :s_und], und_k) and torch.equal(v_buf[:, :s_und], und_v)
+ assert torch.equal(k_buf[:, prefix : prefix + real_len], k_dyn)
+ assert torch.equal(v_buf[:, prefix : prefix + real_len], v_dyn)
+
+ # Same und tensors: the region is not rewritten (a scribble survives); new und tensors re-prime it.
+ k_buf[:, :s_und].fill_(-7.0)
+ cache.fetch_kv_static(
+ 8,
+ max_tokens,
+ tokens,
+ num_heads=num_heads,
+ head_dim=head_dim,
+ device=torch.device("cpu"),
+ dtype=torch.float32,
+ prefix_tokens=prefix,
+ und_kv=(und_k, und_v),
+ )
+ assert torch.all(k_buf[:, :s_und] == -7.0)
+ und_k2 = torch.randn(1, s_und, num_heads, head_dim) # [B,S_und,H,D]
+ cache.fetch_kv_static(
+ 8,
+ max_tokens,
+ tokens,
+ num_heads=num_heads,
+ head_dim=head_dim,
+ device=torch.device("cpu"),
+ dtype=torch.float32,
+ prefix_tokens=prefix,
+ und_kv=(und_k2, und_v),
+ )
+ assert torch.equal(k_buf[:, :s_und], und_k2)
+ # A store invalidates the history region only; the und region and buffer address survive.
+ ptr = k_buf.data_ptr()
+ key9 = torch.full((1, tokens, num_heads, head_dim), 9.0) # [B,S_entry,H,D]
+ cache.store_kv(key9, key9 + 100.0, frame_idx=9)
+ k_buf2, _, real_len2 = cache.fetch_kv_static(
+ 10,
+ max_tokens,
+ tokens,
+ num_heads=num_heads,
+ head_dim=head_dim,
+ device=torch.device("cpu"),
+ dtype=torch.float32,
+ prefix_tokens=prefix,
+ und_kv=(und_k2, und_v),
+ )
+ k_dyn2, _ = cache.fetch_kv(10)
+ assert k_buf2.data_ptr() == ptr and k_dyn2 is not None
+ assert torch.equal(k_buf2[:, prefix : prefix + real_len2], k_dyn2)
+ assert torch.equal(k_buf2[:, :s_und], und_k2)
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_gen_kv_cache_composite_buffer_holds_one_row_per_batched_sample() -> None:
+ """B rows at a fixed stride: und right-aligned per row, history copied per row from the ring, pads zero."""
+ cache_size, tokens, num_heads, head_dim, rows, und_cap, gen_len = 5, 2, 1, 2, 2, 4, 2
+ cache = GenKVCache(cache_size=cache_size, buffer_pool=KVBufferPool(), pool_slot=0, preallocate_ring=True)
+ for entry_idx in range(6): # wraps the 5-slot ring
+ key = torch.randn(rows, tokens, num_heads, head_dim) # [B,S_entry,H,D]
+ cache.store_kv(key, key * 2, frame_idx=entry_idx)
+ und_k = torch.randn(rows, und_cap, num_heads, head_dim) # [B,U,H,D]
+ und_v = torch.randn(rows, und_cap, num_heads, head_dim) # [B,U,H,D]
+ und_lens = (3, 4)
+ max_tokens = (cache_size - 1) * tokens
+ prefix = und_cap + gen_len
+ stride = prefix + max_tokens
+
+ k_buf, v_buf, real_len = cache.fetch_kv_static(
+ 6,
+ max_tokens,
+ tokens,
+ num_heads=num_heads,
+ head_dim=head_dim,
+ device=torch.device("cpu"),
+ dtype=torch.float32,
+ prefix_tokens=prefix,
+ und_kv=(und_k, und_v),
+ batch_rows=rows,
+ und_lens=und_lens,
+ )
+ assert k_buf.shape == (1, rows * stride, num_heads, head_dim)
+ k_dyn, v_dyn = cache.fetch_kv(6) # [B,S_hist,H,D]
+ assert k_dyn is not None and v_dyn is not None and real_len == k_dyn.shape[1]
+ for row in range(rows):
+ base = row * stride
+ u = und_lens[row]
+ assert torch.equal(k_buf[0, base + und_cap - u : base + und_cap], und_k[row, :u])
+ assert torch.equal(v_buf[0, base + und_cap - u : base + und_cap], und_v[row, :u])
+ assert torch.all(k_buf[0, base : base + und_cap - u] == 0) # lead pad
+ assert torch.equal(k_buf[0, base + prefix : base + prefix + real_len], k_dyn[row])
+ assert torch.equal(v_buf[0, base + prefix : base + prefix + real_len], v_dyn[row])
+ assert torch.all(k_buf[0, base + prefix + real_len : base + stride] == 0) # tail pad
+ # A store of B rows invalidates and the rebuild stays per row.
+ key = torch.randn(rows, tokens, num_heads, head_dim) # [B,S_entry,H,D]
+ cache.store_kv(key, key, frame_idx=6)
+ k_buf2, _, real_len2 = cache.fetch_kv_static(
+ 7,
+ max_tokens,
+ tokens,
+ num_heads=num_heads,
+ head_dim=head_dim,
+ device=torch.device("cpu"),
+ dtype=torch.float32,
+ prefix_tokens=prefix,
+ und_kv=(und_k, und_v),
+ batch_rows=rows,
+ und_lens=und_lens,
+ )
+ k_dyn2, _ = cache.fetch_kv(7)
+ assert k_dyn2 is not None and k_buf2.data_ptr() == k_buf.data_ptr()
+ for row in range(rows):
+ base = row * stride
+ assert torch.equal(k_buf2[0, base + prefix : base + prefix + real_len2], k_dyn2[row])
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_ar_memory_state_static_flavor_batched_rows_match_dynamic_batched_flavor() -> None:
+ """Static flavor with B=2: [2B] offsets with zero-query gap entries and per-row regions equal
+ to what the dynamic batched flavor reads for each sample."""
+ rows, gen_len, und_cap, num_heads, head_dim, cache_size = 2, 2, 4, 1, 2, 4
+ und_lens = (3, 4)
+ pool = KVBufferPool()
+ cache = DualKVCache(gen_cache_size=cache_size, buffer_pool=pool, pool_slot=0, preallocate_ring=True)
+ und_k = torch.randn(rows, und_cap, num_heads, head_dim) # [B,U,H,D]
+ und_v = torch.randn(rows, und_cap, num_heads, head_dim) # [B,U,H,D]
+ cache.und_cache.store(und_k, und_v, lengths=und_lens)
+ for entry_idx in range(2):
+ key = torch.randn(rows, gen_len, num_heads, head_dim) # [B,S_gen,H,D]
+ cache.gen_cache.store_kv(key, key + 1.0, frame_idx=entry_idx)
+ hidden = {
+ "sample_offsets": torch.tensor([0, 2, 4]),
+ "_num_full_tokens": rows * gen_len,
+ "_num_causal_tokens": 0,
+ }
+ common = dict(vision_token_shapes=[(1, 1, gen_len)], num_kv_heads=num_heads, head_dim=head_dim, batched=True)
+ static = ARMemoryState([cache], frame_idx=2, for_cuda_graphs=True, **common)
+ static.init(hidden, torch.device("cpu"))
+ dynamic = ARMemoryState([cache], frame_idx=2, batched=True)
+ dynamic.init(hidden, torch.device("cpu"))
+
+ sv, dv = static.read_for_layer(0), dynamic.read_for_layer(0)
+ hist_len = dv.gen_k_hist.shape[1]
+ stride = und_cap + gen_len + (cache_size - 1) * gen_len
+ assert sv.batch_size == rows and sv.gen_len == gen_len and sv.static_row_stride == stride
+ assert (
+ sv.max_seqlen_KV == stride and sv.static_curr_offset == und_cap and sv.static_hist_offset == und_cap + gen_len
+ )
+ assert sv.kv_k_static is not None and sv.kv_k_static.shape == (1, rows * stride, num_heads, head_dim)
+ assert sv.cu_seqlens_q_t.tolist() == [0, gen_len, gen_len, 2 * gen_len]
+ expected_kv = []
+ for row in range(rows):
+ expected_kv += [row * stride + und_cap - und_lens[row], row * stride + und_cap + gen_len + hist_len]
+ assert sv.cu_seqlens_kv_t.tolist() == expected_kv
+ for row in range(rows):
+ base = row * stride
+ u = und_lens[row]
+ assert torch.equal(sv.kv_k_static[0, base + und_cap - u : base + und_cap], dv.und_k_cached[row, :u])
+ hist = base + und_cap + gen_len
+ assert torch.equal(sv.kv_k_static[0, hist : hist + hist_len], dv.gen_k_hist[row])
+ assert torch.equal(sv.kv_v_static[0, hist : hist + hist_len], dv.gen_v_hist[row])
+
+ # Coarse replay flavor: offsets keep their identity, values follow the frame.
+ coarse = ARMemoryState(
+ [cache], frame_idx=2, for_cuda_graphs=True, coarse_cuda_graph=True, stage_gen_cache_writes=True, **common
+ )
+ coarse.init(hidden, torch.device("cpu"))
+ cu_kv = coarse._cu_seqlens_kv_t
+ key = torch.randn(rows, gen_len, num_heads, head_dim) # [B,S_gen,H,D]
+ cache.gen_cache.store_kv(key, key, frame_idx=2)
+ coarse.prepare_for_coarse_cuda_graph_replay(frame_idx=3)
+ assert coarse._cu_seqlens_kv_t is cu_kv
+ assert cu_kv.tolist()[1] == und_cap + gen_len + 3 * gen_len # row 0 now sees 3 history entries
+ assert cu_kv.tolist()[3] == stride + und_cap + gen_len + 3 * gen_len
+
+
+def _fresh_single_row_history(
+ entries: list[tuple[torch.Tensor, torch.Tensor]],
+ cache_size: int,
+ sink: int,
+ frame: int,
+ sink_entries: int,
+ recent: int,
+):
+ """Reference: a brand-new single-row cache fed the same entries; returns (k_hist, v_hist) at ``frame``."""
+ fresh = GenKVCache(
+ cache_size=cache_size, attention_sink_size=sink, buffer_pool=KVBufferPool(), preallocate_ring=True
+ )
+ for idx, (k, v) in enumerate(entries):
+ fresh.store_kv(k, v, frame_idx=idx)
+ plan = fresh.transfer_history_frame_indices(frame, sink_entries, recent)
+ tokens = entries[0][0].shape[1]
+ k_buf, v_buf, real_len = fresh.fetch_kv_static(
+ frame,
+ (cache_size - 1) * tokens,
+ tokens,
+ num_heads=k.shape[2],
+ head_dim=k.shape[3],
+ device=torch.device("cpu"),
+ dtype=torch.float32,
+ history_frame_indices=plan,
+ )
+ return k_buf[0, :real_len].clone(), v_buf[0, :real_len].clone()
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_gen_kv_cache_row_restart_reads_a_fresh_cache_history_for_that_row_only() -> None:
+ """Row 1 restarts at global frame 3: its plan/history equal a fresh single-row cache started there
+ (sink entries mirrored into its pinned slots), row 0 is untouched, and only tensor values changed."""
+ cache_size, sink, tokens, num_heads, head_dim, rows = 8, 2, 2, 1, 2, 2
+ sink_entries, recent = 2, 3
+ cache = GenKVCache(
+ cache_size=cache_size, attention_sink_size=sink, buffer_pool=KVBufferPool(), preallocate_ring=True
+ )
+ entries = [
+ (torch.randn(rows, tokens, num_heads, head_dim), torch.randn(rows, tokens, num_heads, head_dim))
+ for _ in range(12)
+ ]
+ restart_at = 3
+ for idx, (k, v) in enumerate(entries):
+ if idx == restart_at:
+ cache.reset_rows([1], at_frame_idx=restart_at)
+ cache.store_kv(k, v, frame_idx=idx)
+ frame = 12
+ max_tokens = (cache_size - 1) * tokens
+ plan_row0 = cache.row_history_slots(0, frame, sink_entries, recent)
+ plan_row1 = cache.row_history_slots(1, frame, sink_entries, recent)
+ # Row 0: the shared plan, unchanged by the other row's restart.
+ assert plan_row0 == [
+ cache._cache_index(i) for i in cache.transfer_history_frame_indices(frame, sink_entries, recent)
+ ]
+ # Row 1: a fresh cache's plan in its local clock; its sink lives in the mirrored pinned slots.
+ assert plan_row1[:sink_entries] == [0, 1]
+ k_buf, v_buf, _ = cache.fetch_kv_static(
+ frame,
+ max_tokens,
+ tokens,
+ num_heads=num_heads,
+ head_dim=head_dim,
+ device=torch.device("cpu"),
+ dtype=torch.float32,
+ batch_rows=rows,
+ row_slot_lists=[plan_row0, plan_row1],
+ )
+ stride = max_tokens
+ row0_ref = _fresh_single_row_history(
+ [(k[0:1], v[0:1]) for k, v in entries], cache_size, sink, frame, sink_entries, recent
+ )
+ row1_ref = _fresh_single_row_history(
+ [(k[1:2], v[1:2]) for k, v in entries[restart_at:]], cache_size, sink, frame - restart_at, sink_entries, recent
+ )
+ assert torch.equal(k_buf[0, 0 : row0_ref[0].shape[0]], row0_ref[0])
+ assert torch.equal(k_buf[0, stride : stride + row1_ref[0].shape[0]], row1_ref[0])
+ assert torch.equal(v_buf[0, stride : stride + row1_ref[1].shape[0]], row1_ref[1])
+ assert cache._static_real_lens == (row0_ref[0].shape[0], row1_ref[0].shape[0])
+ # Right after the restart the row has an empty history, then grows like a fresh cache.
+ assert cache.row_history_slots(1, restart_at, sink_entries, recent) == []
+ assert len(cache.row_history_slots(1, restart_at + 1, sink_entries, recent)) == 1
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_ar_memory_state_static_batched_offsets_follow_per_row_restarts() -> None:
+ """cu_seqlens_kv row ends shrink for a restarted row; addresses stay; B=1 without restarts keeps the shared plan."""
+ rows, gen_len, und_cap, num_heads, head_dim, cache_size, sink = 2, 2, 3, 1, 2, 8, 2
+ pool = KVBufferPool()
+ cache = DualKVCache(
+ gen_cache_size=cache_size, attention_sink_size=sink, buffer_pool=pool, pool_slot=0, preallocate_ring=True
+ )
+ cache.und_cache.store(
+ torch.randn(rows, und_cap, num_heads, head_dim), torch.randn(rows, und_cap, num_heads, head_dim)
+ )
+ for idx in range(5):
+ key = torch.randn(rows, gen_len, num_heads, head_dim)
+ if idx == 3:
+ cache.gen_cache.reset_rows([1], at_frame_idx=3)
+ cache.gen_cache.store_kv(key, key, frame_idx=idx)
+ hidden = {"sample_offsets": torch.tensor([0, 2, 4]), "_num_full_tokens": rows * gen_len, "_num_causal_tokens": 0}
+ state = ARMemoryState(
+ [cache],
+ frame_idx=5,
+ vision_token_shapes=[(1, 1, gen_len)],
+ for_cuda_graphs=True,
+ num_kv_heads=num_heads,
+ head_dim=head_dim,
+ batched=True,
+ coarse_cuda_graph=True,
+ stage_gen_cache_writes=True,
+ transfer_history_sink_tokens=2 * gen_len,
+ transfer_history_max_tokens=3 * gen_len,
+ )
+ state.init(hidden, torch.device("cpu"))
+ stride = und_cap + gen_len + (cache_size - 1) * gen_len
+ cu_kv = state._cu_seqlens_kv_t.tolist()
+ assert cu_kv[1] == und_cap + gen_len + 5 * gen_len # row 0: sink 2 + recent 3 entries
+ assert cu_kv[3] == stride + und_cap + gen_len + 2 * gen_len # row 1: 2 entries since its restart at 3
+ value = state.read_for_layer(0)
+ assert value.kv_k_static is not None and state._static_row_slots is not None
+ assert len(state._static_row_slots[1]) == 2
+ kv_t = state._cu_seqlens_kv_t
+ key = torch.randn(rows, gen_len, num_heads, head_dim)
+ cache.gen_cache.store_kv(key, key, frame_idx=5)
+ state.prepare_for_coarse_cuda_graph_replay(frame_idx=6)
+ assert state._cu_seqlens_kv_t is kv_t
+ assert kv_t.tolist()[3] == stride + und_cap + gen_len + 3 * gen_len
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_gen_kv_cache_static_rebuild_generic_path_matches_ring_path() -> None:
+ """Clone storage (no ring) takes the decode_many path; both layouts agree with the dynamic history."""
+ cache_size, sink, tokens, num_heads, head_dim, prefix = 5, 1, 3, 2, 2, 4
+ ring = GenKVCache(
+ cache_size=cache_size, attention_sink_size=sink, buffer_pool=KVBufferPool(), preallocate_ring=True
+ )
+ clone = GenKVCache(cache_size=cache_size, attention_sink_size=sink)
+ for entry_idx in range(7):
+ key = torch.randn(1, tokens, num_heads, head_dim) # [B,S_entry,H,D]
+ ring.store_kv(key, key * 2, frame_idx=entry_idx)
+ clone.store_kv(key, key * 2, frame_idx=entry_idx)
+ common = dict(num_heads=num_heads, head_dim=head_dim, device=torch.device("cpu"), dtype=torch.float32)
+ k_ring, v_ring, len_ring = ring.fetch_kv_static(
+ 7, (cache_size - 1) * tokens, tokens, prefix_tokens=prefix, **common
+ )
+ k_clone, v_clone, len_clone = clone.fetch_kv_static(
+ 7, (cache_size - 1) * tokens, tokens, prefix_tokens=prefix, **common
+ )
+ k_dyn, v_dyn = clone.fetch_kv(7)
+ assert k_dyn is not None and v_dyn is not None and len_ring == len_clone == k_dyn.shape[1]
+ assert torch.equal(k_ring[:, prefix : prefix + len_ring], k_dyn)
+ assert torch.equal(k_clone[:, prefix : prefix + len_clone], k_dyn)
+ assert torch.equal(v_ring[:, prefix : prefix + len_ring], v_dyn)
+ assert torch.equal(v_clone[:, prefix : prefix + len_clone], v_dyn)
+
+
+@pytest.mark.L0
+@pytest.mark.CPU
+def test_ar_memory_state_static_flavor_coarse_replay_without_transfer_limits() -> None:
+ """The plain (non-Transfer) static flavor replays with the ring-native history length."""
+ tokens_per_entry, num_heads, head_dim, s_und = 2, 1, 1, 3
+ cache = DualKVCache(gen_cache_size=4)
+ cache.und_cache.store(torch.randn(1, s_und, num_heads, head_dim), torch.randn(1, s_und, num_heads, head_dim))
+ for entry_idx in range(3):
+ key = torch.full((1, tokens_per_entry, num_heads, head_dim), float(entry_idx)) # [B,S_entry,H,D]
+ cache.gen_cache.store_kv(key, key, frame_idx=entry_idx)
+ state = ARMemoryState(
+ [cache],
+ frame_idx=1,
+ vision_token_shapes=[(1, 1, tokens_per_entry)],
+ for_cuda_graphs=True,
+ num_kv_heads=num_heads,
+ head_dim=head_dim,
+ coarse_cuda_graph=True,
+ stage_gen_cache_writes=False,
+ )
+ state.init({"_num_full_tokens": tokens_per_entry}, torch.device("cpu"))
+ assert state._real_gen_cache_len_t is not None and int(state._real_gen_cache_len_t.item()) == tokens_per_entry
+ state.prepare_for_coarse_cuda_graph_replay(frame_idx=3)
+ assert int(state._real_gen_cache_len_t.item()) == 3 * tokens_per_entry # min(3, cache_size-1) entries
+ value = state.read_for_layer(0)
+ assert value.kv_k_static is not None
+ hist = value.static_hist_offset
+ torch.testing.assert_close(
+ value.kv_k_static[:, hist : hist + 3 * tokens_per_entry],
+ torch.cat([torch.full((1, tokens_per_entry, num_heads, head_dim), float(i)) for i in range(3)], dim=1),
+ )
diff --git a/cosmos_framework/model/generator/utils/sr_latent_noise.py b/cosmos_framework/model/generator/utils/sr_latent_noise.py
new file mode 100644
index 000000000..9c9519cd5
--- /dev/null
+++ b/cosmos_framework/model/generator/utils/sr_latent_noise.py
@@ -0,0 +1,103 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+"""L1: Gaussian noise on the low-resolution conditioning latent of super-resolution samples.
+
+Cascaded generators that condition on their own (or a degraded) low-resolution input add noise to
+that conditioning at training time so the model does not copy its artifacts (noise-conditioning
+augmentation in Imagen / CDM, FlashVideo's latent blend, SeedVR's LR-latent noise). Here the noise
+is applied after VAE encoding to the conditioning vision items of samples that came from the SR
+data streams, and only during training.
+
+The mixing follows the rectified-flow interpolation used by the model, ``z' = (1 - t) z + t eps``
+with ``t ~ U[t_min, t_max]`` drawn once per sample. With ``persist_over_time`` the same ``eps`` is
+shared by every latent frame of the clip (temporal-persist, as in the Transfer1 decoder), otherwise
+``eps`` is i.i.d. per latent frame. The applied ``t`` is not fed to the network; a level embedding
+is a follow-up.
+"""
+
+from __future__ import annotations
+
+from typing import Sequence
+
+import attrs
+import torch
+
+
+@attrs.define(slots=False)
+class SRLatentConditionNoiseConfig:
+ """Config for L1 latent conditioning noise. ``None`` on the model config disables it."""
+
+ prob: float = 0.8
+ """Probability of noising the conditioning latents of an eligible sample."""
+
+ t_min: float = 0.05
+ t_max: float = 0.40
+ """Range of the interpolation weight ``t`` toward pure noise; calibrated by experiment E3."""
+
+ persist_over_time: bool = True
+ """Share one noise realisation across latent frames (True) or draw i.i.d. per frame (False)."""
+
+ dataset_names: tuple[str, ...] = ("video_sr", "image_sr")
+ """Only samples whose ``dataset_name`` is in this set are eligible (SR streams only)."""
+
+
+def sr_sample_mask(dataset_names: Sequence[str] | str | None, batch_size: int, eligible: Sequence[str]) -> list[bool]:
+ """Per-sample eligibility from ``data_batch["dataset_name"]`` (a list per sample, or one string)."""
+ if dataset_names is None:
+ return [False] * batch_size
+ if isinstance(dataset_names, str):
+ return [dataset_names in eligible] * batch_size
+ names = list(dataset_names)
+ if len(names) != batch_size:
+ raise ValueError(f"dataset_name has {len(names)} entries for batch_size {batch_size}")
+ return [str(n) in eligible for n in names]
+
+
+def apply_sr_latent_condition_noise(
+ x0_tokens_vision: list[torch.Tensor], # flattened items, each [1,C,T,H,W]
+ num_vision_items_per_sample: list[int] | None,
+ eligible_samples: Sequence[bool],
+ cfg: SRLatentConditionNoiseConfig,
+ generator: torch.Generator | None = None,
+) -> tuple[list[torch.Tensor], list[float | None]]:
+ """Noise the conditioning (non-last) vision items of eligible multi-item samples.
+
+ Returns the new latent list (non-eligible items are the same tensor objects) and the ``t``
+ used per sample (``None`` when the sample was skipped), for logging.
+ """
+ batch_size = len(eligible_samples)
+ if num_vision_items_per_sample is None:
+ # Single-item samples have no conditioning item to noise.
+ return list(x0_tokens_vision), [None] * batch_size
+ if len(num_vision_items_per_sample) != batch_size:
+ raise ValueError(
+ f"num_vision_items_per_sample has {len(num_vision_items_per_sample)} entries for batch_size {batch_size}"
+ )
+ if sum(num_vision_items_per_sample) != len(x0_tokens_vision):
+ raise ValueError(f"{len(x0_tokens_vision)} vision items but samples declare {sum(num_vision_items_per_sample)}")
+ out = list(x0_tokens_vision)
+ applied_t: list[float | None] = []
+ offset = 0
+ for eligible, num_items in zip(eligible_samples, num_vision_items_per_sample):
+ start, end = offset, offset + num_items
+ offset = end
+ if not eligible or num_items < 2:
+ applied_t.append(None)
+ continue
+ ref = x0_tokens_vision[start]
+ draw = torch.rand(2, generator=generator, device=ref.device if generator is None else generator.device) # [2]
+ if float(draw[0]) >= cfg.prob:
+ applied_t.append(None)
+ continue
+ t = cfg.t_min + float(draw[1]) * (cfg.t_max - cfg.t_min)
+ applied_t.append(t)
+ for item_idx in range(start, end - 1): # every item but the generated (last) one
+ z = x0_tokens_vision[item_idx] # [1,C,T,H,W]
+ if cfg.persist_over_time:
+ eps = torch.randn(
+ (z.shape[0], z.shape[1], 1, *z.shape[3:]), generator=generator, device=z.device, dtype=z.dtype
+ ).expand_as(z) # [1,C,T,H,W]
+ else:
+ eps = torch.randn(z.shape, generator=generator, device=z.device, dtype=z.dtype) # [1,C,T,H,W]
+ out[item_idx] = (1.0 - t) * z + t * eps # [1,C,T,H,W]
+ return out, applied_t
diff --git a/cosmos_framework/model/generator/utils/sr_latent_noise_test.py b/cosmos_framework/model/generator/utils/sr_latent_noise_test.py
new file mode 100644
index 000000000..71b6cb7e8
--- /dev/null
+++ b/cosmos_framework/model/generator/utils/sr_latent_noise_test.py
@@ -0,0 +1,78 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+
+import pytest
+import torch
+
+from cosmos_framework.model.generator.utils.sr_latent_noise import (
+ SRLatentConditionNoiseConfig,
+ apply_sr_latent_condition_noise,
+ sr_sample_mask,
+)
+
+pytestmark = [pytest.mark.L0, pytest.mark.CPU]
+
+
+def _latents() -> list[torch.Tensor]:
+ g = torch.Generator().manual_seed(0)
+ # Sample 0: SR pair (LR small grid, HR big grid). Sample 1: single item. Sample 2: transfer pair.
+ return [
+ torch.randn(1, 16, 3, 8, 13, generator=g),
+ torch.randn(1, 16, 3, 15, 26, generator=g),
+ torch.randn(1, 16, 3, 15, 26, generator=g),
+ torch.randn(1, 16, 3, 15, 26, generator=g),
+ torch.randn(1, 16, 3, 15, 26, generator=g),
+ ]
+
+
+def test_mask_from_dataset_names() -> None:
+ eligible = ("video_sr", "image_sr")
+ assert sr_sample_mask(["video_sr", "video_data", "video_transfer"], 3, eligible) == [True, False, False]
+ assert sr_sample_mask("image_sr", 2, eligible) == [True, True]
+ assert sr_sample_mask(None, 2, eligible) == [False, False]
+ with pytest.raises(ValueError):
+ sr_sample_mask(["video_sr"], 2, eligible)
+
+
+def test_only_sr_conditioning_items_are_noised_and_generated_items_untouched() -> None:
+ latents = _latents()
+ cfg = SRLatentConditionNoiseConfig(prob=1.0, t_min=0.3, t_max=0.3)
+ gen = torch.Generator().manual_seed(1)
+ out, applied = apply_sr_latent_condition_noise(latents, [2, 1, 2], [True, False, False], cfg, generator=gen)
+ assert applied[0] == pytest.approx(0.3) and applied[1] is None and applied[2] is None
+ assert not torch.equal(out[0], latents[0]) # LR of the SR sample is noised
+ for i in (1, 2, 3, 4):
+ assert out[i] is latents[i] # HR of SR sample, single item, transfer pair: untouched
+ # z' = 0.7 z + 0.3 eps -> variance 0.49 + 0.09 for unit-variance inputs
+ assert out[0].var().item() == pytest.approx(0.58, abs=0.05)
+
+
+def test_persist_over_time_shares_noise_across_latent_frames() -> None:
+ z = torch.zeros(1, 16, 4, 6, 6)
+ cfg_persist = SRLatentConditionNoiseConfig(prob=1.0, t_min=1.0, t_max=1.0, persist_over_time=True)
+ out, _ = apply_sr_latent_condition_noise([z, z.clone()], [2], [True], cfg_persist, torch.Generator().manual_seed(0))
+ assert torch.equal(out[0][:, :, 0], out[0][:, :, 3]) # pure noise, identical over time
+ cfg_iid = SRLatentConditionNoiseConfig(prob=1.0, t_min=1.0, t_max=1.0, persist_over_time=False)
+ out, _ = apply_sr_latent_condition_noise([z, z.clone()], [2], [True], cfg_iid, torch.Generator().manual_seed(0))
+ assert not torch.equal(out[0][:, :, 0], out[0][:, :, 3])
+
+
+def test_probability_and_range_are_respected() -> None:
+ latents = _latents()
+ cfg = SRLatentConditionNoiseConfig(prob=0.5, t_min=0.1, t_max=0.2)
+ gen = torch.Generator().manual_seed(3)
+ ts = []
+ for _ in range(200):
+ _, applied = apply_sr_latent_condition_noise(latents[:2], [2], [True], cfg, generator=gen)
+ ts.append(applied[0])
+ skipped = sum(t is None for t in ts)
+ assert 60 < skipped < 140 # prob 0.5
+ assert all(0.1 <= t <= 0.2 for t in ts if t is not None)
+
+
+def test_single_item_batches_and_bad_shapes() -> None:
+ latents = _latents()
+ out, applied = apply_sr_latent_condition_noise(latents, None, [True] * 5, SRLatentConditionNoiseConfig())
+ assert all(a is None for a in applied) and all(o is z for o, z in zip(out, latents))
+ with pytest.raises(ValueError):
+ apply_sr_latent_condition_noise(latents, [2, 2], [True, True], SRLatentConditionNoiseConfig())
diff --git a/cosmos_framework/trainer/__init__.py b/cosmos_framework/trainer/__init__.py
index 347caff72..94aea9449 100644
--- a/cosmos_framework/trainer/__init__.py
+++ b/cosmos_framework/trainer/__init__.py
@@ -5,6 +5,7 @@
import inspect
import os
import signal
+from collections.abc import Iterator
from dataclasses import dataclass
from typing import Any
@@ -95,6 +96,7 @@ def __init__(self, config):
"""
super().__init__()
self.config = config
+ self._validation_iterator: Iterator[Any] | None = None
# Set up the distributed computing environment.
with distributed_init():
distributed.init()
@@ -315,6 +317,13 @@ def train(
# callbacks) so data-augmentation randomness starts from a deterministic state
# regardless of how much RNG state init consumed.
misc.set_random_seed(seed=self.config.trainer.seed, by_rank=True)
+ if (
+ self.config.trainer.run_validation
+ and getattr(self.config.trainer, "prefetch_validation", False)
+ and self.config.trainer.max_val_iter is not None
+ and self._validation_iterator is None
+ ):
+ self._validation_iterator = iter(dataloader_val)
with (
maybe_enable_profiling(self.config, global_step=iteration) as torch_profiler,
maybe_enable_memory_snapshot(self.config, global_step=iteration) as memory_profiler,
@@ -374,6 +383,8 @@ def train(
)
# If the gradients are still being accumulated, continue to load the next training batch.
if grad_accum_iter != 0:
+ # Release this microstep's outputs before the next forward/backward.
+ del output_batch, loss
continue
# Do the following when an actual optimizer (update) step has been made.
iteration += 1
@@ -381,6 +392,9 @@ def train(
if iteration % self.config.checkpoint.save_iter == 0:
self.checkpointer.save(model, optimizer, scheduler, grad_scaler, iteration=iteration)
self.callbacks.on_training_step_end(model, data_batch, output_batch, loss, iteration=iteration)
+ # Callback consumers have finished; do not retain GPU outputs through
+ # validation or the next training step. Callback-owned references remain valid.
+ del output_batch, loss
# Validation.
if self.config.trainer.run_validation and iteration % self.config.trainer.validation_iter == 0:
self.validate(model, dataloader_val, iteration=iteration)
@@ -395,6 +409,7 @@ def train(
nsys_profiler.step()
if _end_training:
break
+ self._validation_iterator = None
log.success("Done with training.")
if sm_carveout:
torch._C._set_sm_carveout_experimental(None)
@@ -498,7 +513,7 @@ def _zero_grad(self, model: torch.nn.Module, optimizer: torch.optim.Optimizer, i
@torch.no_grad()
def validate(self, model: ImaginaireModel, dataloader_val: torch.utils.data.DataLoader, iteration: int = 0) -> None:
- """Validate on the full validation dataset.
+ """Validate on at most the configured number of validation batches.
Args:
model (ImaginaireModel): The PyTorch model.
@@ -507,13 +522,32 @@ def validate(self, model: ImaginaireModel, dataloader_val: torch.utils.data.Data
"""
self.callbacks.on_validation_start(model, dataloader_val, iteration=iteration)
model.eval()
- # Evaluate on the full validation set.
+ max_val_iter = self.config.trainer.max_val_iter
+ reuse_iterator = getattr(self.config.trainer, "prefetch_validation", False) and max_val_iter is not None
+ validation_iterator = self._validation_iterator if reuse_iterator else None
+ if validation_iterator is None:
+ validation_iterator = iter(dataloader_val)
+ if reuse_iterator:
+ self._validation_iterator = validation_iterator
+
+ val_iter = 0
+ restarted_iterator = False
with ema.ema_scope(model, enabled=model.config.ema.enabled):
- for val_iter, data_batch in enumerate(dataloader_val):
- if self.config.trainer.max_val_iter is not None and val_iter >= self.config.trainer.max_val_iter:
+ while max_val_iter is None or val_iter < max_val_iter:
+ try:
+ data_batch = next(validation_iterator)
+ except StopIteration:
+ if reuse_iterator:
+ self._validation_iterator = None
+ if val_iter == 0 and not restarted_iterator:
+ validation_iterator = iter(dataloader_val)
+ self._validation_iterator = validation_iterator
+ restarted_iterator = True
+ continue
break
data_batch = misc.to(data_batch, device="cuda")
self.callbacks.on_validation_step_start(model, data_batch, iteration=iteration)
output_batch, loss = model.validation_step(data_batch, iteration)
self.callbacks.on_validation_step_end(model, data_batch, output_batch, loss, iteration=iteration)
+ val_iter += 1
self.callbacks.on_validation_end(model, iteration=iteration)
diff --git a/cosmos_framework/utils/config.py b/cosmos_framework/utils/config.py
index 6c30fd298..04ccbe4ca 100644
--- a/cosmos_framework/utils/config.py
+++ b/cosmos_framework/utils/config.py
@@ -463,6 +463,8 @@ class TrainerConfig:
run_validation: bool = True
# How often we evaluate on the validation set.
validation_iter: int = 999999999
+ # Keep a bounded validation iterator alive so its worker can prefetch between validation calls.
+ prefetch_validation: bool = False
# Whether to run the validation on the start of the training.
run_validation_on_start: bool = False
# Kill the process after N seconds since the last iteration (usually means dead job).
diff --git a/cosmos_framework/utils/generator/deterministic_rng.py b/cosmos_framework/utils/generator/deterministic_rng.py
index a53a83f27..b02504e28 100644
--- a/cosmos_framework/utils/generator/deterministic_rng.py
+++ b/cosmos_framework/utils/generator/deterministic_rng.py
@@ -22,6 +22,9 @@ def sample_identity_seed(root: object, key: object, epoch: int = 0) -> int:
def sample_seed(data: dict[str, Any]) -> int:
url = data.get("__url__", "")
root = getattr(url, "root", url)
+ local_data_prefix = getattr(getattr(url, "meta", None), "opts", {}).get("local_data_prefix")
+ if isinstance(root, str) and isinstance(local_data_prefix, str):
+ root = root.removeprefix(local_data_prefix.rstrip("/") + "/")
sample_meta = getattr(url, "sample_meta", None)
epoch = int(data.get("sample_epoch", getattr(sample_meta, "sample_epoch", 0)))
return sample_identity_seed(root, data.get("__key__", ""), epoch)
diff --git a/cosmos_framework/utils/launch.py b/cosmos_framework/utils/launch.py
index 250a7eca6..31112e453 100644
--- a/cosmos_framework/utils/launch.py
+++ b/cosmos_framework/utils/launch.py
@@ -5,9 +5,10 @@
import os
import sys
import time
+from collections.abc import Mapping
import torch
-from omegaconf import OmegaConf
+import yaml
from cosmos_framework.utils.config import Config
from cosmos_framework.utils import distributed, log
@@ -20,6 +21,12 @@
S3_READY = False
+def _save_reproducibility_yaml(data: Mapping[str, object], path: str) -> None:
+ """Save metadata as literal YAML without interpreting OmegaConf interpolation syntax."""
+ with open(path, "w") as file:
+ yaml.safe_dump(data, file, sort_keys=False, default_flow_style=False, allow_unicode=True)
+
+
def log_reproducible_setup(config: Config, args: argparse.Namespace) -> None:
"""
Configures the environment for reproducibility of experiments by setting up
@@ -49,7 +56,7 @@ def log_reproducible_setup(config: Config, args: argparse.Namespace) -> None:
(like saving configurations) are restricted to the main node (rank 0).
- It uses the 'easy_io' module for interacting with S3, ensuring files are written and
read correctly from the object store.
- - It leverages OmegaConf for saving YAML configurations
+ - It saves metadata as literal YAML so shell syntax in captured git diffs remains unchanged.
- git information is read from 'git_commit.txt' and 'git_branch.txt' files if they exist.
- snapshot codebase is saved as 'codebase.zip' if it exists in the current directory.
@@ -114,10 +121,8 @@ def log_reproducible_setup(config: Config, args: argparse.Namespace) -> None:
job_info["git_diff"] = f.read().strip()
log.critical(f"git diff: {job_info['git_diff']}")
- with open(f"{job_local_path}/job_env.yaml", "w") as f:
- OmegaConf.save(job_info, f)
- with open(f"{job_local_path}/launch_info.yaml", "w") as f:
- OmegaConf.save(launch_info, f)
+ _save_reproducibility_yaml(job_info, f"{job_local_path}/job_env.yaml")
+ _save_reproducibility_yaml(launch_info, f"{job_local_path}/launch_info.yaml")
set_wandb_job_info(job_info)
# by default, we upload run in ngc and slurm