diff --git a/.file_mapping.json b/.file_mapping.json
index 9cbaf068d..b96ddff35 100644
--- a/.file_mapping.json
+++ b/.file_mapping.json
@@ -1,7 +1,7 @@
{
- "_source_commit": "4d9b6cfd3731fdbbca883184937d40a64d2dca52-dirty",
- "_dest_commit": "c23e51f2f157ae3e51cfcd86ebfb5464850894f2",
- "_generated_at": "2026-09-20T05:50:07Z",
+ "_source_commit": "461431bae40d4a8ccf27d812fc9d1757fe8d3b96-dirty",
+ "_dest_commit": "0460be81f16883aa380e716dc6f58c1189481172",
+ "_generated_at": "2026-09-23T04:47:38Z",
"files": {
"imaginaire/__init__.py": "cosmos_framework/__init__.py",
"imaginaire/attention/__init__.py": "cosmos_framework/model/attention/__init__.py",
@@ -169,6 +169,7 @@
"imaginaire/utils/one_logger/one_logger_global_vars.py": "cosmos_framework/utils/one_logger/one_logger_global_vars.py",
"imaginaire/utils/one_logger/one_logger_override_utils.py": "cosmos_framework/utils/one_logger/one_logger_override_utils.py",
"imaginaire/utils/one_logger/one_logger_utils.py": "cosmos_framework/utils/one_logger/one_logger_utils.py",
+ "imaginaire/utils/one_logger/one_logger_utils_test.py": "cosmos_framework/utils/one_logger/one_logger_utils_test.py",
"imaginaire/utils/optim_instantiate.py": "cosmos_framework/utils/optim_instantiate.py",
"imaginaire/utils/profiling.py": "cosmos_framework/utils/profiling.py",
"imaginaire/utils/progress_bar.py": "cosmos_framework/utils/progress_bar.py",
@@ -292,20 +293,6 @@
"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",
@@ -326,7 +313,9 @@
"projects/cosmos3/cosmos3/datasets/augmentors/reasoner/format_hot_fixes.py": "cosmos_framework/data/generator/augmentors/reasoner/format_hot_fixes.py",
"projects/cosmos3/cosmos3/datasets/augmentors/reasoner/nvlm_data_to_conversation.py": "cosmos_framework/data/generator/augmentors/reasoner/nvlm_data_to_conversation.py",
"projects/cosmos3/cosmos3/datasets/augmentors/reasoner/prompt_format.py": "cosmos_framework/data/generator/augmentors/reasoner/prompt_format.py",
+ "projects/cosmos3/cosmos3/datasets/augmentors/reasoner/prompt_format_test.py": "cosmos_framework/data/generator/augmentors/reasoner/prompt_format_test.py",
"projects/cosmos3/cosmos3/datasets/augmentors/reasoner/shuffle_text_media_order.py": "cosmos_framework/data/generator/augmentors/reasoner/shuffle_text_media_order.py",
+ "projects/cosmos3/cosmos3/datasets/augmentors/reasoner/source_timestamps_test.py": "cosmos_framework/data/generator/augmentors/reasoner/source_timestamps_test.py",
"projects/cosmos3/cosmos3/datasets/augmentors/reasoner/timestamp.py": "cosmos_framework/data/generator/augmentors/reasoner/timestamp.py",
"projects/cosmos3/cosmos3/datasets/augmentors/reasoner/timestamp_test.py": "cosmos_framework/data/generator/augmentors/reasoner/timestamp_test.py",
"projects/cosmos3/cosmos3/datasets/augmentors/reasoner/timestamp_with_subject_tracking.py": "cosmos_framework/data/generator/augmentors/reasoner/timestamp_with_subject_tracking.py",
@@ -488,9 +477,11 @@
"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/VIDEO_TIMESTAMPS.md": "cosmos_framework/data/generator/processors/VIDEO_TIMESTAMPS.md",
"projects/cosmos3/cosmos3/processors/__init__.py": "cosmos_framework/data/generator/processors/__init__.py",
"projects/cosmos3/cosmos3/processors/audio_utils.py": "cosmos_framework/data/generator/processors/audio_utils.py",
"projects/cosmos3/cosmos3/processors/base.py": "cosmos_framework/data/generator/processors/base.py",
+ "projects/cosmos3/cosmos3/processors/base_test.py": "cosmos_framework/data/generator/processors/base_test.py",
"projects/cosmos3/cosmos3/processors/cosmos3_edge_processing.py": "cosmos_framework/data/generator/processors/cosmos3_edge_processing.py",
"projects/cosmos3/cosmos3/processors/cosmos3_edge_processing_test.py": "cosmos_framework/data/generator/processors/cosmos3_edge_processing_test.py",
"projects/cosmos3/cosmos3/processors/nemotron3densevl_processor.py": "cosmos_framework/data/generator/processors/nemotron3densevl_processor.py",
@@ -504,6 +495,7 @@
"projects/cosmos3/cosmos3/processors/qwen3vl_nemo_chat_processor.py": "cosmos_framework/data/generator/processors/qwen3vl_nemo_chat_processor.py",
"projects/cosmos3/cosmos3/processors/qwen3vl_nemo_chat_processor_test.py": "cosmos_framework/data/generator/processors/qwen3vl_nemo_chat_processor_test.py",
"projects/cosmos3/cosmos3/processors/qwen3vl_processor.py": "cosmos_framework/data/generator/processors/qwen3vl_processor.py",
+ "projects/cosmos3/cosmos3/processors/source_video_timing_test.py": "cosmos_framework/data/generator/processors/source_video_timing_test.py",
"projects/cosmos3/cosmos3/scripts/multiview_auto/multiview_collage.py": "cosmos_framework/scripts/multiview_collage.py",
"projects/cosmos3/cosmos3/scripts/multiview_auto/multiview_collage_test.py": "cosmos_framework/scripts/multiview_collage_test.py",
"projects/cosmos3/cosmos3/sequence_packing/__init__.py": "cosmos_framework/data/generator/sequence_packing/__init__.py",
@@ -576,8 +568,10 @@
"projects/cosmos3/cosmos3/utils/reasoner/pretrained_models_downloader.py": "cosmos_framework/utils/generator/reasoner/pretrained_models_downloader.py",
"projects/cosmos3/cosmos3/utils/reasoner/pretrained_models_downloader_test.py": "cosmos_framework/utils/generator/reasoner/pretrained_models_downloader_test.py",
"projects/cosmos3/cosmos3/utils/reasoner/true_packing.py": "cosmos_framework/utils/generator/reasoner/true_packing.py",
+ "projects/cosmos3/cosmos3/utils/source_video_timing.py": "cosmos_framework/utils/generator/source_video_timing.py",
"projects/cosmos3/cosmos3/utils/video_frame_sampling.py": "cosmos_framework/utils/generator/video_frame_sampling.py",
"projects/cosmos3/cosmos3/utils/video_preprocess.py": "cosmos_framework/utils/generator/video_preprocess.py",
+ "projects/cosmos3/cosmos3/utils/video_source_metadata.py": "cosmos_framework/utils/generator/video_source_metadata.py",
"projects/cosmos3/interactive/configs/defaults/flex_attention.py": "cosmos_framework/configs/base/defaults/causal_flex_attention.py",
"projects/cosmos3/interactive/configs/defaults/replay_attention.py": "cosmos_framework/configs/base/defaults/replay_attention.py",
"projects/cosmos3/interactive/models/attention_io_layout.py": "cosmos_framework/model/generator/attention_io_layout.py",
diff --git a/cosmos_framework/configs/base/defaults/multiview_attention.py b/cosmos_framework/configs/base/defaults/multiview_attention.py
index 3b3c7cb55..89b71d2fb 100644
--- a/cosmos_framework/configs/base/defaults/multiview_attention.py
+++ b/cosmos_framework/configs/base/defaults/multiview_attention.py
@@ -78,8 +78,11 @@ def resolve_caption_scope(access: CaptionAccess, *, per_view_captions: bool) ->
# ``decomposed_temporal_window_seconds`` is set: the two streams do not share a frame
# index, but they do share real capture time, which the window compares instead.
#
-# Read by the ``flex_*`` backends only. The ``"maskless"`` backend is its own attention pattern
-# and does not take a scope -- see ``BackendPreference``.
+# Read by every backend, but not the same way. A ``flex_*`` backend expresses the scope as a mask.
+# The ``"maskless"`` backend expresses ``"same_view"`` and ``"decomposed"`` as partitions of the
+# GEN stream and refuses ``"all_views"``, which is not a partition at all -- so there the scope
+# decides whether the cross-instant pass exists rather than describing one attention two ways. See
+# ``BackendPreference`` and ``models.mot.multiview_maskless_attention.MASKLESS_ATTENTION_SCOPES``.
AttentionScope = Literal["all_views", "same_view", "decomposed"]
# The scopes of ``AttentionScope`` at runtime, which the annotation itself is not.
diff --git a/cosmos_framework/configs/base/reasoner/defaults/augmentors.py b/cosmos_framework/configs/base/reasoner/defaults/augmentors.py
index 24553ad23..bd498701f 100644
--- a/cosmos_framework/configs/base/reasoner/defaults/augmentors.py
+++ b/cosmos_framework/configs/base/reasoner/defaults/augmentors.py
@@ -37,6 +37,7 @@ def create_data_augmentor_config() -> dict[str, Any]:
max_fps_thres=60,
target_fps="${data_setting.qwen_target_fps}", # type: ignore
video_temporal_mode="${data_setting.qwen_video_temporal_mode}",
+ video_timestamp_mode="${data_setting.video_timestamp_mode}",
max_video_token_length="${data_setting.qwen_max_video_token_length}", # type: ignore
processor=processor,
extract_audio="${model.config.sound_und}",
@@ -45,6 +46,7 @@ def create_data_augmentor_config() -> dict[str, Any]:
"prompt_format": L(PromptFormat)( # takes text_keys and output "conversation"
input_keys=["texts"],
text_chat_order="${data_setting.text_chat_order}",
+ strip_thinking_prob="${data_setting.strip_thinking_prob}",
),
"shuffle_text_media_order": L(ShuffleTextMediaOrder)(),
"format_hot_fixes": L(FormatHotFixes)(),
@@ -130,6 +132,7 @@ def create_data_augmentor_config() -> dict[str, Any]:
custom_system_prompt="${data_setting.custom_system_prompt}",
strip_original_system_prompt="${data_setting.strip_original_system_prompt}",
video_temporal_mode="${data_setting.qwen_video_temporal_mode}",
+ video_timestamp_mode="${data_setting.video_timestamp_mode}",
text_only=False,
sound_und="${model.config.sound_und}",
audio_encoder_type="${model.config.sound_und_config.audio_encoder_type}",
@@ -168,6 +171,7 @@ def create_data_augmentor_config() -> dict[str, Any]:
processor = L(build_processor_lazy)(
tokenizer_type="${model.config.policy.backbone.model_name}",
+ use_native_edge_processor="${data_setting.use_native_edge_processor}",
credentials="${checkpoint.load_from_object_store.credentials}",
bucket="${checkpoint.load_from_object_store.bucket}",
)
diff --git a/cosmos_framework/configs/base/reasoner/defaults/config.py b/cosmos_framework/configs/base/reasoner/defaults/config.py
index 8d7524d78..dc0bec2a1 100644
--- a/cosmos_framework/configs/base/reasoner/defaults/config.py
+++ b/cosmos_framework/configs/base/reasoner/defaults/config.py
@@ -17,6 +17,7 @@ class DataSetting:
qwen_max_video_token_length: Maximum video token length.
qwen_target_fps: Target fps for video sampling.
text_chat_order: Order of text items in user messages.
+ strip_thinking_prob: Per-sample probability of converting thinking data into non-thinking data.
custom_system_prompt: System prompt injected when a conversation has no leading system message.
strip_original_system_prompt: Remove existing system messages before optional custom prompt injection.
distributor_type: "with_replace" (WeightedShardlistBasic) or "no_replace" (NoReplaceShardlistBasic).
@@ -30,6 +31,11 @@ class DataSetting:
qwen_max_video_token_length: int = 8192
qwen_max_image_token_length: int = 8192
qwen_target_fps: float = 4.0
+ use_native_edge_processor: bool = False
+ video_timestamp_mode: str = attrs.field(
+ default="qwen_index",
+ validator=attrs.validators.in_({"qwen_index", "legacy_fps", "source_pts"}),
+ )
qwen_video_temporal_mode: str = attrs.field(
default="native", validator=attrs.validators.in_({"native", "framewise"})
)
@@ -38,6 +44,14 @@ class DataSetting:
default="text_end",
validator=attrs.validators.in_({"text_end", "text_start", "random"}),
)
+ strip_thinking_prob: float = attrs.field(
+ default=0.0,
+ validator=attrs.validators.and_(
+ attrs.validators.instance_of((int, float)),
+ attrs.validators.ge(0.0),
+ attrs.validators.le(1.0),
+ ),
+ )
custom_system_prompt: str | None = "You are a helpful assistant."
strip_original_system_prompt: bool = False
temporal_localization_output_format: str = attrs.field(
diff --git a/cosmos_framework/data/generator/action/utils/domain_utils.py b/cosmos_framework/data/generator/action/utils/domain_utils.py
index ad56547ff..4d00f36ed 100644
--- a/cosmos_framework/data/generator/action/utils/domain_utils.py
+++ b/cosmos_framework/data/generator/action/utils/domain_utils.py
@@ -54,6 +54,10 @@
# RoboCasa PandaOmron mobile manipulation (10/15/20D raw action per
# ``use_base_action`` / ``base_encoding``); appended above the maximum.
"robocasa": 30,
+ # embodiment_b nvidia-20260828 ingestion: a new one-shot dataset, distinct from
+ # "embodiment_b" (domain 9, an earlier unrelated sample drop with its own 30D
+ # contract).
+ "embodiment_b_20260828": 32,
}
@@ -88,6 +92,7 @@
"so101-bimanual-midtrain-conditional": 20,
"geniesim3_g2a": 29,
"geniesim3_g2a_joint": 16,
+ "embodiment_b_20260828": 50,
# NOTE: ``libero`` (7/10/13 depending on ``rotation_space``), ``hand_pose``
# (variable with ``keypoint_option`` and ``rotation_format``) and ``robocasa``
# (10 arm-only, 15/20 with the mobile base, per ``use_base_action`` /
diff --git a/cosmos_framework/data/generator/augmentor_provider.py b/cosmos_framework/data/generator/augmentor_provider.py
index 7fc757b65..d2d6b6e4a 100644
--- a/cosmos_framework/data/generator/augmentor_provider.py
+++ b/cosmos_framework/data/generator/augmentor_provider.py
@@ -1543,112 +1543,3 @@ def _insert_before(augmentors: dict, anchor_keys: tuple[str, ...], new_key: str,
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
deleted file mode 100644
index b4cb9e0f3..000000000
--- a/cosmos_framework/data/generator/augmentors/hr_lr_degradation/__init__.py
+++ /dev/null
@@ -1,35 +0,0 @@
-# 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
deleted file mode 100644
index 0abcd8c08..000000000
--- a/cosmos_framework/data/generator/augmentors/hr_lr_degradation/augmentor.py
+++ /dev/null
@@ -1,236 +0,0 @@
-# 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
deleted file mode 100644
index 2a272e50b..000000000
--- a/cosmos_framework/data/generator/augmentors/hr_lr_degradation/augmentor_test.py
+++ /dev/null
@@ -1,298 +0,0 @@
-# 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
deleted file mode 100644
index c51c653ba..000000000
--- a/cosmos_framework/data/generator/augmentors/hr_lr_degradation/bench.py
+++ /dev/null
@@ -1,160 +0,0 @@
-# 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
deleted file mode 100644
index 762b75df4..000000000
--- a/cosmos_framework/data/generator/augmentors/hr_lr_degradation/codec.py
+++ /dev/null
@@ -1,219 +0,0 @@
-# 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
deleted file mode 100644
index d353d9600..000000000
--- a/cosmos_framework/data/generator/augmentors/hr_lr_degradation/contact_sheet.py
+++ /dev/null
@@ -1,145 +0,0 @@
-# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
-# SPDX-License-Identifier: OpenMDW-1.1
-"""E4: visual sanity sheet. For each input clip or image and each profile, render first / middle / last
-frames of HR and LR plus a 4x zoom crop into one HTML page, with the degradation record alongside.
-
-Inputs are video files (decoded with PyAV) or images (PNG/JPEG). Example::
-
- PYTHONPATH=. python -m cosmos_framework.data.generator.augmentors.hr_lr_degradation.contact_sheet \
- --inputs clips/*.mp4 --profiles p0_clean_bicubic p1_first_order p1_second_order p3_video_codec \
- --max-frames 33 --out sheet/index.html
-"""
-
-from __future__ import annotations
-
-import argparse
-import base64
-import html
-import io
-import json
-import sys
-from pathlib import Path
-
-import av
-import numpy as np
-import torch
-from PIL import Image
-
-from cosmos_framework.data.generator.augmentors.hr_lr_degradation.degrade import degrade_hr_to_lr
-
-_IMAGE_SUFFIXES = {".png", ".jpg", ".jpeg", ".webp", ".bmp"}
-
-
-def load_media(path: Path, max_frames: int, max_side: int | None) -> torch.Tensor: # returns [3,T,H,W] uint8
- if path.suffix.lower() in _IMAGE_SUFFIXES:
- img = Image.open(path).convert("RGB")
- if max_side and max(img.size) > max_side:
- scale = max_side / max(img.size)
- img = img.resize((round(img.width * scale), round(img.height * scale)), Image.LANCZOS)
- arr = torch.from_numpy(np.array(img)).permute(2, 0, 1) # [3,H,W]
- return arr.unsqueeze(1) # [3,1,H,W]
- frames: list[np.ndarray] = []
- with av.open(str(path)) as container:
- stream = container.streams.video[0]
- stream.thread_type = "AUTO"
- stream.thread_count = 4 # bounded: decoders otherwise size their pool to the whole host
- for frame in container.decode(stream):
- frames.append(frame.reformat(format="rgb24", threads=4).to_ndarray()) # [H,W,3], bounded swscale pool
- if len(frames) >= max_frames:
- break
- if not frames:
- raise ValueError(f"No frames decoded from {path}")
- clip = torch.from_numpy(np.stack(frames)).permute(3, 0, 1, 2) # [3,T,H,W]
- if max_side and max(clip.shape[-2:]) > max_side:
- scale = max_side / max(clip.shape[-2:])
- size = (round(clip.shape[-2] * scale), round(clip.shape[-1] * scale))
- clip = torch.nn.functional.interpolate(
- clip.permute(1, 0, 2, 3).float(), size=size, mode="bicubic", antialias=True, align_corners=False
- ) # [T,3,h,w]
- clip = clip.clamp(0, 255).round().to(torch.uint8).permute(1, 0, 2, 3) # [3,T,h,w]
- return clip
-
-
-def _to_png_b64(frame: torch.Tensor, scale: float = 1.0) -> str: # frame: [3,H,W] uint8
- img = Image.fromarray(frame.permute(1, 2, 0).numpy())
- if scale != 1.0:
- img = img.resize((round(img.width * scale), round(img.height * scale)), Image.NEAREST)
- buf = io.BytesIO()
- img.save(buf, format="PNG")
- return base64.b64encode(buf.getvalue()).decode()
-
-
-def _crop(frame: torch.Tensor, frac: float, cy: float, cx: float) -> torch.Tensor: # frame: [3,H,W]
- h, w = frame.shape[-2:]
- ch, cw = max(8, int(h * frac)), max(8, int(w * frac))
- top = min(max(0, int(cy * h - ch / 2)), h - ch)
- left = min(max(0, int(cx * w - cw / 2)), w - cw)
- return frame[:, top : top + ch, left : left + cw]
-
-
-def render_row(name: str, hr: torch.Tensor, profile: str, seed: int, scale: float, display_width: int) -> str:
- result = degrade_hr_to_lr(hr, profile, scale=scale, seed=seed)
- lr = result.lr # [3,T,h,w]
- t = hr.shape[1]
- idxs = sorted({0, t // 2, t - 1})
- cells: list[str] = []
- for i in idxs:
- hr_f, lr_f = hr[:, i], lr[:, i] # [3,H,W], [3,h,w]
- disp_hr = display_width / hr_f.shape[-1]
- disp_lr = display_width / lr_f.shape[-1] # LR is shown upscaled to the same width for comparison
- crop_hr = _crop(hr_f, 0.15, 0.5, 0.5)
- crop_lr = _crop(lr_f, 0.15, 0.5, 0.5)
- zoom_hr = display_width / 2 / crop_hr.shape[-1]
- zoom_lr = display_width / 2 / crop_lr.shape[-1]
- cells.append(
- f"
frame {i}: HR {hr_f.shape[-2]}x{hr_f.shape[-1]} "
- f" "
- f"LR {lr_f.shape[-2]}x{lr_f.shape[-1]} (shown at HR width) "
- f" "
- f"centre crop, HR | LR "
- f" "
- f"}) | "
- )
- record = html.escape(json.dumps(result.record, indent=1))
- return (
- f"{html.escape(name)} {html.escape(profile)} seed {seed} | "
- + "".join(cells)
- + f"{record} |
"
- )
-
-
-def main(argv: list[str] | None = None) -> int:
- parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
- parser.add_argument("--inputs", nargs="+", required=True, help="video or image files")
- parser.add_argument("--profiles", nargs="+", default=["p0_clean_bicubic", "p1_first_order", "p1_second_order"])
- parser.add_argument("--scale", type=float, default=2.0)
- parser.add_argument("--seed", type=int, default=0)
- parser.add_argument("--max-frames", type=int, default=33)
- parser.add_argument("--max-side", type=int, default=None, help="downscale inputs whose longest side exceeds this")
- parser.add_argument("--display-width", type=int, default=480)
- parser.add_argument("--out", type=str, required=True)
- args = parser.parse_args(argv)
-
- rows: list[str] = []
- for path_str in args.inputs:
- path = Path(path_str)
- hr = load_media(path, args.max_frames, args.max_side) # [3,T,H,W]
- for k, profile in enumerate(args.profiles):
- rows.append(render_row(path.name, hr, profile, args.seed + k, args.scale, args.display_width))
- print(f"rendered {path.name} / {profile}", flush=True)
- page = (
- ""
- f"HR-to-LR degradation contact sheet (scale {args.scale})
"
- )
- out = Path(args.out)
- out.parent.mkdir(parents=True, exist_ok=True)
- out.write_text(page)
- print(f"wrote {out}")
- return 0
-
-
-if __name__ == "__main__":
- sys.exit(main())
diff --git a/cosmos_framework/data/generator/augmentors/hr_lr_degradation/contact_sheet_test.py b/cosmos_framework/data/generator/augmentors/hr_lr_degradation/contact_sheet_test.py
deleted file mode 100644
index 118d70355..000000000
--- a/cosmos_framework/data/generator/augmentors/hr_lr_degradation/contact_sheet_test.py
+++ /dev/null
@@ -1,59 +0,0 @@
-# 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
deleted file mode 100644
index 229bbbaf0..000000000
--- a/cosmos_framework/data/generator/augmentors/hr_lr_degradation/degrade.py
+++ /dev/null
@@ -1,456 +0,0 @@
-# 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
deleted file mode 100644
index 208cdfccd..000000000
--- a/cosmos_framework/data/generator/augmentors/hr_lr_degradation/degrade_test.py
+++ /dev/null
@@ -1,650 +0,0 @@
-# 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
deleted file mode 100644
index fdd3edaab..000000000
--- a/cosmos_framework/data/generator/augmentors/hr_lr_degradation/diffjpeg.py
+++ /dev/null
@@ -1,163 +0,0 @@
-# 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
deleted file mode 100644
index 6c744557d..000000000
--- a/cosmos_framework/data/generator/augmentors/hr_lr_degradation/kernels.py
+++ /dev/null
@@ -1,182 +0,0 @@
-# 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
deleted file mode 100644
index bc4a234ba..000000000
--- a/cosmos_framework/data/generator/augmentors/hr_lr_degradation/ops.py
+++ /dev/null
@@ -1,213 +0,0 @@
-# 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
deleted file mode 100644
index 9219772a7..000000000
--- a/cosmos_framework/data/generator/augmentors/hr_lr_degradation/packing_test.py
+++ /dev/null
@@ -1,65 +0,0 @@
-# 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
deleted file mode 100644
index d00d21d9c..000000000
--- a/cosmos_framework/data/generator/augmentors/hr_lr_degradation/profiles.py
+++ /dev/null
@@ -1,392 +0,0 @@
-# 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/reasoner/bytes_to_media.py b/cosmos_framework/data/generator/augmentors/reasoner/bytes_to_media.py
index e0545a6b1..5d9d6fdf0 100644
--- a/cosmos_framework/data/generator/augmentors/reasoner/bytes_to_media.py
+++ b/cosmos_framework/data/generator/augmentors/reasoner/bytes_to_media.py
@@ -21,7 +21,14 @@
from cosmos_framework.utils import log
from cosmos_framework.data.generator.reasoner.video_decoder_qwen import VideoTemporalMode, _video_decoder_qwen_func
from cosmos_framework.data.generator.processors.qwen3vl_processor import Qwen3VLProcessor
+from cosmos_framework.utils.generator.source_video_timing import (
+ SOURCE_VIDEO_TIMING_KEY,
+ require_source_pts_processor,
+ validate_source_video_timing,
+ validate_video_timestamp_mode,
+)
from cosmos_framework.utils.generator.video_preprocess import tensor_to_pil_images
+from cosmos_framework.utils.generator.video_source_metadata import VIDEO_METADATA_KEY
class BytesToMedia(Augmentor):
@@ -55,6 +62,7 @@ def __init__(
processor: Qwen3VLProcessor = None,
extract_audio: bool = False,
audio_sample_rate: int = 16_000,
+ video_timestamp_mode: str = "qwen_index",
video_temporal_mode: VideoTemporalMode = "native",
) -> None:
"""
@@ -74,6 +82,12 @@ def __init__(
extract_audio (bool): Whether to decode the audio stream from video containers.
audio_sample_rate (int): Target sample rate for decoded mono audio.
"""
+ validate_video_timestamp_mode(video_timestamp_mode)
+ self.video_timestamp_mode: str = video_timestamp_mode
+ if video_timestamp_mode == "source_pts":
+ require_source_pts_processor(processor)
+ if extract_audio:
+ raise ValueError("source_pts does not support audio extraction")
self.input_key = input_key
self.output_key = output_key
if video_temporal_mode not in ("native", "framewise"):
@@ -94,6 +108,7 @@ def __init__(
self.processor = processor
self.extract_audio = extract_audio
self.audio_sample_rate = audio_sample_rate
+ self.video_decoder_params["video_timestamp_mode"] = video_timestamp_mode
def _is_video_key(self, name: str) -> bool:
"""Returns whether the media key will be decoded as video."""
@@ -236,11 +251,17 @@ def _bytes_to_video_frames(
),
)
if result is None:
+ if self.video_timestamp_mode == "source_pts":
+ raise ValueError("source_pts decoder returned no frames or timing record")
log.warning(f"Skipping item '{identifier}': Video decoder returned None.")
return None
+ if self.video_timestamp_mode == "source_pts":
+ validate_source_video_timing(result.get(SOURCE_VIDEO_TIMING_KEY), result["videos"].shape[1])
result["videos"] = tensor_to_pil_images(result["videos"]) # 3,T,H,W -> list of PIL images
return result
except Exception as e:
+ if self.video_timestamp_mode == "source_pts":
+ raise ValueError(f"source_pts failed to decode and align video {identifier!r}") from e
log.warning(f"Skipping item '{identifier}': Error decoding video bytes: {e}")
return None
@@ -303,6 +324,8 @@ def __call__(self, data_dict: Dict) -> Dict:
output_data = {}
if isinstance(data, dict):
+ if self.video_timestamp_mode == "source_pts" and any(self._is_audio_key(name) for name in data):
+ raise ValueError("source_pts does not support audio media")
video_count = sum(1 for name, item in data.items() if isinstance(item, bytes) and self._is_video_key(name))
video_durations = self._get_video_durations(data, data_dict) if video_count > 1 else {}
total_video_duration = sum(video_durations.values()) if video_durations else None
@@ -352,6 +375,8 @@ def __call__(self, data_dict: Dict) -> Dict:
)
if audio is not None:
result["audio"] = audio
+ if VIDEO_METADATA_KEY in result:
+ result["audio_start_seconds"] = start_frame / result[VIDEO_METADATA_KEY]["fps"]
output_data[name] = result
elif (
diff --git a/cosmos_framework/data/generator/augmentors/reasoner/prompt_format.py b/cosmos_framework/data/generator/augmentors/reasoner/prompt_format.py
index 717803705..631d01d58 100644
--- a/cosmos_framework/data/generator/augmentors/reasoner/prompt_format.py
+++ b/cosmos_framework/data/generator/augmentors/reasoner/prompt_format.py
@@ -4,26 +4,34 @@
"""Visual-Text Transformations or Augmentations."""
import random
-from typing import Dict, Literal
+import re
+from typing import Any, Literal
from cosmos_framework.data.imaginaire.webdataset.augmentors.augmentor import Augmentor
+_THINK_RE = re.compile(r".*?\s*|.*\Z", re.DOTALL)
+
class PromptFormat(Augmentor):
def __init__(
self,
- input_keys: list = ["texts"],
+ input_keys: list[str] = ["texts"],
text_chat_order: Literal["text_end", "text_start", "random"] = "text_end",
+ strip_thinking_prob: float = 0.0,
) -> None:
"""
Args:
- input_keys (list): List of input keys.
- text_chat_order (Literal["text_end", "text_start", "random"]): Order of text items in user messages.
+ input_keys: List of input keys.
+ text_chat_order: Order of text items in user messages.
+ strip_thinking_prob: Per-sample probability of dropping assistant thinking traces.
"""
+ if not 0.0 <= strip_thinking_prob <= 1.0:
+ raise ValueError(f"strip_thinking_prob must be in [0, 1], got {strip_thinking_prob}")
self.input_keys = input_keys
self.text_chat_order = text_chat_order
+ self.strip_thinking_prob = strip_thinking_prob
- def __call__(self, data_dict: Dict) -> Dict:
+ def __call__(self, data_dict: dict[str, Any]) -> dict[str, Any] | None:
conversation_key = self.input_keys[0]
# retrive conversations from dict
@@ -56,16 +64,48 @@ def __call__(self, data_dict: Dict) -> Dict:
if "reasoning_content" in message and isinstance(message["reasoning_content"], str):
message["reasoning_content"] = [{"type": "text", "text": message["reasoning_content"]}]
- # Merge reasoning_content into assistant message content
- for message in selected_conversation:
- if message.get("role") == "assistant" and message.get("reasoning_content"):
- # Wrap reasoning items in ... tags
- reasoning_items = message["reasoning_content"]
- think_start = [{"type": "text", "text": "\n"}]
- think_end = [{"type": "text", "text": "\n\n\n"}]
- message["content"] = think_start + reasoning_items + think_end + message["content"]
- del message["reasoning_content"]
-
+ is_thinking_stripped = False
+ if random.random() < self.strip_thinking_prob:
+ for message in selected_conversation:
+ if message.get("role") != "assistant":
+ continue
+ if message.pop("reasoning_content", None):
+ is_thinking_stripped = True
+ content = message.get("content", [])
+ for item in content:
+ if not isinstance(item, dict) or item.get("type") != "text":
+ continue
+ text = item.get("text")
+ if not isinstance(text, str):
+ continue
+ stripped_text = _THINK_RE.sub("", text).lstrip()
+ if stripped_text != text:
+ is_thinking_stripped = True
+ item["text"] = stripped_text
+ has_text = any(
+ isinstance(item, dict)
+ and item.get("type") == "text"
+ and isinstance(item.get("text"), str)
+ and item["text"].strip()
+ for item in content
+ )
+ has_media = any(
+ isinstance(item, dict) and item.get("type") in ("image", "video", "audio") for item in content
+ )
+ if not has_text and not has_media:
+ return None
+ else:
+ # Merge reasoning_content into assistant message content
+ for message in selected_conversation:
+ if message.get("role") == "assistant" and message.get("reasoning_content"):
+ # Wrap reasoning items in ... tags
+ reasoning_items = message["reasoning_content"]
+ think_start = [{"type": "text", "text": "\n"}]
+ think_end = [{"type": "text", "text": "\n\n\n"}]
+ message["content"] = think_start + reasoning_items + think_end + message["content"]
+ del message["reasoning_content"]
+
+ data_dict["is_thinking_stripped"] = is_thinking_stripped
data_dict["conversation"] = selected_conversation
del data_dict[conversation_key]
@@ -75,7 +115,7 @@ def __call__(self, data_dict: Dict) -> Dict:
return data_dict
- def _enforce_text_chat_order(self, conversation: list) -> None:
+ def _enforce_text_chat_order(self, conversation: list[dict[str, Any]]) -> None:
"""
Reorder text content within user messages based on text_chat_order setting.
NOTE (maxzhaoshuol): this does NOT work for interleaved data!!!!!!
diff --git a/cosmos_framework/data/generator/augmentors/reasoner/prompt_format_test.py b/cosmos_framework/data/generator/augmentors/reasoner/prompt_format_test.py
new file mode 100644
index 000000000..9e89da57c
--- /dev/null
+++ b/cosmos_framework/data/generator/augmentors/reasoner/prompt_format_test.py
@@ -0,0 +1,75 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+
+import pytest
+
+from cosmos_framework.data.generator.augmentors.reasoner.prompt_format import PromptFormat
+
+pytestmark = [pytest.mark.L0, pytest.mark.CPU]
+
+
+def test_strip_thinking_removes_reasoning_and_inline_trace() -> None:
+ formatter = PromptFormat(strip_thinking_prob=1.0)
+
+ result = formatter(
+ {
+ "texts": [
+ {"role": "user", "content": "Question"},
+ {
+ "role": "assistant",
+ "reasoning_content": "Hidden reasoning",
+ "content": "Inline reasoning\nFinal answer",
+ },
+ ]
+ }
+ )
+
+ assert result is not None
+ assert result["is_thinking_stripped"] is True
+ assert result["conversation"][0]["content"] == [{"type": "text", "text": "Question"}]
+ assert result["conversation"][1]["content"] == [{"type": "text", "text": "Final answer"}]
+ assert "reasoning_content" not in result["conversation"][1]
+
+
+def test_zero_probability_preserves_thinking_without_prompt_injection() -> None:
+ formatter = PromptFormat(strip_thinking_prob=0.0)
+
+ result = formatter(
+ {
+ "texts": [
+ {"role": "user", "content": "Question"},
+ {"role": "assistant", "reasoning_content": "Reasoning", "content": "Final answer"},
+ ]
+ }
+ )
+
+ assert result is not None
+ assert result["is_thinking_stripped"] is False
+ assert result["conversation"][0]["content"] == [{"type": "text", "text": "Question"}]
+ assert result["conversation"][1]["content"] == [
+ {"type": "text", "text": "\n"},
+ {"type": "text", "text": "Reasoning"},
+ {"type": "text", "text": "\n\n\n"},
+ {"type": "text", "text": "Final answer"},
+ ]
+
+
+def test_strip_thinking_drops_sample_without_assistant_supervision() -> None:
+ formatter = PromptFormat(strip_thinking_prob=1.0)
+
+ result = formatter(
+ {
+ "texts": [
+ {"role": "user", "content": "Question"},
+ {"role": "assistant", "reasoning_content": "Only reasoning", "content": ""},
+ ]
+ }
+ )
+
+ assert result is None
+
+
+@pytest.mark.parametrize("probability", [-0.1, 1.1])
+def test_strip_thinking_rejects_invalid_probability(probability: float) -> None:
+ with pytest.raises(ValueError, match=r"must be in \[0, 1\]"):
+ PromptFormat(strip_thinking_prob=probability)
diff --git a/cosmos_framework/data/generator/augmentors/reasoner/source_timestamps_test.py b/cosmos_framework/data/generator/augmentors/reasoner/source_timestamps_test.py
new file mode 100644
index 000000000..5df1e78b1
--- /dev/null
+++ b/cosmos_framework/data/generator/augmentors/reasoner/source_timestamps_test.py
@@ -0,0 +1,243 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: OpenMDW-1.1
+
+"""Training labels and paired cropped audio share the video source clock."""
+
+from types import SimpleNamespace
+from typing import Any
+
+import numpy as np
+import pytest
+from PIL import Image
+
+from cosmos_framework.data.generator.augmentors.reasoner.timestamp import overlay_text
+from cosmos_framework.data.generator.augmentors.reasoner.tokenize_data import TokenizeData
+from cosmos_framework.data.generator.augmentors.reasoner.tokenize_data_test import (
+ _FakeAudioProcessor,
+ _FakeVLMProcessor,
+)
+from cosmos_framework.data.generator.processors.audio_utils import (
+ AUDIO_END_TOKEN,
+ AUDIO_PAD_TOKEN,
+ AUDIO_START_TOKEN,
+ get_audio_segment_token_lengths,
+)
+from cosmos_framework.data.generator.processors.base import maybe_parse_video_content
+from cosmos_framework.utils.generator.video_source_metadata import calculate_video_timestamps
+
+pytestmark = [pytest.mark.L1, pytest.mark.CPU]
+
+
+@pytest.mark.parametrize("temporal_patch_size,expected", [(1, [1.0, 1.3, 1.7]), (2, [1.2, 1.2, 1.7])])
+def test_overlay_uses_temporal_patch_size_not_spatial_merge(temporal_patch_size: int, expected: list[float]) -> None:
+ frames = [Image.new("RGB", (64, 64))] * 3
+ processor = SimpleNamespace(
+ name="/local/edge", temporal_patch_size=temporal_patch_size, merge_size=4, USES_SOURCE_VIDEO_TIMESTAMPS=True
+ )
+ result, times = overlay_text(
+ frames,
+ 4.0,
+ processor=processor,
+ video_metadata={"fps": 30.0, "total_num_frames": 100, "frames_indices": [30, 40, 50]},
+ )
+ assert result is frames
+ assert times == expected
+
+
+def test_cropped_audio_uses_same_origin_and_token_partition() -> None:
+ processor = _FakeVLMProcessor()
+ tokenize = TokenizeData(
+ processor=processor,
+ sound_und=True,
+ audio_processor=_FakeAudioProcessor(token_lengths=(10,), timestamp_stride=0.11),
+ )
+ metadata = {"fps": 10.0, "total_num_frames": 100, "frames_indices": [50, 52, 56, 58]}
+ data = {
+ "__key__": "crop",
+ "__url__": SimpleNamespace(root="test", path="crop"),
+ "media": {
+ "video": {
+ "videos": [Image.new("RGB", (32, 32))] * 4,
+ "fps": 4.0,
+ "video_metadata": metadata,
+ "audio": np.zeros(16000, dtype=np.float32),
+ "audio_start_seconds": 5.0,
+ }
+ },
+ "conversation": [
+ {"role": "user", "content": [{"type": "video", "video": "video"}, {"type": "audio", "audio": "video"}]},
+ {"role": "assistant", "content": [{"type": "text", "text": "answer"}]},
+ ],
+ }
+ assert tokenize(data) is not None
+ content = processor.last_conversation[0]["content"]
+ assert content[0]["video_metadata"] == metadata
+ assert (
+ content[1]["text"]
+ == f"{AUDIO_START_TOKEN}<5.1 seconds>{AUDIO_PAD_TOKEN * 4}<5.7 seconds>{AUDIO_PAD_TOKEN * 6}{AUDIO_END_TOKEN}"
+ )
+
+
+@pytest.mark.parametrize(
+ "metadata,mode",
+ [
+ (None, "qwen_index"),
+ ({"fps": 30.0, "total_num_frames": 100, "frames_indices": [0]}, "qwen_index"),
+ ({"fps": 30.0, "total_num_frames": 100, "frames_indices": [0, 1, 2, 3]}, "legacy_fps"),
+ ],
+)
+def test_tokenize_rejects_invalid_explicit_source_metadata(metadata: object, mode: str) -> None:
+ processor = _FakeVLMProcessor()
+ data = {
+ "__key__": "invalid",
+ "__url__": SimpleNamespace(root="test", path="invalid"),
+ "media": {"video": {"videos": [Image.new("RGB", (32, 32))] * 4, "fps": 4.0, "video_metadata": metadata}},
+ "conversation": [{"role": "user", "content": [{"type": "video", "video": "video"}]}],
+ }
+ with pytest.raises(ValueError, match="video_metadata"):
+ TokenizeData(processor=processor, video_timestamp_mode=mode)(data)
+
+
+def test_repeated_source_frames_partition_audio_without_losing_tokens() -> None:
+ assert get_audio_segment_token_lengths(5, [1.0, 1.0, 1.2], audio_token_timestamps=[1.0, 1.05, 1.1, 1.15, 1.2]) == [
+ 1,
+ 2,
+ 2,
+ ]
+ with pytest.raises(ValueError, match="nondecreasing"):
+ get_audio_segment_token_lengths(5, [1.2, 1.0])
+
+
+def _multi_video_audio_sample(content_order: list[tuple[str, str]]) -> dict[str, Any]:
+ media = {}
+ for key, start, value in (("video_a", 50, 2.0), ("video_b", 200, 7.0)):
+ media[key] = {
+ "videos": [Image.new("RGB", (32, 32))] * 4,
+ "fps": 4.0,
+ "video_metadata": {
+ "fps": 10.0,
+ "total_num_frames": 300,
+ "frames_indices": [start + index for index in (0, 2, 6, 8)],
+ },
+ "audio": np.full(320, value, dtype=np.float32),
+ "audio_start_seconds": start / 10.0,
+ }
+ return {
+ "__key__": "multi-video-crops",
+ "__url__": SimpleNamespace(root="test", path="multi-video-crops"),
+ "media": media,
+ "conversation": [
+ {"role": "user", "content": [{"type": kind, kind: key} for kind, key in content_order]},
+ {"role": "assistant", "content": [{"type": "text", "text": "answer"}]},
+ ],
+ }
+
+
+def test_multi_video_audio_uses_its_own_source_clock() -> None:
+ processor = _FakeVLMProcessor()
+ data = _multi_video_audio_sample(
+ [("video", "video_a"), ("video", "video_b"), ("audio", "video_a"), ("audio", "video_b")]
+ )
+ output = TokenizeData(
+ processor=processor,
+ sound_und=True,
+ audio_processor=_FakeAudioProcessor(token_lengths=(10, 10), timestamp_stride=0.11),
+ )(data)
+ assert output is not None
+ content = processor.last_conversation[0]["content"]
+ for item, start in zip(content[2:], (5, 20), strict=True):
+ assert item["text"] == (
+ f"{AUDIO_START_TOKEN}<{start}.1 seconds>{AUDIO_PAD_TOKEN * 4}"
+ f"<{start}.7 seconds>{AUDIO_PAD_TOKEN * 6}{AUDIO_END_TOKEN}"
+ )
+ assert output["audio_features"][:, 0, 0].tolist() == [2.0, 7.0]
+
+
+@pytest.mark.parametrize("audio_layout", ["separate_with_timestamps", "separate_no_timestamps", "interleaved_av"])
+def test_audio_cannot_borrow_another_video_clock_before_its_own_video(audio_layout: str) -> None:
+ processor = _FakeVLMProcessor()
+ data = _multi_video_audio_sample([("video", "video_b"), ("audio", "video_a"), ("video", "video_a")])
+ output = TokenizeData(
+ processor=processor,
+ sound_und=True,
+ audio_layout=audio_layout,
+ audio_processor=_FakeAudioProcessor(token_lengths=(10,)),
+ )(data)
+ assert output is None
+ assert processor.last_conversation is None
+
+
+def test_interleaved_audio_cannot_attach_to_an_unrelated_adjacent_video() -> None:
+ processor = _FakeVLMProcessor()
+ data = _multi_video_audio_sample([("video", "video_a"), ("video", "video_b"), ("audio", "video_a")])
+ output = TokenizeData(
+ processor=processor,
+ sound_und=True,
+ audio_layout="interleaved_av",
+ audio_processor=_FakeAudioProcessor(token_lengths=(10,)),
+ )(data)
+ assert output is None
+ assert processor.last_conversation is None
+
+
+def test_interleaved_cropped_pairs_keep_each_clock_and_audio_feature_order() -> None:
+ processor = _FakeVLMProcessor()
+ data = _multi_video_audio_sample(
+ [("video", "video_a"), ("audio", "video_a"), ("video", "video_b"), ("audio", "video_b")]
+ )
+ output = TokenizeData(
+ processor=processor,
+ sound_und=True,
+ audio_layout="interleaved_av",
+ audio_processor=_FakeAudioProcessor(token_lengths=(10, 10), timestamp_stride=0.11),
+ )(data)
+ assert output is not None
+ tokens = processor.tokenizer.vocabulary
+ chunk = [tokens[key] for key in ("", "<|vision_start|>", "