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})

{''.join(rows)}
" + ) + 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