From bd6ec1c10ddf9fc460fc6158efba772b18e9e8de Mon Sep 17 00:00:00 2001 From: Codex Date: Mon, 21 Sep 2026 18:39:16 +0800 Subject: [PATCH 1/4] feat(training): decouple Nano Reasoner with offline K/V conditioning --- cosmos_framework/callbacks/load_pretrained.py | 36 + cosmos_framework/checkpoint/reasoner_only.py | 210 +++ .../checkpoint/reasoner_only_test.py | 226 +++ .../configs/base/defaults/model_config.py | 51 +- .../configs/toml_config/sft_config.py | 86 +- .../configs/toml_config/sft_config_test.py | 96 + .../configs/toml_config/toml_config_helper.py | 32 +- .../generator/local_datasets/sft_dataset.py | 340 +++- .../sft_dataset_caption_test.py | 18 + .../local_datasets/sft_reasoner_documents.py | 195 ++ .../sft_reasoner_documents_test.py | 352 ++++ .../generator/mot/cosmos3_vfm_network.py | 39 +- ...mos3_vfm_network_reasoner_features_test.py | 52 + .../generator/mot/include_gen_pathway_test.py | 190 +- .../model/generator/mot/unified_mot.py | 295 ++- .../model/generator/omni_mot_causal_model.py | 33 +- .../model/generator/omni_mot_model.py | 393 +++- .../omni_mot_reasoner_conditioning_test.py | 390 ++++ .../model/generator/reasoner_feature_cache.py | 1663 +++++++++++++++++ .../generator/reasoner_feature_cache_test.py | 596 ++++++ .../model/generator/reasoner_features.py | 870 +++++++++ .../model/generator/reasoner_features_test.py | 501 +++++ .../scripts/extract_reasoner_features.py | 669 +++++++ .../scripts/extract_reasoner_features_test.py | 267 +++ docs/nano_sft_decoupled_reasoner.md | 670 +++++++ 25 files changed, 8019 insertions(+), 251 deletions(-) create mode 100644 cosmos_framework/checkpoint/reasoner_only.py create mode 100644 cosmos_framework/checkpoint/reasoner_only_test.py create mode 100644 cosmos_framework/data/generator/local_datasets/sft_reasoner_documents.py create mode 100644 cosmos_framework/data/generator/local_datasets/sft_reasoner_documents_test.py create mode 100644 cosmos_framework/model/generator/mot/cosmos3_vfm_network_reasoner_features_test.py create mode 100644 cosmos_framework/model/generator/omni_mot_reasoner_conditioning_test.py create mode 100644 cosmos_framework/model/generator/reasoner_feature_cache.py create mode 100644 cosmos_framework/model/generator/reasoner_feature_cache_test.py create mode 100644 cosmos_framework/model/generator/reasoner_features.py create mode 100644 cosmos_framework/model/generator/reasoner_features_test.py create mode 100644 cosmos_framework/scripts/extract_reasoner_features.py create mode 100644 cosmos_framework/scripts/extract_reasoner_features_test.py create mode 100644 docs/nano_sft_decoupled_reasoner.md diff --git a/cosmos_framework/callbacks/load_pretrained.py b/cosmos_framework/callbacks/load_pretrained.py index 45f81d8c0..d6267e301 100644 --- a/cosmos_framework/callbacks/load_pretrained.py +++ b/cosmos_framework/callbacks/load_pretrained.py @@ -5,6 +5,27 @@ from cosmos_framework.utils.callback import Callback +def _warm_start_skips_complete_ema(patterns: list[str], ema_state_fqns: list[str] | None = None) -> bool: + """Return whether DCP substring filters cover every ``net_ema.*`` state leaf.""" + ema_root_fqn = "net_ema." + if any(pattern and pattern in ema_root_fqn for pattern in patterns): + return True + if not ema_state_fqns: + return False + return all(any(pattern and pattern in fqn for pattern in patterns) for fqn in ema_state_fqns) + + +def _warm_start_partially_skips_ema( + patterns: list[str], + ema_state_fqns: list[str], +) -> bool: + """Return whether DCP filters match some, but not all, EMA state leaves.""" + if not patterns or not ema_state_fqns or _warm_start_skips_complete_ema(patterns, ema_state_fqns): + return False + matched = sum(any(pattern and pattern in fqn for pattern in patterns) for fqn in ema_state_fqns) + return 0 < matched < len(ema_state_fqns) + + class LoadPretrained(Callback): """Load HF understanding-pathway weights after DCP resume, gated by checkpoint state. @@ -23,7 +44,22 @@ def on_train_start(self, model: ImaginaireModel, iteration: int = 0) -> None: from cosmos_framework.checkpoint.dcp import DistributedCheckpointer probe = DistributedCheckpointer(self.config.checkpoint, self.config.job, callbacks=None, disable_async=True) + # DCP matches these patterns as substrings of flattened FQNs. Only a + # pattern broad enough to match the EMA root proves that the complete + # warm-start EMA subtree was skipped; partial EMA skips must not trigger + # an unconditional regular->EMA overwrite. + skip_patterns = self.config.checkpoint.keys_to_skip_loading + ema_state_fqns: list[str] = [] + net_ema = getattr(model, "net_ema", None) + if net_ema is not None: + ema_state_fqns.extend(f"net_ema.{name}" for name, _ in net_ema.named_parameters()) + ema_state_fqns.extend(f"net_ema.{name}" for name, _ in net_ema.named_buffers()) + warm_start_ema_skipped = _warm_start_skips_complete_ema(skip_patterns, ema_state_fqns) + warm_start_ema_partially_skipped = _warm_start_partially_skips_ema(skip_patterns, ema_state_fqns) model.load_pretrained_model_if_needed( has_resumable_checkpoint=probe.has_resumable_checkpoint(), has_load_path=probe.load_path is not None, + warm_start_ema_skipped=warm_start_ema_skipped, + warm_start_ema_partially_skipped=warm_start_ema_partially_skipped, + warm_start_strict_resume=self.config.checkpoint.strict_resume, ) diff --git a/cosmos_framework/checkpoint/reasoner_only.py b/cosmos_framework/checkpoint/reasoner_only.py new file mode 100644 index 000000000..2d55ee37e --- /dev/null +++ b/cosmos_framework/checkpoint/reasoner_only.py @@ -0,0 +1,210 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +"""Load only the Reasoner subtree from a Cosmos training DCP checkpoint. + +The training checkpoint stores the two model copies under distinct roots:: + + net.language_model.* + net_ema.language_model.* + +This module maps exactly one of those roots onto an already-instantiated +Reasoner module. It deliberately does not construct ``OmniMoTModel`` (and +therefore never constructs the Generator, VAE, or EMA model). +""" + +from __future__ import annotations + +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Literal, cast + +import torch +import torch.distributed.checkpoint as dcp +from torch import nn +from torch.distributed.checkpoint import FileSystemReader +from torch.distributed.checkpoint.state_dict import ( + StateDictOptions, + get_model_state_dict, + set_model_state_dict, +) + +ReasonerCheckpointSource = Literal["regular", "ema"] + +_SOURCE_PREFIXES: dict[ReasonerCheckpointSource, str] = { + "regular": "net.language_model.", + "ema": "net_ema.language_model.", +} + + +@dataclass(frozen=True) +class ReasonerCheckpointLoadResult: + """Summary of one successful Reasoner-only load.""" + + checkpoint_path: Path + source: ReasonerCheckpointSource + checkpoint_prefix: str + num_state_leaves: int + + +def _resolve_dcp_model_path(checkpoint_path: str | Path) -> Path: + """Resolve either a DCP model-component directory or its iteration root.""" + raw_path = str(checkpoint_path) + if "://" in raw_path: + raise ValueError(f"Reasoner-only checkpoint loading currently supports local DCP paths only; got {raw_path!r}") + + path = Path(checkpoint_path) + direct_metadata = path / ".metadata" + nested_metadata = path / "model" / ".metadata" + if direct_metadata.is_file(): + return path + if nested_metadata.is_file(): + return path / "model" + raise FileNotFoundError(f"Could not find DCP metadata. Expected either {direct_metadata} or {nested_metadata}.") + + +def _checkpoint_prefix(source: str) -> tuple[ReasonerCheckpointSource, str]: + if source not in _SOURCE_PREFIXES: + raise ValueError(f"source must be explicitly set to 'regular' or 'ema', got {source!r}") + typed_source = cast(ReasonerCheckpointSource, source) + return typed_source, _SOURCE_PREFIXES[typed_source] + + +def _is_generation_pathway_fqn(fqn: str) -> bool: + """Whether an MoT language-model FQN belongs exclusively to the Generator tower.""" + return "moe_gen" in fqn + + +def _is_visual_pathway_fqn(fqn: str) -> bool: + """Whether an MoT language-model FQN belongs to the optional visual tower.""" + return fqn == "visual" or fqn.startswith("visual.") + + +def _reasoner_target_state(reasoner: nn.Module) -> dict[str, Any]: + state = dict( + get_model_state_dict( + reasoner, + options=StateDictOptions(strict=True), + ) + ) + if not state: + raise ValueError("Reasoner module has an empty state dict") + + generation_leaves = sorted(name for name in state if _is_generation_pathway_fqn(name)) + if generation_leaves: + raise ValueError( + "Reasoner-only load target still contains Generator pathway state. Construct the language model " + f"with include_gen_pathway=False; generator leaves={generation_leaves}" + ) + + unsupported = [name for name, value in state.items() if not torch.is_tensor(value)] + if unsupported: + raise TypeError( + "Reasoner-only DCP loading currently supports tensor state leaves only; " + f"non-tensor leaves={sorted(unsupported)}" + ) + meta_leaves = [ + name for name, value in state.items() if isinstance(value, torch.Tensor) and value.device.type == "meta" + ] + if meta_leaves: + raise ValueError(f"Reasoner must be materialized before loading; meta state leaves={sorted(meta_leaves)}") + return state + + +def load_reasoner_only_dcp( + reasoner: nn.Module, + checkpoint_path: str | Path, + *, + source: ReasonerCheckpointSource, +) -> ReasonerCheckpointLoadResult: + """Load one complete Reasoner subtree from a local training DCP checkpoint. + + Args: + reasoner: The already-instantiated language-model/Reasoner module. It is + the load target itself, not an Omni model or a ``net`` wrapper. + checkpoint_path: Either the DCP ``model/`` component directory or its + parent iteration directory. + source: Explicitly select ``"regular"`` (``net.language_model.*``) or + ``"ema"`` (``net_ema.language_model.*``). There is no implicit + fallback between the two sources. + + Raises: + ValueError: If the source is invalid, the target is unmaterialized, or + the selected checkpoint subtree differs from the target state. + FileNotFoundError: If no local DCP metadata is present. + + The key-set comparison happens before any tensor is read, so a missing or + unexpected Reasoner leaf cannot produce a partially initialized model. + Shape compatibility remains enforced by PyTorch's strict DCP planner. + """ + typed_source, prefix = _checkpoint_prefix(source) + model_path = _resolve_dcp_model_path(checkpoint_path) + target_state = _reasoner_target_state(reasoner) + + reader = FileSystemReader(str(model_path)) + metadata = reader.read_metadata() + checkpoint_keys = set(metadata.state_dict_metadata) + selected_keys = {key for key in checkpoint_keys if key.startswith(prefix)} + if not selected_keys: + raise KeyError( + f"DCP checkpoint {model_path} contains no Reasoner state under the explicitly selected prefix {prefix!r}" + ) + + expected_keys = {f"{prefix}{name}" for name in target_state} + missing = sorted(expected_keys - selected_keys) + # A full MoT language-model checkpoint contains both the frozen Reasoner + # and ``*_moe_gen`` Generator tower. Shipped Nano checkpoints also contain + # the optional visual encoder. Text-only cache extraction deliberately + # constructs neither, so those source-only leaves are safe to omit. Visual + # extras are allowed only when the target has no visual module at all; any + # other extra leaf signals an architecture/config mismatch. + allow_source_visual = not hasattr(reasoner, "visual") + unexpected = sorted( + key + for key in selected_keys - expected_keys + if not ( + _is_generation_pathway_fqn(key.removeprefix(prefix)) + or (allow_source_visual and _is_visual_pathway_fqn(key.removeprefix(prefix))) + ) + ) + if missing or unexpected: + raise ValueError( + f"Reasoner checkpoint/target FQN mismatch under {prefix!r}: missing={missing}, unexpected={unexpected}" + ) + + prefixed_target_state = {f"{prefix}{name}": value for name, value in target_state.items()} + dcp.load( + state_dict=prefixed_target_state, + storage_reader=reader, + planner=dcp.DefaultLoadPlanner(allow_partial_load=False), + # Every extraction worker owns a complete Reasoner replica. Independent + # reads avoid nested collectives and collective-order deadlocks when one + # worker encounters an I/O or materialization error. + no_dist=True, + ) + + loaded_state = {name: prefixed_target_state[f"{prefix}{name}"] for name in target_state} + incompatible = set_model_state_dict( + reasoner, + model_state_dict=loaded_state, + options=StateDictOptions(strict=True), + ) + if incompatible.missing_keys or incompatible.unexpected_keys: + raise RuntimeError( + "Strict Reasoner state installation unexpectedly reported incompatible keys: " + f"missing={incompatible.missing_keys}, unexpected={incompatible.unexpected_keys}" + ) + + return ReasonerCheckpointLoadResult( + checkpoint_path=model_path, + source=typed_source, + checkpoint_prefix=prefix, + num_state_leaves=len(target_state), + ) + + +__all__ = [ + "ReasonerCheckpointLoadResult", + "ReasonerCheckpointSource", + "load_reasoner_only_dcp", +] diff --git a/cosmos_framework/checkpoint/reasoner_only_test.py b/cosmos_framework/checkpoint/reasoner_only_test.py new file mode 100644 index 000000000..9b8618a99 --- /dev/null +++ b/cosmos_framework/checkpoint/reasoner_only_test.py @@ -0,0 +1,226 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +from __future__ import annotations + +from pathlib import Path +from typing import cast + +import pytest +import torch +import torch.distributed.checkpoint as dcp +from torch import nn +from torch.distributed.checkpoint.state_dict import ( + StateDictOptions, + get_model_state_dict, + set_model_state_dict, +) + +from cosmos_framework.checkpoint.reasoner_only import ( + ReasonerCheckpointSource, + load_reasoner_only_dcp, +) + + +class _TinyReasoner(nn.Module): + def __init__(self) -> None: + super().__init__() + self.proj = nn.Linear(3, 4) + self.norm = nn.LayerNorm(4) + self.register_buffer("persistent_scale", torch.ones(1)) + + +class _FullMoTLanguageModel(_TinyReasoner): + """Training shape: Reasoner state plus source-only Generator/visual state.""" + + def __init__(self) -> None: + super().__init__() + self.proj_moe_gen = nn.Linear(3, 4) + self.norm_moe_gen = nn.LayerNorm(4) + self.visual = nn.Linear(6, 7) + + +class _TrainingPath(nn.Module): + def __init__(self) -> None: + super().__init__() + self.language_model = _FullMoTLanguageModel() + self.generator = nn.Linear(5, 6) + self.vae = nn.Linear(7, 8) + + +class _TrainingCheckpointModel(nn.Module): + def __init__(self) -> None: + super().__init__() + self.net = _TrainingPath() + self.net_ema = _TrainingPath() + + +def _fill_module(module: nn.Module, value: float) -> None: + with torch.no_grad(): + for tensor in module.state_dict().values(): + if tensor.is_floating_point(): + tensor.fill_(value) + + +def _save_training_checkpoint(path: Path, model: nn.Module) -> None: + state = get_model_state_dict(model, options=StateDictOptions(strict=True)) + dcp.save(state_dict=state, checkpoint_id=path) + + +def _full_dcp_load(path: Path, model: nn.Module) -> None: + state = get_model_state_dict(model, options=StateDictOptions(strict=True)) + dcp.load(state_dict=state, checkpoint_id=path) + incompatible = set_model_state_dict( + model, + model_state_dict=state, + options=StateDictOptions(strict=True), + ) + assert incompatible.missing_keys == [] + assert incompatible.unexpected_keys == [] + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +@pytest.mark.parametrize( + ("source", "expected_prefix", "expected_value"), + [ + ("regular", "net.language_model.", 1.0), + ("ema", "net_ema.language_model.", 2.0), + ], +) +def test_reasoner_only_dcp_load_matches_full_model_dcp_load( + tmp_path: Path, + source: ReasonerCheckpointSource, + expected_prefix: str, + expected_value: float, +) -> None: + source_model = _TrainingCheckpointModel() + _fill_module(source_model.net, 1.0) + _fill_module(source_model.net_ema, 2.0) + checkpoint_path = tmp_path / "model" + _save_training_checkpoint(checkpoint_path, source_model) + + full_target = _TrainingCheckpointModel() + _fill_module(full_target, -1.0) + _full_dcp_load(checkpoint_path, full_target) + full_reasoner = full_target.net.language_model if source == "regular" else full_target.net_ema.language_model + + reasoner_only_target = _TinyReasoner() + _fill_module(reasoner_only_target, -2.0) + result = load_reasoner_only_dcp( + reasoner_only_target, + tmp_path, + source=source, + ) + + assert result.checkpoint_path == checkpoint_path + assert result.source == source + assert result.checkpoint_prefix == expected_prefix + assert result.num_state_leaves == len(reasoner_only_target.state_dict()) + for name, actual in reasoner_only_target.state_dict().items(): + torch.testing.assert_close(actual, full_reasoner.state_dict()[name], rtol=0, atol=0) + if actual.is_floating_point(): + torch.testing.assert_close(actual, torch.full_like(actual, expected_value), rtol=0, atol=0) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_reasoner_only_dcp_load_rejects_missing_and_unexpected_selected_fqns_before_mutation( + tmp_path: Path, +) -> None: + source = _TinyReasoner() + source_state = source.state_dict() + checkpoint_state = { + f"net.language_model.{name}": tensor.clone() for name, tensor in source_state.items() if name != "norm.bias" + } + checkpoint_state["net.language_model.not_in_target"] = torch.ones(1) + checkpoint_state["net.generator.unrelated"] = torch.full((2,), 99.0) + checkpoint_path = tmp_path / "model" + dcp.save(state_dict=checkpoint_state, checkpoint_id=checkpoint_path) + + target = _TinyReasoner() + _fill_module(target, -3.0) + before = {name: tensor.clone() for name, tensor in target.state_dict().items()} + + with pytest.raises(ValueError, match="Reasoner checkpoint/target FQN mismatch") as error: + load_reasoner_only_dcp(target, checkpoint_path, source="regular") + + assert "net.language_model.norm.bias" in str(error.value) + assert "net.language_model.not_in_target" in str(error.value) + for name, actual in target.state_dict().items(): + torch.testing.assert_close(actual, before[name], rtol=0, atol=0) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_reasoner_only_dcp_load_rejects_target_that_still_contains_generator_tower(tmp_path: Path) -> None: + source_model = _TrainingCheckpointModel() + checkpoint_path = tmp_path / "model" + _save_training_checkpoint(checkpoint_path, source_model) + + with pytest.raises(ValueError, match="include_gen_pathway=False"): + load_reasoner_only_dcp(_FullMoTLanguageModel(), checkpoint_path, source="regular") + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_reasoner_only_dcp_load_never_falls_back_between_regular_and_ema(tmp_path: Path) -> None: + ema_only = { + f"net_ema.language_model.{name}": tensor.clone() for name, tensor in _TinyReasoner().state_dict().items() + } + checkpoint_path = tmp_path / "model" + dcp.save(state_dict=ema_only, checkpoint_id=checkpoint_path) + + with pytest.raises(KeyError, match="net.language_model"): + load_reasoner_only_dcp(_TinyReasoner(), checkpoint_path, source="regular") + + result = load_reasoner_only_dcp(_TinyReasoner(), checkpoint_path, source="ema") + assert result.source == "ema" + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_reasoner_only_dcp_load_rejects_invalid_source_and_non_dcp_path(tmp_path: Path) -> None: + with pytest.raises(ValueError, match="regular.*ema"): + load_reasoner_only_dcp( + _TinyReasoner(), + tmp_path, + source=cast(ReasonerCheckpointSource, "automatic"), + ) + + with pytest.raises(FileNotFoundError, match="Could not find DCP metadata"): + load_reasoner_only_dcp(_TinyReasoner(), tmp_path, source="regular") + + with pytest.raises(ValueError, match="local DCP paths only"): + load_reasoner_only_dcp(_TinyReasoner(), "s3://bucket/model", source="regular") + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_reasoner_only_dcp_load_rejects_meta_target(tmp_path: Path) -> None: + source_model = _TrainingCheckpointModel() + checkpoint_path = tmp_path / "model" + _save_training_checkpoint(checkpoint_path, source_model) + + with torch.device("meta"): + target = _TinyReasoner() + with pytest.raises(ValueError, match="must be materialized"): + load_reasoner_only_dcp(target, checkpoint_path, source="regular") + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_reasoner_only_dcp_load_casts_fp32_master_weights_into_bf16_compute_target(tmp_path: Path) -> None: + source_model = _TrainingCheckpointModel() + _fill_module(source_model.net, 1.25) + checkpoint_path = tmp_path / "model" + _save_training_checkpoint(checkpoint_path, source_model) + + target = _TinyReasoner().to(dtype=torch.bfloat16) + result = load_reasoner_only_dcp(target, checkpoint_path, source="regular") + + assert result.num_state_leaves == len(target.state_dict()) + for tensor in target.state_dict().values(): + if tensor.is_floating_point(): + assert tensor.dtype == torch.bfloat16 + torch.testing.assert_close(tensor, torch.full_like(tensor, 1.25), rtol=0, atol=0) diff --git a/cosmos_framework/configs/base/defaults/model_config.py b/cosmos_framework/configs/base/defaults/model_config.py index 777f8a59f..32899b742 100644 --- a/cosmos_framework/configs/base/defaults/model_config.py +++ b/cosmos_framework/configs/base/defaults/model_config.py @@ -5,7 +5,6 @@ import attrs -from cosmos_framework.utils.lazy_config import LazyDict from cosmos_framework.configs.base.defaults.activation_checkpointing import ActivationCheckpointingConfig from cosmos_framework.configs.base.defaults.compile import CompileConfig from cosmos_framework.configs.base.defaults.ema import EMAConfig @@ -17,6 +16,7 @@ from cosmos_framework.model.generator.mot.action_io_projector import ACTION_IO_PROJECTOR_TYPES from cosmos_framework.model.generator.utils.load_balancing_stats import LBLConfig from cosmos_framework.model.generator.utils.sr_latent_noise import SRLatentConditionNoiseConfig +from cosmos_framework.utils.lazy_config import LazyDict # Mirrors ``cosmos3.common.args.AttentionIOLayout``. Defined locally on purpose: importing # the ``cosmos3`` workspace package at module scope makes the whole cosmos3 config tree @@ -24,6 +24,7 @@ # ``cosmos3``), which breaks the benchmark-request config-check CI job. Keep in sync with # ``packages/cosmos3/cosmos3/common/args.py``. AttentionIOLayout = Literal["sequence_sharded", "replicated"] +ReasonerConditioningBackend = Literal["joint", "inline", "offline", "remote", "read_through"] @attrs.define(slots=False) @@ -153,6 +154,50 @@ class FixedStepSamplerConfig: sample_type: str = "sde" +@attrs.define(slots=False) +class ReasonerConditioningConfig: + """How frozen Reasoner features are supplied to generator SFT. + + ``joint`` preserves the existing one-pass MoT forward. ``inline`` is the + two-pass numerical-reference path that captures Reasoner K/V locally and + then executes the generator-only path. The remaining backends remove the + Reasoner parameters from training ranks and obtain the same per-layer K/V + tensors from storage and/or a dedicated service. + """ + + backend: ReasonerConditioningBackend = attrs.field( + default="joint", + validator=attrs.validators.in_({"joint", "inline", "offline", "remote", "read_through"}), + ) + cache_root: str | None = None + endpoint: str | None = None + reasoner_fingerprint: str | None = None + tokenizer_fingerprint: str | None = None + framing_fingerprint: str | None = None + strict_fingerprint: bool = True + prefetch_batches: int = attrs.field(default=2, validator=attrs.validators.ge(0)) + layerwise_h2d: bool = False + request_timeout_s: float = attrs.field(default=300.0, validator=attrs.validators.gt(0.0)) + + def __attrs_post_init__(self) -> None: + if self.backend in {"offline", "read_through"} and not self.cache_root: + raise ValueError(f"reasoner_conditioning.backend={self.backend!r} requires cache_root") + if self.backend in {"remote", "read_through"} and not self.endpoint: + raise ValueError(f"reasoner_conditioning.backend={self.backend!r} requires endpoint") + if self.backend in {"offline", "remote", "read_through"} and not self.strict_fingerprint: + raise ValueError(f"reasoner_conditioning.backend={self.backend!r} requires strict_fingerprint=true") + if self.backend in {"offline", "remote", "read_through"}: + missing = [ + name + for name in ("reasoner_fingerprint", "tokenizer_fingerprint", "framing_fingerprint") + if not getattr(self, name) + ] + if missing: + raise ValueError( + f"reasoner_conditioning.backend={self.backend!r} with strict_fingerprint=true requires {missing}" + ) + + # Don't have any defaults and init only in config file. @attrs.define(slots=False) class OmniMoTModelConfig: @@ -240,6 +285,10 @@ class OmniMoTModelConfig: # Optional fixed-step sampler for distilled models (None for base models). fixed_step_sampler_config: FixedStepSamplerConfig | None = None + # Frozen Reasoner decoupling. The default keeps the existing joint model + # and is intentionally a no-op for every current recipe. + reasoner_conditioning: ReasonerConditioningConfig = ReasonerConditioningConfig() + # Model configs vlm_config: VLMConfig = VLMConfig() diffusion_expert_config: DiffusionExpertConfig = DiffusionExpertConfig() diff --git a/cosmos_framework/configs/toml_config/sft_config.py b/cosmos_framework/configs/toml_config/sft_config.py index 04d0efbc5..6cdf8abdd 100644 --- a/cosmos_framework/configs/toml_config/sft_config.py +++ b/cosmos_framework/configs/toml_config/sft_config.py @@ -11,10 +11,10 @@ from __future__ import annotations from pathlib import Path -from typing import Any, Optional +from typing import Any, Literal, Optional import tomllib -from pydantic import BaseModel, ConfigDict, Field +from pydantic import BaseModel, ConfigDict, Field, model_validator from cosmos_framework.configs.toml_config.toml_config_helper import ( TASK_TO_BASE_CONFIG, @@ -245,6 +245,54 @@ class ActivationCheckpointingConfig(BaseModel): ) +class ReasonerConditioningConfig(BaseModel): + """Frozen Reasoner feature source for VFM generator SFT.""" + + model_config = _PYDANTIC_MODEL_CONFIG + + backend: Literal["joint", "inline", "offline", "remote", "read_through"] = Field( + default="joint", + description=( + "'joint' preserves the current dual-path forward; 'inline' captures Reasoner K/V locally for parity; " + "'offline' reads a precomputed cache; 'remote' calls dedicated Reasoner workers; 'read_through' " + "uses the cache first and sends misses to the service." + ), + ) + cache_root: Optional[str] = Field(default=None, description="Root of the immutable Reasoner K/V cache.") + endpoint: Optional[str] = Field(default=None, description="Reasoner feature-service endpoint.") + reasoner_fingerprint: Optional[str] = Field(default=None, description="Pinned Reasoner checkpoint digest.") + tokenizer_fingerprint: Optional[str] = Field(default=None, description="Pinned tokenizer/special-token digest.") + framing_fingerprint: Optional[str] = Field(default=None, description="Pinned prompt-framing/schema digest.") + strict_fingerprint: bool = Field( + default=True, + description="Reject feature entries whose model/tokenizer/framing fingerprint does not match.", + ) + prefetch_batches: int = Field(default=2, ge=0, description="Number of future batches to prefetch.") + layerwise_h2d: bool = Field( + default=False, + description="Stage K/V one layer at a time instead of keeping the entire batch resident on GPU.", + ) + request_timeout_s: float = Field(default=300.0, gt=0.0, description="Remote request deadline in seconds.") + + @model_validator(mode="after") + def validate_backend_inputs(self) -> "ReasonerConditioningConfig": + if self.backend in {"offline", "read_through"} and not self.cache_root: + raise ValueError(f"backend={self.backend!r} requires cache_root") + if self.backend in {"remote", "read_through"} and not self.endpoint: + raise ValueError(f"backend={self.backend!r} requires endpoint") + if self.backend in {"offline", "remote", "read_through"} and not self.strict_fingerprint: + raise ValueError(f"backend={self.backend!r} requires strict_fingerprint=true") + if self.backend in {"offline", "remote", "read_through"}: + missing = [ + name + for name in ("reasoner_fingerprint", "tokenizer_fingerprint", "framing_fingerprint") + if not getattr(self, name) + ] + if missing: + raise ValueError(f"backend={self.backend!r} with strict_fingerprint=true requires {missing}") + return self + + class ModelTokenizerConfig(BaseModel): """Video tokenizer (VAE) settings. VFM only — VLM skips this sub-tree.""" @@ -357,8 +405,7 @@ class ModelConfig(BaseModel): lora_rank: int = Field( default=16, description=( - "LoRA rank `r`. Adapter shape is (rank × hidden_dim) per target " - "module. Standard values are 4, 8, 16, 32." + "LoRA rank `r`. Adapter shape is (rank × hidden_dim) per target module. Standard values are 4, 8, 16, 32." ), ) lora_alpha: int = Field( @@ -379,9 +426,8 @@ class ModelConfig(BaseModel): ema: EMAConfig = Field(default_factory=EMAConfig) parallelism: ParallelismConfig = Field(default_factory=ParallelismConfig) compile: CompileConfig = Field(default_factory=CompileConfig) - activation_checkpointing: ActivationCheckpointingConfig = Field( - default_factory=ActivationCheckpointingConfig - ) + activation_checkpointing: ActivationCheckpointingConfig = Field(default_factory=ActivationCheckpointingConfig) + reasoner_conditioning: ReasonerConditioningConfig = Field(default_factory=ReasonerConditioningConfig) tokenizer: ModelTokenizerConfig = Field(default_factory=ModelTokenizerConfig) backbone: BackboneConfig = Field(default_factory=BackboneConfig) @@ -473,15 +519,12 @@ class SchedulerConfig(BaseModel): ) f_start: list[float] = Field( default_factory=lambda: [1.0e-6], - description=( - "Initial LR multiplier at step 0, before warmup ramps up." - ), + description=("Initial LR multiplier at step 0, before warmup ramps up."), ) verbosity_interval: int = Field( default=0, description=( - "How often the scheduler logs the current LR (in optimizer " - "steps). 0 = silent. VFM only — skipped on VLM." + "How often the scheduler logs the current LR (in optimizer steps). 0 = silent. VFM only — skipped on VLM." ), ) warm_up_steps: list[int] = Field( @@ -533,8 +576,7 @@ class GradClipCallback(BaseModel): clip_norm: float = Field( default=1.0, description=( - "Maximum global L2 norm of the gradient. Steps with a larger " - "norm are rescaled so ||grad|| ≤ clip_norm." + "Maximum global L2 norm of the gradient. Steps with a larger norm are rescaled so ||grad|| ≤ clip_norm." ), ) force_finite: bool = Field( @@ -567,8 +609,7 @@ class TrainerConfig(BaseModel): distributed_parallelism: str = Field( default="fsdp", description=( - "Distributed strategy. 'fsdp' (the only supported value today) " - "routes through cosmos's FSDP wrapper." + "Distributed strategy. 'fsdp' (the only supported value today) routes through cosmos's FSDP wrapper." ), ) grad_accum_iter: int = Field( @@ -658,10 +699,7 @@ class DataloaderTrainConfig(BaseModel): ) seed: int = Field( default=42, - description=( - "Dataloader RNG seed. Skipped on VLM (CosmosDataLoader has " - "no seed ctor kwarg there)." - ), + description=("Dataloader RNG seed. Skipped on VLM (CosmosDataLoader has no seed ctor kwarg there)."), ) @@ -746,8 +784,7 @@ def load_experiment_from_toml( base_config_path = TASK_TO_BASE_CONFIG[task] except KeyError as e: raise ValueError( - f"{toml_path}: [job].task={task!r} is not supported. " - f"Valid values: {sorted(TASK_TO_BASE_CONFIG)}" + f"{toml_path}: [job].task={task!r} is not supported. Valid values: {sorted(TASK_TO_BASE_CONFIG)}" ) from e overrides = build_hydra_overrides(raw) @@ -759,10 +796,7 @@ def load_experiment_from_toml( if not o or o == "--": continue if "=" not in o: - raise ValueError( - f"extra override {o!r} must be Hydra dotted-path syntax " - f"(e.g. 'optimizer.lr=1e-5')." - ) + raise ValueError(f"extra override {o!r} must be Hydra dotted-path syntax (e.g. 'optimizer.lr=1e-5').") overrides.append(o) # Import lazily so this module stays cheap to import in non-training contexts. diff --git a/cosmos_framework/configs/toml_config/sft_config_test.py b/cosmos_framework/configs/toml_config/sft_config_test.py index 8ff8f3f04..155043e72 100644 --- a/cosmos_framework/configs/toml_config/sft_config_test.py +++ b/cosmos_framework/configs/toml_config/sft_config_test.py @@ -76,6 +76,76 @@ def test_custom_does_not_loosen_sibling_validation(self) -> None: } ) + @pytest.mark.parametrize( + ("backend", "fields"), + [ + ("joint", {}), + ("inline", {}), + ( + "offline", + { + "cache_root": "/features", + "reasoner_fingerprint": "reasoner", + "tokenizer_fingerprint": "tokenizer", + "framing_fingerprint": "framing", + }, + ), + ( + "remote", + { + "endpoint": "dns:///reasoner:50051", + "reasoner_fingerprint": "reasoner", + "tokenizer_fingerprint": "tokenizer", + "framing_fingerprint": "framing", + }, + ), + ( + "read_through", + { + "cache_root": "/features", + "endpoint": "dns:///reasoner:50051", + "reasoner_fingerprint": "reasoner", + "tokenizer_fingerprint": "tokenizer", + "framing_fingerprint": "framing", + }, + ), + ], + ) + def test_reasoner_conditioning_valid_backends(self, backend: str, fields: dict[str, str]) -> None: + raw = { + "job": {"task": "vfm", "experiment": "vision_sft_nano"}, + "model": {"reasoner_conditioning": {"backend": backend, **fields}}, + } + cfg = SFTExperimentConfig.model_validate(raw) + assert cfg.model.reasoner_conditioning.backend == backend + + @pytest.mark.parametrize( + "settings", + [ + {"backend": "offline"}, + {"backend": "remote"}, + {"backend": "read_through", "cache_root": "/features"}, + {"backend": "read_through", "endpoint": "dns:///reasoner:50051"}, + {"backend": "offline", "cache_root": "/features"}, + { + "backend": "offline", + "cache_root": "/features", + "reasoner_fingerprint": "reasoner", + "tokenizer_fingerprint": "tokenizer", + "framing_fingerprint": "framing", + "strict_fingerprint": False, + }, + ], + ) + def test_reasoner_conditioning_requires_backend_inputs(self, settings: dict[str, str]) -> None: + with pytest.raises(ValidationError): + SFTExperimentConfig.model_validate( + { + "job": {"task": "vfm", "experiment": "vision_sft_nano"}, + "model": {"reasoner_conditioning": settings}, + } + ) + # --------------------------------------------------------------------------- # # 2. build_hydra_overrides must NOT emit [custom] as per-leaf overrides # @@ -101,6 +171,32 @@ def test_other_keys_still_emitted(self) -> None: assert "experiment=vision_sft_nano" in overrides assert any(o.startswith("optimizer.lr=") for o in overrides), overrides + def test_reasoner_conditioning_routes_only_to_vfm(self) -> None: + settings = { + "backend": "offline", + "cache_root": "/features", + "strict_fingerprint": True, + "reasoner_fingerprint": "reasoner", + "tokenizer_fingerprint": "tokenizer", + "framing_fingerprint": "framing", + } + vfm = build_hydra_overrides( + { + "job": {"task": "vfm", "experiment": "vision_sft_nano"}, + "model": {"reasoner_conditioning": settings}, + } + ) + assert "model.config.reasoner_conditioning.backend=offline" in vfm + assert "model.config.reasoner_conditioning.cache_root=/features" in vfm + + vlm = build_hydra_overrides( + { + "job": {"task": "vlm", "experiment": "dummy"}, + "model": {"reasoner_conditioning": settings}, + } + ) + assert all("reasoner_conditioning" not in override for override in vlm) + # --------------------------------------------------------------------------- # # 3. end-to-end load_experiment_from_toml on the shipped vision_sft_nano recipe # diff --git a/cosmos_framework/configs/toml_config/toml_config_helper.py b/cosmos_framework/configs/toml_config/toml_config_helper.py index 4d1535c55..24bc7d0e5 100644 --- a/cosmos_framework/configs/toml_config/toml_config_helper.py +++ b/cosmos_framework/configs/toml_config/toml_config_helper.py @@ -18,7 +18,6 @@ from typing import Any - # Maps ``job.task`` to the base Hydra config that ``make_config()`` lives in. TASK_TO_BASE_CONFIG: dict[str, str] = { "vfm": "cosmos_framework/configs/base/config.py", @@ -51,11 +50,16 @@ # not config.job.* — hoist it out of the job section. ("job", "upload_reproducible_setup"): ("upload_reproducible_setup",), ("model", "attn_implementation"): None, - ("model", "backbone"): None, # VLM-only — VFM has no model.config.backbone + ("model", "backbone"): None, # VLM-only — VFM has no model.config.backbone # Per-caption token cap lives on the nested SFT dataset, not a top-level # dataloader scalar — route it to the get_sft_dataset node. ("dataloader_train", "max_caption_tokens"): ( - "dataloader_train", "dataloader", "datasets", "video", "dataset", "max_caption_tokens", + "dataloader_train", + "dataloader", + "datasets", + "video", + "dataset", + "max_caption_tokens", ), ("model",): ("model", "config"), }, @@ -76,11 +80,16 @@ ("model", "lora_rank"): None, ("model", "lora_alpha"): None, ("model", "lora_target_modules"): None, - ("model", "tokenizer"): None, # blocks model.tokenizer.* + ("model", "reasoner_conditioning"): None, + ("model", "tokenizer"): None, # blocks model.tokenizer.* ("dataloader_train", "seed"): None, - ("optimizer", "eps"): None, # VLM_OPTIMIZER_KWARGS has no eps field - ("scheduler", "verbosity_interval"): None, # VLM_LAMBDACOSINE_KWARGS has no verbosity_interval - ("trainer", "callbacks", "compile_tokenizer"): None, # VFM-only callback (VLM has no torch.compile of the tokenizer) + ("optimizer", "eps"): None, # VLM_OPTIMIZER_KWARGS has no eps field + ("scheduler", "verbosity_interval"): None, # VLM_LAMBDACOSINE_KWARGS has no verbosity_interval + ( + "trainer", + "callbacks", + "compile_tokenizer", + ): None, # VFM-only callback (VLM has no torch.compile of the tokenizer) # Rename / re-route to the VLM path ("model", "attn_implementation"): ("model", "config", "policy", "attn_implementation"), ("model", "ema"): ("model", "config", "ema"), @@ -89,7 +98,7 @@ # PoolPackingBatcher (dataloader_train.batcher.*), not flat on the loader. ("dataloader_train", "max_samples_per_batch"): ("dataloader_train", "batcher", "max_batch_size"), ("dataloader_train", "max_sequence_length"): ("dataloader_train", "batcher", "max_tokens"), - ("dataloader_train", "max_caption_tokens"): None, # VFM-only knob — VLM packer caps via max_sequence_length + ("dataloader_train", "max_caption_tokens"): None, # VFM-only knob — VLM packer caps via max_sequence_length # Catch-all for any other model.* sub-keys ("model",): ("model", "config"), }, @@ -136,10 +145,7 @@ def build_hydra_overrides(toml_dict: dict) -> list[str]: overrides.append(f"experiment={experiment_name}") if task not in PATH_REMAPS: - raise ValueError( - f"[job].task={task!r} has no remap rules. " - f"Valid values: {sorted(PATH_REMAPS)}" - ) + raise ValueError(f"[job].task={task!r} has no remap rules. Valid values: {sorted(PATH_REMAPS)}") rules = PATH_REMAPS[task] overlay = dict(toml_dict) @@ -208,5 +214,3 @@ def _hydra_format(v: Any, in_list: bool = False) -> str: return f"'{v}'" return v return str(v) - - diff --git a/cosmos_framework/data/generator/local_datasets/sft_dataset.py b/cosmos_framework/data/generator/local_datasets/sft_dataset.py index 1a9fcdbcf..2b364ccc8 100644 --- a/cosmos_framework/data/generator/local_datasets/sft_dataset.py +++ b/cosmos_framework/data/generator/local_datasets/sft_dataset.py @@ -9,6 +9,8 @@ import os import random import tempfile +from collections.abc import Callable +from dataclasses import dataclass from pathlib import Path from typing import Any, Optional @@ -63,6 +65,68 @@ CAPTION_WEIGHTS = list(CAPTION_TYPES_AND_WEIGHTS.values()) +@dataclass(frozen=True) +class SFTWindowFraming: + """Deterministic frame geometry shared by SFT loading and Reasoner extraction.""" + + sample_key: str + window_index: int + original_fps: float + total_frames: int + start_frame: int + end_frame: int + temporal_interval: int + num_frames: int + target_height: int + target_width: int + + +def sft_metadata_sort_key(metadata: dict[str, Any]) -> str: + """Return the stable ordering key used for SFT metadata.""" + + return hashlib.sha256(metadata["uuid"].encode("utf-8")).hexdigest() + + +def _format_caption(t2w_window: dict, caption_key: str) -> tuple[str, str, bool]: + """Normalize one selected caption exactly as the training path expects.""" + + raw = t2w_window[caption_key] + if isinstance(raw, dict): + return caption_key, caption_json_to_prompt(raw), True + if caption_key == CAPTION_JSON_KEY: + return caption_key, str(raw).strip(), True + return caption_key, raw.strip().rstrip(".") + ".", False + + +def enumerate_sft_captions(t2w_window: dict) -> tuple[tuple[str, str, bool], ...]: + """Return every caption selection reachable by :func:`_select_caption`. + + Priority captions shadow all fallbacks. When training reaches the weighted + caption family, only positive-weight choices are reachable and they are + returned in the same stable order as ``CAPTION_TYPES``. + """ + + for caption_key in (CAPTION_JSON_KEY, "qwen3_32b_rewrite-dense", "caption"): + if caption_key in t2w_window: + return (_format_caption(t2w_window, caption_key),) + + available_types = [ + caption_type + for caption_type in CAPTION_TYPES + if caption_type in t2w_window and CAPTION_TYPES_AND_WEIGHTS[caption_type] > 0 + ] + if available_types: + return tuple(_format_caption(t2w_window, caption_type) for caption_type in available_types) + + zero_weight_types = [caption_type for caption_type in CAPTION_TYPES if caption_type in t2w_window] + if zero_weight_types: + raise ValueError( + "SFT window only contains zero-weight caption types, so training cannot select a caption: " + f"{zero_weight_types}" + ) + return () + + def _select_caption(t2w_window: dict) -> tuple[str, str, bool] | None: """Pick a window's caption: ``(caption_key, caption_text, used_structured_json)``. @@ -74,25 +138,163 @@ def _select_caption(t2w_window: dict) -> tuple[str, str, bool] | None: which would append a stray ``.`` after the closing ``}``. Returns ``None`` when the window has no known caption key. """ - if CAPTION_JSON_KEY in t2w_window: - caption_key = CAPTION_JSON_KEY - elif "qwen3_32b_rewrite-dense" in t2w_window: - caption_key = "qwen3_32b_rewrite-dense" - elif "caption" in t2w_window: - caption_key = "caption" + for caption_key in (CAPTION_JSON_KEY, "qwen3_32b_rewrite-dense", "caption"): + if caption_key in t2w_window: + return _format_caption(t2w_window, caption_key) + + available_types = [caption_type for caption_type in CAPTION_TYPES if caption_type in t2w_window] + if not available_types: + return None + caption_key = random.choices( + available_types, + weights=[CAPTION_TYPES_AND_WEIGHTS[caption_type] for caption_type in available_types], + k=1, + )[0] + return _format_caption(t2w_window, caption_key) + + +def resolve_sft_window_framing( + metadata: dict[str, Any], + window_index: int, + *, + original_fps: float, + total_frames: int, + decoded_total_frames: int | None = None, + num_video_frames: int, + temporal_interval_mode: str, + frame_selection_mode: str, + temporal_compression_factor: int, + target_height: int, + target_width: int, + random_frame_selector: Callable[[int, int], int] | None = None, +) -> SFTWindowFraming | None: + """Resolve one retained SFT window without decoding video pixels. + + ``total_frames`` is the ffprobe count used for the training window + selection. ``decoded_total_frames`` optionally bounds the retained-frame + count to what ffmpeg actually yielded; keeping those values separate avoids + changing fixed-window selection when ffprobe overestimates a damaged/VFR + video. ``None`` mirrors the training path's insufficient/empty-window skip. + A caller using ``frame_selection_mode='random'`` must explicitly provide + the selector; deterministic offline enumeration intentionally does not draw + from global random state. + """ + + if not 0 <= window_index < len(metadata["t2w_windows"]): + raise IndexError(f"window_index={window_index} is out of range for {metadata['uuid']!r}") + if original_fps <= 0: + raise ValueError(f"original_fps must be positive, got {original_fps}") + if temporal_compression_factor < 1: + raise ValueError(f"temporal_compression_factor must be >= 1, got {temporal_compression_factor}") + if total_frames <= 0: + return None + + t2w_window = metadata["t2w_windows"][window_index] + window_start = int(t2w_window["start_frame"]) + window_end = int(t2w_window["end_frame"]) + actual_end = min(window_end, total_frames - 1) + frames_in_window = actual_end - window_start + 1 + if frames_in_window <= 0: + return None + + if num_video_frames == -1: + temporal_interval = int(t2w_window["temporal_interval"]) + start_frame = window_start + end_frame = actual_end else: - available_types = [ct for ct in CAPTION_TYPES if ct in t2w_window] - if not available_types: + if frames_in_window < num_video_frames: return None - available_weights = [CAPTION_TYPES_AND_WEIGHTS[ct] for ct in available_types] - caption_key = random.choices(available_types, weights=available_weights, k=1)[0] + if temporal_interval_mode == "force_one": + temporal_interval = 1 + elif temporal_interval_mode == "max_30fps": + temporal_interval = max(1, int(original_fps / 30.0)) + elif temporal_interval_mode == "entire_chunk": + temporal_interval = max(1, frames_in_window // num_video_frames) + else: + raise ValueError(f"Unknown temporal_interval_mode: {temporal_interval_mode}") + + num_frames_before_downsample = (num_video_frames - 1) * temporal_interval + 1 + if frame_selection_mode == "first": + start_frame = window_start + elif frame_selection_mode == "center": + start_frame = window_start + (frames_in_window - num_frames_before_downsample) // 2 + elif frame_selection_mode == "random": + if random_frame_selector is None: + raise ValueError("frame_selection_mode='random' requires an explicit random_frame_selector") + max_offset = frames_in_window - num_frames_before_downsample + start_frame = window_start + random_frame_selector(0, max(0, max_offset)) + else: + raise ValueError(f"Unknown frame_selection_mode: {frame_selection_mode}") + end_frame = start_frame + num_frames_before_downsample - 1 + + if temporal_interval <= 0: + raise ValueError(f"temporal_interval must be positive, got {temporal_interval}") + decode_frame_count = total_frames if decoded_total_frames is None else decoded_total_frames + if decode_frame_count <= 0: + return None + decode_end = min(end_frame, decode_frame_count - 1) + first_decoded_frame = start_frame + if first_decoded_frame < 0: + first_decoded_frame += (-first_decoded_frame + temporal_interval - 1) // temporal_interval * temporal_interval + if first_decoded_frame > decode_end: + return None + sampled_frames = (decode_end - first_decoded_frame) // temporal_interval + 1 + num_frames = (sampled_frames - 1) // temporal_compression_factor * temporal_compression_factor + 1 + return SFTWindowFraming( + sample_key=f"{metadata['uuid']}_w{window_index}", + window_index=window_index, + original_fps=original_fps, + total_frames=total_frames, + start_frame=start_frame, + end_frame=end_frame, + temporal_interval=temporal_interval, + num_frames=num_frames, + target_height=target_height, + target_width=target_width, + ) - raw = t2w_window[caption_key] - if isinstance(raw, dict): - return caption_key, caption_json_to_prompt(raw), True - if caption_key == CAPTION_JSON_KEY: - return caption_key, str(raw).strip(), True - return caption_key, raw.strip().rstrip(".") + ".", False + +def render_sft_caption( + caption: str, + *, + used_structured_json: bool, + cfg_dropped: bool, + cfg_dropout_keep_metadata: bool, + caption_suffix: str, + append_duration_fps_timestamps: bool, + append_resolution_info: bool, + num_frames: int, + conditioning_fps: float, + target_height: int, + target_width: int, +) -> str: + """Apply the training caption suffix, CFG and metadata framing deterministically.""" + + if caption_suffix and not used_structured_json: + caption = (caption + " " + caption_suffix).strip() + if cfg_dropout_keep_metadata and cfg_dropped: + caption = "" + + if append_duration_fps_timestamps and not used_structured_json: + duration = num_frames / conditioning_fps + caption = caption + " " + _DURATION_TEMPLATE.format(duration=duration, fps=conditioning_fps) + if append_resolution_info and not used_structured_json: + caption = caption + " " + _RESOLUTION_TEMPLATE.format(height=target_height, width=target_width) + caption = caption.strip() + + if not cfg_dropout_keep_metadata and cfg_dropped: + caption = "" + return caption + + +def enumerate_sft_cfg_variants(cfg_dropout_rate: float) -> tuple[bool, ...]: + """Return reachable ``cfg_dropped`` states in deterministic order.""" + + if cfg_dropout_rate <= 0: + return (False,) + if cfg_dropout_rate >= 1: + return (True,) + return False, True class SFTDataset(torch.utils.data.IterableDataset): @@ -184,8 +386,6 @@ def process_one_sample(self, metadata: dict) -> dict | None: windows = metadata["t2w_windows"] win_idx = random.randrange(len(windows)) t2w_window = windows[win_idx] - window_start = t2w_window["start_frame"] - window_end = t2w_window["end_frame"] # Compute output resolution input_w, input_h = metadata["width"], metadata["height"] @@ -207,47 +407,25 @@ def process_one_sample(self, metadata: dict) -> dict | None: video_info = get_video_metadata(input_video_path) original_fps = video_info["fps"] total_frames = video_info["total_frames"] - - # Constrain to the t2w window - actual_end = min(window_end, total_frames - 1) - frames_in_window = actual_end - window_start + 1 - - if self.num_video_frames == -1: - # Native chunk mode: use start/end/interval directly from the window - temporal_interval = t2w_window["temporal_interval"] - start_frame = window_start - end_frame = actual_end - else: - if frames_in_window < self.num_video_frames: - log.warning( - f"Not enough frames in window: {metadata['uuid']}, " - f"frames_in_window: {frames_in_window}, required: {self.num_video_frames}" - ) - return None - - # Compute temporal interval - if self.temporal_interval_mode == "force_one": - temporal_interval = 1 - elif self.temporal_interval_mode == "max_30fps": - temporal_interval = max(1, int(original_fps / 30.0)) - elif self.temporal_interval_mode == "entire_chunk": - temporal_interval = frames_in_window // self.num_video_frames - temporal_interval = max(1, temporal_interval) - else: - raise ValueError(f"Unknown temporal_interval_mode: {self.temporal_interval_mode}") - - num_frames_before_downsample = (self.num_video_frames - 1) * temporal_interval + 1 - if self.frame_selection_mode == "first": - start_frame = window_start - elif self.frame_selection_mode == "center": - start_frame = window_start + (frames_in_window - num_frames_before_downsample) // 2 - elif self.frame_selection_mode == "random": - max_offset = frames_in_window - num_frames_before_downsample - start_frame = window_start + random.randint(0, max(0, max_offset)) - else: - raise ValueError(f"Unknown frame_selection_mode: {self.frame_selection_mode}") - end_frame = start_frame + num_frames_before_downsample - 1 - + framing = resolve_sft_window_framing( + metadata, + win_idx, + original_fps=original_fps, + total_frames=total_frames, + num_video_frames=self.num_video_frames, + temporal_interval_mode=self.temporal_interval_mode, + frame_selection_mode=self.frame_selection_mode, + temporal_compression_factor=self.temporal_compression_factor, + target_height=target_h, + target_width=target_w, + random_frame_selector=random.randint if self.frame_selection_mode == "random" else None, + ) + if framing is None: + log.warning(f"Window is empty or too short: {metadata['uuid']}_w{win_idx}") + return None + temporal_interval = framing.temporal_interval + start_frame = framing.start_frame + end_frame = framing.end_frame fps = original_fps / temporal_interval video_chunk = [] @@ -298,36 +476,24 @@ def process_one_sample(self, metadata: dict) -> dict | None: if self.conditioning_fps_noise_std > 0: noise_factor = np.exp(np.random.randn() * self.conditioning_fps_noise_std) cond_fps = cond_fps * noise_factor - - if self.caption_suffix and not used_structured_json: - caption = (caption + " " + self.caption_suffix).strip() - - # CFG dropout: when cfg_dropout_keep_metadata is True, dropout fires - # before appending resolution/duration/FPS so that metadata text is - # preserved even under unconditional guidance. - if self.cfg_dropout_keep_metadata and self.cfg_dropout_rate > 0: - if random.random() < self.cfg_dropout_rate: - caption = "" - - # Structured-JSON captions already carry duration/fps/resolution inside the - # JSON, so skip the natural-language metadata suffixes for them. This also - # makes the training prompt byte-match the inference prompt. - if self.append_duration_fps_timestamps and not used_structured_json: - duration = num_decoded_frames / cond_fps - suffix = _DURATION_TEMPLATE.format(duration=duration, fps=cond_fps) - caption = caption + " " + suffix - if self.append_resolution_info and not used_structured_json: - suffix = _RESOLUTION_TEMPLATE.format(height=target_h, width=target_w) - caption = caption + " " + suffix - caption = caption.strip() - - if not self.cfg_dropout_keep_metadata and self.cfg_dropout_rate > 0: - if random.random() < self.cfg_dropout_rate: - caption = "" + cfg_dropped = self.cfg_dropout_rate > 0 and random.random() < self.cfg_dropout_rate + caption = render_sft_caption( + caption, + used_structured_json=used_structured_json, + cfg_dropped=cfg_dropped, + cfg_dropout_keep_metadata=self.cfg_dropout_keep_metadata, + caption_suffix=self.caption_suffix, + append_duration_fps_timestamps=self.append_duration_fps_timestamps, + append_resolution_info=self.append_resolution_info, + num_frames=num_decoded_frames, + conditioning_fps=cond_fps, + target_height=target_h, + target_width=target_w, + ) text_ids, caption = self._tokenize_caption(caption) ret = dict( - __key__=f"{metadata['uuid']}_w{win_idx}", + __key__=framing.sample_key, __url__=metadata["vision_path"], fps=original_fps, n_orig_video_frames=total_frames, @@ -688,7 +854,7 @@ def get_sft_dataset( log.info(f"sample_by_window=True: flattened to {len(metadata_list)} samples (one per window)") # Deterministic shuffle based on the sha256 hash of uuid - metadata_list.sort(key=lambda x: hashlib.sha256(x["uuid"].encode("utf-8")).hexdigest()) + metadata_list.sort(key=sft_metadata_sort_key) dataset = SFTDataset( metadata=metadata_list, diff --git a/cosmos_framework/data/generator/local_datasets/sft_dataset_caption_test.py b/cosmos_framework/data/generator/local_datasets/sft_dataset_caption_test.py index 045608835..4ce21c4ce 100644 --- a/cosmos_framework/data/generator/local_datasets/sft_dataset_caption_test.py +++ b/cosmos_framework/data/generator/local_datasets/sft_dataset_caption_test.py @@ -4,6 +4,7 @@ import json +from cosmos_framework.data.generator.local_datasets import sft_dataset as sft_dataset_module from cosmos_framework.data.generator.local_datasets.sft_dataset import _select_caption from cosmos_framework.inference.structured_caption import CAPTION_JSON_KEY @@ -51,5 +52,22 @@ def test_weighted_caption_types_fallback(): assert text.endswith(".") +def test_weighted_selection_only_formats_the_chosen_caption(monkeypatch): + monkeypatch.setattr( + sft_dataset_module.random, + "choices", + lambda *_args, **_kwargs: ["qwen3_235b_dense"], + ) + + key, text, used_json = _select_caption( + { + "qwen3_235b_dense": "valid caption", + "qwen3_32b_dense": None, + } + ) + + assert (key, text, used_json) == ("qwen3_235b_dense", "valid caption.", False) + + def test_no_known_caption_key_returns_none(): assert _select_caption({"start_frame": 0, "end_frame": 84}) is None diff --git a/cosmos_framework/data/generator/local_datasets/sft_reasoner_documents.py b/cosmos_framework/data/generator/local_datasets/sft_reasoner_documents.py new file mode 100644 index 000000000..e148bbc40 --- /dev/null +++ b/cosmos_framework/data/generator/local_datasets/sft_reasoner_documents.py @@ -0,0 +1,195 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +"""Finite, deterministic Reasoner documents derived from the generator SFT dataset.""" + +import tempfile +from collections.abc import Callable, Iterator, Mapping +from dataclasses import dataclass +from typing import Any + +import boto3 + +from cosmos_framework.data.generator.local_datasets.helper import ( + client_config, + download_from_s3, + ffmpeg_decode_video, + get_video_metadata, +) +from cosmos_framework.data.generator.local_datasets.sft_dataset import ( + SFTDataset, + SFTWindowFraming, + enumerate_sft_captions, + enumerate_sft_cfg_variants, + render_sft_caption, + resolve_sft_window_framing, + sft_metadata_sort_key, +) +from cosmos_framework.utils import log + +SFTVideoInfoResolver = Callable[[dict[str, Any]], Mapping[str, Any]] + + +@dataclass(frozen=True) +class SFTReasonerDocument: + """One tokenized SFT prompt variant before BOS/EOS/mRoPE sequence packing.""" + + sample_key: str + vision_path: str + window_index: int + caption_key: str + used_structured_json: bool + cfg_dropped: bool + caption: str + text_token_ids: tuple[int, ...] + conditioning_fps: float + framing: SFTWindowFraming + + +class SFTVideoInfoProbe: + """Resolve ffprobe metadata and the actual decoded frame count used by SFT.""" + + def __init__( + self, + s3_credentials: Mapping[str, Any], + output_sizes: Mapping[str, tuple[int, int]] | None = None, + ): + self._s3_client = boto3.client("s3", **dict(s3_credentials), config=client_config) + self._output_sizes = output_sizes + + def __call__(self, metadata: dict[str, Any]) -> Mapping[str, Any]: + vision_path = metadata["vision_path"] + video_bytes = download_from_s3(self._s3_client, vision_path) + if video_bytes is None: + raise RuntimeError(f"Failed to download video while resolving SFT framing: {vision_path}") + with tempfile.NamedTemporaryFile(suffix=".mp4", delete=True) as tmp_input: + tmp_input.write(video_bytes) + tmp_input.flush() + video_info = get_video_metadata(tmp_input.name) + scale_hw = None + if self._output_sizes is not None: + target_width, target_height = self._output_sizes[metadata["aspect_ratio"]] + input_width, input_height = metadata["width"], metadata["height"] + resize_ratio = max(target_width / input_width, target_height / input_height) + scale_hw = round(input_height * resize_ratio), round(input_width * resize_ratio) + decoded_total_frames = sum(1 for _ in ffmpeg_decode_video(tmp_input.name, scale_hw=scale_hw, num_threads=2)) + return {**video_info, "decoded_total_frames": decoded_total_frames} + + +def iter_sft_reasoner_documents( + dataset: SFTDataset, + *, + video_info_resolver: SFTVideoInfoResolver | None = None, +) -> Iterator[SFTReasonerDocument]: + """Enumerate every reachable Nano SFT single-caption document exactly once. + + Unlike ``SFTDataset.__iter__``, this function is finite and does not repeat, + shuffle, randomly choose a window, or sample CFG. Metadata uses the same + stable UUID-hash ordering as ``get_sft_dataset``; windows, caption choices, + and reachable CFG states are expanded in deterministic order. + + A custom ``video_info_resolver`` must return ``fps``, ffprobe + ``total_frames``, and the independently measured ``decoded_total_frames``. + The default resolver downloads each unique video once and follows the same + scaled ffmpeg decode path as training so duration text remains byte-identical + when ffprobe overestimates a damaged or variable-frame-rate source. + + Positive conditioning-FPS noise is deliberately unsupported because it + creates an unbounded set of text prompts. Fixed-length random frame + selection is also rejected: the current cache identity has no frame-offset + variant, while Nano's standard native-window recipe does not need one. + """ + + if dataset.conditioning_fps_noise_std > 0: + raise ValueError( + "Offline SFT Reasoner document enumeration requires conditioning_fps_noise_std=0; " + f"got {dataset.conditioning_fps_noise_std}" + ) + if dataset.num_video_frames != -1 and dataset.frame_selection_mode == "random": + raise ValueError( + "Offline SFT Reasoner document enumeration does not support random fixed-length frame selection" + ) + + if getattr(dataset, "is_initialized", False): + raise ValueError( + "Reasoner documents must be enumerated before SFTDataset.__iter__ mutates, pads, and shards metadata" + ) + + resolve_video_info = ( + video_info_resolver + if video_info_resolver is not None + else SFTVideoInfoProbe(dataset.s3_credentials, dataset.output_sizes) + ) + cfg_variants = enumerate_sft_cfg_variants(dataset.cfg_dropout_rate) + video_info_cache: dict[str, Mapping[str, Any]] = {} + + for metadata in sorted(dataset.metadata, key=sft_metadata_sort_key): + vision_path = metadata["vision_path"] + if vision_path not in video_info_cache: + video_info_cache[vision_path] = resolve_video_info(metadata) + video_info = video_info_cache[vision_path] + original_fps = float(video_info["fps"]) + total_frames = int(video_info["total_frames"]) + if "decoded_total_frames" not in video_info: + raise ValueError( + "SFT video_info_resolver must report decoded_total_frames so offline caption framing " + f"matches the training decoder: {vision_path}" + ) + decoded_total_frames = int(video_info["decoded_total_frames"]) + target_width, target_height = dataset.output_sizes[metadata["aspect_ratio"]] + + for window_index, t2w_window in enumerate(metadata["t2w_windows"]): + framing = resolve_sft_window_framing( + metadata, + window_index, + original_fps=original_fps, + total_frames=total_frames, + decoded_total_frames=decoded_total_frames, + num_video_frames=dataset.num_video_frames, + temporal_interval_mode=dataset.temporal_interval_mode, + frame_selection_mode=dataset.frame_selection_mode, + temporal_compression_factor=dataset.temporal_compression_factor, + target_height=target_height, + target_width=target_width, + ) + if framing is None: + log.warning(f"Skipping empty or too-short SFT window during Reasoner enumeration: {metadata['uuid']}") + continue + + caption_variants = enumerate_sft_captions(t2w_window) + if not caption_variants: + log.warning( + f"Skipping SFT window with no selectable caption during Reasoner enumeration: {framing.sample_key}" + ) + continue + + effective_fps = original_fps / framing.temporal_interval + conditioning_fps = float(effective_fps if dataset.conditioning_fps < 0 else dataset.conditioning_fps) + for caption_key, base_caption, used_structured_json in caption_variants: + for cfg_dropped in cfg_variants: + caption = render_sft_caption( + base_caption, + used_structured_json=used_structured_json, + cfg_dropped=cfg_dropped, + cfg_dropout_keep_metadata=dataset.cfg_dropout_keep_metadata, + caption_suffix=dataset.caption_suffix, + append_duration_fps_timestamps=dataset.append_duration_fps_timestamps, + append_resolution_info=dataset.append_resolution_info, + num_frames=framing.num_frames, + conditioning_fps=conditioning_fps, + target_height=target_height, + target_width=target_width, + ) + text_ids, caption = dataset._tokenize_caption(caption) + yield SFTReasonerDocument( + sample_key=framing.sample_key, + vision_path=vision_path, + window_index=window_index, + caption_key=caption_key, + used_structured_json=used_structured_json, + cfg_dropped=cfg_dropped, + caption=caption, + text_token_ids=tuple(int(token_id) for token_id in text_ids), + conditioning_fps=conditioning_fps, + framing=framing, + ) diff --git a/cosmos_framework/data/generator/local_datasets/sft_reasoner_documents_test.py b/cosmos_framework/data/generator/local_datasets/sft_reasoner_documents_test.py new file mode 100644 index 000000000..b1c7168e0 --- /dev/null +++ b/cosmos_framework/data/generator/local_datasets/sft_reasoner_documents_test.py @@ -0,0 +1,352 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 +"""Tests for finite, deterministic SFT Reasoner document enumeration.""" + +from collections.abc import Mapping +from typing import Any + +import numpy as np +import pytest + +from cosmos_framework.data.generator.local_datasets import sft_dataset as sft_dataset_module +from cosmos_framework.data.generator.local_datasets import sft_reasoner_documents as reasoner_documents_module +from cosmos_framework.data.generator.local_datasets.sft_dataset import ( + SFTDataset, + enumerate_sft_captions, + enumerate_sft_cfg_variants, + render_sft_caption, + resolve_sft_window_framing, + sft_metadata_sort_key, +) +from cosmos_framework.data.generator.local_datasets.sft_reasoner_documents import ( + SFTVideoInfoProbe, + iter_sft_reasoner_documents, +) + + +def _metadata(uuid: str, *windows: dict[str, Any]) -> dict[str, Any]: + return { + "uuid": uuid, + "vision_path": f"/videos/{uuid}.mp4", + "width": 256, + "height": 144, + "aspect_ratio": "test", + "t2w_windows": list(windows), + } + + +def _dataset(metadata: list[dict[str, Any]], **overrides: Any) -> SFTDataset: + dataset = object.__new__(SFTDataset) + defaults = { + "metadata": metadata, + "s3_credentials": {}, + "num_video_frames": -1, + "temporal_interval_mode": "max_30fps", + "frame_selection_mode": "first", + "temporal_compression_factor": 4, + "output_sizes": {"test": (256, 144)}, + "cfg_dropout_rate": 0.1, + "cfg_dropout_keep_metadata": False, + "caption_suffix": "quality suffix", + "append_duration_fps_timestamps": True, + "append_resolution_info": True, + "conditioning_fps": -1, + "conditioning_fps_noise_std": 0.0, + "conditioning_config": None, + "is_initialized": False, + "s3_client": object(), + } + defaults.update(overrides) + for name, value in defaults.items(): + setattr(dataset, name, value) + + def tokenize(caption: str) -> tuple[list[int], str]: + return [len(caption), sum(ord(char) for char in caption)], caption + + dataset._tokenize_caption = tokenize + return dataset + + +def _video_info(_: dict[str, Any]) -> Mapping[str, Any]: + return {"fps": 20.0, "total_frames": 40, "decoded_total_frames": 40} + + +def test_reasoner_documents_are_finite_stably_sorted_and_expand_windows_and_cfg(): + metadata = [ + _metadata( + "video-b", + {"start_frame": 0, "end_frame": 8, "temporal_interval": 2, "caption": "first"}, + {"start_frame": 10, "end_frame": 18, "temporal_interval": 2, "caption": "second"}, + ), + _metadata( + "video-a", + {"start_frame": 20, "end_frame": 28, "temporal_interval": 2, "caption": "third"}, + ), + ] + dataset = _dataset(metadata) + + first = list(iter_sft_reasoner_documents(dataset, video_info_resolver=_video_info)) + dataset.metadata = list(reversed(dataset.metadata)) + second = list(iter_sft_reasoner_documents(dataset, video_info_resolver=_video_info)) + + assert first == second + assert len(first) == 6 + expected_keys = [ + f"{item['uuid']}_w{window_index}" + for item in sorted(metadata, key=sft_metadata_sort_key) + for window_index in range(len(item["t2w_windows"])) + for _ in range(2) + ] + assert [document.sample_key for document in first] == expected_keys + assert [document.cfg_dropped for document in first] == [False, True] * 3 + assert all(document.window_index == document.framing.window_index for document in first) + assert all(document.conditioning_fps == 10.0 for document in first) + assert all(document.framing.num_frames == 5 for document in first) + assert all(document.caption == "" for document in first if document.cfg_dropped) + assert all("0.5 seconds" in document.caption for document in first if not document.cfg_dropped) + + +def test_reasoner_document_matches_training_sample_caption_window_and_tokens(monkeypatch): + metadata = _metadata( + "video", + {"start_frame": 0, "end_frame": 8, "temporal_interval": 1, "caption": "same prompt"}, + ) + metadata["width"] = 8 + metadata["height"] = 4 + dataset = _dataset([metadata], cfg_dropout_rate=0.0, output_sizes={"test": (8, 4)}) + + monkeypatch.setattr(sft_dataset_module, "download_from_s3", lambda *_args, **_kwargs: b"video") + monkeypatch.setattr( + sft_dataset_module, + "get_video_metadata", + lambda _path: {"fps": 20.0, "total_frames": 40}, + ) + monkeypatch.setattr( + sft_dataset_module, + "ffmpeg_decode_video", + lambda *_args, **_kwargs: iter(np.zeros((7, 4, 8, 3), dtype=np.uint8)), + ) + monkeypatch.setattr(sft_dataset_module.random, "randrange", lambda _size: 0) + + training_sample = dataset.process_one_sample(metadata) + documents = list( + iter_sft_reasoner_documents( + dataset, + video_info_resolver=lambda _metadata: { + "fps": 20.0, + "total_frames": 40, + "decoded_total_frames": 7, + }, + ) + ) + + assert training_sample is not None + assert len(documents) == 1 + document = documents[0] + assert training_sample["__key__"] == document.sample_key + assert training_sample["ai_caption"] == document.caption + assert tuple(training_sample["text_token_ids"].tolist()) == document.text_token_ids + assert training_sample["frame_start"] == document.framing.start_frame + assert training_sample["frame_end"] == document.framing.end_frame + assert training_sample["num_frames"] == document.framing.num_frames + + +def test_default_video_probe_counts_frames_from_the_real_decode_path(monkeypatch): + metadata = _metadata( + "video", + {"start_frame": 0, "end_frame": 8, "temporal_interval": 1, "caption": "x"}, + ) + seen_decode_args: dict[str, Any] = {} + monkeypatch.setattr(reasoner_documents_module.boto3, "client", lambda *_args, **_kwargs: object()) + monkeypatch.setattr(reasoner_documents_module, "download_from_s3", lambda *_args, **_kwargs: b"video") + monkeypatch.setattr( + reasoner_documents_module, + "get_video_metadata", + lambda _path: {"width": 256, "height": 144, "fps": 20.0, "total_frames": 40}, + ) + + def decode(_path, *, scale_hw, num_threads): + seen_decode_args.update(scale_hw=scale_hw, num_threads=num_threads) + return iter(range(7)) + + monkeypatch.setattr(reasoner_documents_module, "ffmpeg_decode_video", decode) + + info = SFTVideoInfoProbe({}, {"test": (256, 144)})(metadata) + + assert info["decoded_total_frames"] == 7 + assert seen_decode_args == {"scale_hw": (144, 256), "num_threads": 2} + + +def test_reasoner_documents_enumerate_every_positive_weight_caption_in_stable_order(): + window = { + "start_frame": 0, + "end_frame": 8, + "temporal_interval": 2, + "qwen3_235b_dense": "dense", + "qwen3_32b_short": "short", + "qwen3_235b_temporal": "unreachable", + } + dataset = _dataset([_metadata("video", window)], cfg_dropout_rate=0.0) + + documents = list(iter_sft_reasoner_documents(dataset, video_info_resolver=_video_info)) + + assert [document.caption_key for document in documents] == ["qwen3_32b_short", "qwen3_235b_dense"] + assert [document.cfg_dropped for document in documents] == [False, False] + + +def test_caption_priority_shadows_weighted_fallbacks_during_enumeration(): + selections = enumerate_sft_captions( + { + "caption": "preferred", + "qwen3_235b_dense": "fallback", + } + ) + + assert selections == (("caption", "preferred.", False),) + + +def test_conditioning_fps_noise_fails_before_resolving_any_video(): + dataset = _dataset( + [_metadata("video", {"start_frame": 0, "end_frame": 8, "temporal_interval": 2, "caption": "x"})], + conditioning_fps_noise_std=0.1, + ) + resolver_called = False + + def resolver(_: dict[str, Any]) -> Mapping[str, Any]: + nonlocal resolver_called + resolver_called = True + return _video_info({}) + + with pytest.raises(ValueError, match="conditioning_fps_noise_std=0"): + list(iter_sft_reasoner_documents(dataset, video_info_resolver=resolver)) + assert not resolver_called + + +def test_custom_video_info_resolver_must_report_actual_decoded_frame_count(): + dataset = _dataset([_metadata("video", {"start_frame": 0, "end_frame": 8, "temporal_interval": 1, "caption": "x"})]) + + with pytest.raises(ValueError, match="decoded_total_frames"): + list( + iter_sft_reasoner_documents( + dataset, + video_info_resolver=lambda _metadata: {"fps": 20.0, "total_frames": 40}, + ) + ) + + +def test_random_fixed_length_selection_fails_closed(): + dataset = _dataset( + [_metadata("video", {"start_frame": 0, "end_frame": 20, "temporal_interval": 1, "caption": "x"})], + num_video_frames=5, + frame_selection_mode="random", + ) + + with pytest.raises(ValueError, match="random fixed-length"): + list(iter_sft_reasoner_documents(dataset, video_info_resolver=_video_info)) + + +def test_cfg_endpoint_variants_and_flattened_window_key_match_training_semantics(): + metadata = _metadata( + "video_w3", + {"start_frame": 0, "end_frame": 8, "temporal_interval": 2, "caption": "x"}, + ) + dataset = _dataset([metadata], cfg_dropout_rate=1.0) + + documents = list(iter_sft_reasoner_documents(dataset, video_info_resolver=_video_info)) + + assert enumerate_sft_cfg_variants(0.0) == (False,) + assert enumerate_sft_cfg_variants(0.5) == (False, True) + assert enumerate_sft_cfg_variants(1.0) == (True,) + assert [(document.sample_key, document.cfg_dropped) for document in documents] == [("video_w3_w0", True)] + + +def test_enumerator_rejects_dataset_after_infinite_iterator_initialization(): + dataset = _dataset([], is_initialized=True) + + with pytest.raises(ValueError, match="before SFTDataset.__iter__"): + list(iter_sft_reasoner_documents(dataset, video_info_resolver=_video_info)) + + +def test_shared_caption_renderer_preserves_metadata_for_cfg_when_configured(): + caption = render_sft_caption( + "description.", + used_structured_json=False, + cfg_dropped=True, + cfg_dropout_keep_metadata=True, + caption_suffix="quality suffix", + append_duration_fps_timestamps=True, + append_resolution_info=True, + num_frames=5, + conditioning_fps=10.0, + target_height=144, + target_width=256, + ) + + assert caption == ("The video is 0.5 seconds long and is of 10 FPS. This video is of 144x256 resolution.") + + +def test_shared_caption_renderer_keeps_structured_json_byte_stable(): + caption = render_sft_caption( + '{"fps": 5}', + used_structured_json=True, + cfg_dropped=False, + cfg_dropout_keep_metadata=False, + caption_suffix="must not be appended", + append_duration_fps_timestamps=True, + append_resolution_info=True, + num_frames=5, + conditioning_fps=10.0, + target_height=144, + target_width=256, + ) + + assert caption == '{"fps": 5}' + + +def test_shared_window_framing_matches_native_window_clamp_and_truncation(): + metadata = _metadata( + "video", + {"start_frame": 5, "end_frame": 25, "temporal_interval": 3, "caption": "x"}, + ) + + framing = resolve_sft_window_framing( + metadata, + 0, + original_fps=30.0, + total_frames=20, + num_video_frames=-1, + temporal_interval_mode="max_30fps", + frame_selection_mode="first", + temporal_compression_factor=4, + target_height=144, + target_width=256, + ) + + assert framing is not None + assert framing.sample_key == "video_w0" + assert (framing.start_frame, framing.end_frame, framing.temporal_interval) == (5, 19, 3) + assert framing.num_frames == 5 + + +def test_shared_window_framing_counts_only_frames_the_training_decoder_can_reach(): + metadata = _metadata( + "video", + {"start_frame": 0, "end_frame": 99, "temporal_interval": 1, "caption": "x"}, + ) + + framing = resolve_sft_window_framing( + metadata, + 0, + original_fps=60.0, + total_frames=100, + num_video_frames=93, + temporal_interval_mode="max_30fps", + frame_selection_mode="first", + temporal_compression_factor=4, + target_height=144, + target_width=256, + ) + + assert framing is not None + assert (framing.start_frame, framing.end_frame, framing.temporal_interval) == (0, 184, 2) + assert framing.num_frames == 49 diff --git a/cosmos_framework/model/generator/mot/cosmos3_vfm_network.py b/cosmos_framework/model/generator/mot/cosmos3_vfm_network.py index 65d683408..f5ab184d2 100644 --- a/cosmos_framework/model/generator/mot/cosmos3_vfm_network.py +++ b/cosmos_framework/model/generator/mot/cosmos3_vfm_network.py @@ -10,12 +10,19 @@ from transformers.configuration_utils import PretrainedConfig from transformers.modeling_utils import PreTrainedModel -from cosmos_framework.utils import log from cosmos_framework.configs.base.defaults.joint_attention import packing_layout from cosmos_framework.configs.base.defaults.multiview_attention import ( MultiviewAttentionConfig, ResolvedBackend, ) +from cosmos_framework.data.generator.sequence_packing import ModalityData, PackedSequence +from cosmos_framework.data.generator.sequence_packing.natten import verify_natten_parameter_list +from cosmos_framework.data.generator.sequence_packing.runtime import ( + SequencePack, + get_caption_seq_offsets, + get_causal_seq, + get_full_only_seq, +) from cosmos_framework.model.generator.mot.action_io_projector import ( ACTION_IO_PROJECTOR_DOMAIN_AWARE, ACTION_IO_PROJECTOR_TYPES, @@ -43,14 +50,7 @@ build_multiview_maskless_plan, ) from cosmos_framework.model.generator.utils.memory import MemoryState -from cosmos_framework.data.generator.sequence_packing import ModalityData, PackedSequence -from cosmos_framework.data.generator.sequence_packing.natten import verify_natten_parameter_list -from cosmos_framework.data.generator.sequence_packing.runtime import ( - SequencePack, - get_caption_seq_offsets, - get_causal_seq, - get_full_only_seq, -) +from cosmos_framework.utils import log class Cosmos3VFMNetworkConfig(PretrainedConfig): @@ -700,7 +700,12 @@ def _encode_text( self, packed_seq: PackedSequence, ) -> tuple[torch.Tensor, torch.dtype]: - """Embed text tokens and initialize packed_sequence. + """Embed text tokens and initialize ``packed_sequence``. + + A generator-only model receives the Reasoner's per-layer K/V through a + ``MemoryState`` and deliberately has no token-embedding table. In that + mode the causal rows are zero placeholders used only to preserve the + existing packed layout; decoder layers never consume them. Args: packed_seq: PackedSequence containing text_ids and text_indexes. @@ -708,7 +713,19 @@ def _encode_text( Returns: tuple of (packed_sequence, target_dtype) where packed_sequence has text embeddings filled in. """ - packed_text_embedding = self.language_model.model.embed_tokens(packed_seq.text_ids) # [N_text,hidden_size] + language_model = self.language_model.model + if not getattr(language_model, "include_und_pathway", True): + if not hasattr(language_model, "norm_moe_gen"): + raise RuntimeError("A generator-only language model must retain norm_moe_gen.") + reference_param = next(language_model.norm_moe_gen.parameters()) + packed_sequence = torch.zeros( + (packed_seq.sequence_length, self.hidden_size), + device=packed_seq.text_ids.device, + dtype=reference_param.dtype, + ) + return packed_sequence, reference_param.dtype + + packed_text_embedding = language_model.embed_tokens(packed_seq.text_ids) # [N_text,hidden_size] packed_sequence = packed_text_embedding.new_zeros( size=(packed_seq.sequence_length, self.hidden_size) ) # [N_total,hidden_size] diff --git a/cosmos_framework/model/generator/mot/cosmos3_vfm_network_reasoner_features_test.py b/cosmos_framework/model/generator/mot/cosmos3_vfm_network_reasoner_features_test.py new file mode 100644 index 000000000..76ed249f2 --- /dev/null +++ b/cosmos_framework/model/generator/mot/cosmos3_vfm_network_reasoner_features_test.py @@ -0,0 +1,52 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +from types import SimpleNamespace + +import pytest +import torch + +from cosmos_framework.model.generator.mot.cosmos3_vfm_network import Cosmos3VFMNetwork + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_encode_text_uses_zero_layout_for_generator_only_model() -> None: + hidden_size = 8 + generator_norm = torch.nn.LayerNorm(hidden_size, dtype=torch.bfloat16) + network = SimpleNamespace( + hidden_size=hidden_size, + language_model=SimpleNamespace( + model=SimpleNamespace( + include_und_pathway=False, + norm_moe_gen=generator_norm, + ) + ), + ) + packed_seq = SimpleNamespace( + sequence_length=11, + text_ids=torch.tensor([3, 4, 5], dtype=torch.long), + ) + + packed, dtype = Cosmos3VFMNetwork._encode_text(network, packed_seq) + + assert packed.shape == (11, hidden_size) + assert packed.dtype == torch.bfloat16 + assert dtype == torch.bfloat16 + assert torch.count_nonzero(packed) == 0 + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_encode_text_requires_generator_norm_for_generator_only_model() -> None: + network = SimpleNamespace( + hidden_size=8, + language_model=SimpleNamespace(model=SimpleNamespace(include_und_pathway=False)), + ) + packed_seq = SimpleNamespace( + sequence_length=4, + text_ids=torch.tensor([1], dtype=torch.long), + ) + + with pytest.raises(RuntimeError, match="must retain norm_moe_gen"): + Cosmos3VFMNetwork._encode_text(network, packed_seq) diff --git a/cosmos_framework/model/generator/mot/include_gen_pathway_test.py b/cosmos_framework/model/generator/mot/include_gen_pathway_test.py index 2e767bba9..805adfddf 100644 --- a/cosmos_framework/model/generator/mot/include_gen_pathway_test.py +++ b/cosmos_framework/model/generator/mot/include_gen_pathway_test.py @@ -1,6 +1,6 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: OpenMDW-1.1 -"""Unit tests for ``include_gen_pathway``. +"""Unit tests for reciprocal MoT pathway inclusion flags. Reasoner-only inference never runs the MoT generation tower, so it can leave the ``*_moe_gen`` duplicates unbuilt. ``dcp.load`` is pull-based (it requests only @@ -11,17 +11,33 @@ behaviour and the unchanged default. """ +from __future__ import annotations + +from typing import Any + +import pytest +import torch import torch.nn as nn +from cosmos_framework.data.generator.sequence_packing.runtime import ( + SequencePack, + from_und_gen_splits, + get_gen_seq, + get_und_seq, +) +from cosmos_framework.model.generator.mot.attention import build_packed_sequence from cosmos_framework.model.generator.mot.unified_mot import ( LayerTypes, MoTDecoderLayer, Nemotron3DenseVLMoTConfig, + Nemotron3DenseVLTextForCausalLM, PackedAttentionMoT, + prune_und_pathway_, ) from cosmos_framework.model.generator.reasoner.nemotron_3_dense_vl.configuration_nemotron_3_dense_vl import ( Nemotron3DenseVLTextConfig, ) +from cosmos_framework.model.generator.utils.memory import KVToStore, MemoryState, MemoryValue NUM_Q_HEADS = 4 NUM_KV_HEADS = 2 @@ -41,6 +57,8 @@ "input_layernorm_moe_gen", "post_attention_layernorm_moe_gen", ) +_ATTN_UND_MODULES = ("q_proj", "k_proj", "v_proj", "o_proj", "q_norm", "k_norm") +_LAYER_UND_MODULES = ("mlp", "input_layernorm", "post_attention_layernorm") def _tiny_config() -> Nemotron3DenseVLTextConfig: @@ -53,8 +71,12 @@ def _tiny_config() -> Nemotron3DenseVLTextConfig: ) -def _make_attention(*, include_gen_pathway: bool | None = None) -> PackedAttentionMoT: +def _make_attention( + *, include_gen_pathway: bool | None = None, include_und_pathway: bool | None = None +) -> PackedAttentionMoT: kwargs = {} if include_gen_pathway is None else {"include_gen_pathway": include_gen_pathway} + if include_und_pathway is not None: + kwargs["include_und_pathway"] = include_und_pathway return PackedAttentionMoT( _tiny_config(), layer_idx=0, @@ -65,8 +87,10 @@ def _make_attention(*, include_gen_pathway: bool | None = None) -> PackedAttenti ) -def _make_layer(*, include_gen_pathway: bool | None = None) -> MoTDecoderLayer: +def _make_layer(*, include_gen_pathway: bool | None = None, include_und_pathway: bool | None = None) -> MoTDecoderLayer: kwargs = {} if include_gen_pathway is None else {"include_gen_pathway": include_gen_pathway} + if include_und_pathway is not None: + kwargs["include_und_pathway"] = include_und_pathway return MoTDecoderLayer( config=_tiny_config(), layer_idx=0, @@ -121,3 +145,163 @@ def test_mot_config_includes_gen_pathway_by_default() -> None: def test_mot_config_forwards_disabled_flag() -> None: assert Nemotron3DenseVLMoTConfig({}, include_gen_pathway=False).include_gen_pathway is False + + +def test_attention_omits_und_modules_when_disabled() -> None: + attn = _make_attention(include_und_pathway=False) + + for name in _ATTN_UND_MODULES: + assert not hasattr(attn, name), f"{name} should not be built when include_und_pathway=False" + assert attn.k_norm_und_for_gen is None + for name in _ATTN_GEN_MODULES: + assert hasattr(attn, name), f"{name} must remain in a generator-only attention module" + + +def test_decoder_layer_omits_und_modules_when_disabled() -> None: + layer = _make_layer(include_und_pathway=False) + + for name in _LAYER_UND_MODULES: + assert not hasattr(layer, name), f"{name} should not be built when include_und_pathway=False" + for name in _LAYER_GEN_MODULES: + assert hasattr(layer, name), f"{name} must remain in a generator-only decoder layer" + assert not [name for name, _ in layer.named_parameters() if "moe_gen" not in name] + + +def test_und_pathway_is_included_by_default() -> None: + attn = _make_attention() + layer = _make_layer() + + assert Nemotron3DenseVLMoTConfig({}).include_und_pathway is True + for name in _ATTN_UND_MODULES: + assert hasattr(attn, name) + for name in _LAYER_UND_MODULES: + assert hasattr(layer, name) + + +def _tiny_mot_config(*, include_und_pathway: bool = True) -> Nemotron3DenseVLMoTConfig: + return Nemotron3DenseVLMoTConfig( + { + "vocab_size": 32, + "hidden_size": NUM_Q_HEADS * HEAD_DIM, + "intermediate_size": 128, + "num_hidden_layers": 1, + "num_attention_heads": NUM_Q_HEADS, + "num_key_value_heads": NUM_KV_HEADS, + "head_dim": HEAD_DIM, + "enable_mrope": False, + }, + include_und_pathway=include_und_pathway, + ) + + +def _assert_generator_only_structure(causal_lm: nn.Module) -> None: + assert not hasattr(causal_lm, "lm_head") + assert not hasattr(causal_lm.model, "embed_tokens") + assert not hasattr(causal_lm.model, "norm") + assert hasattr(causal_lm.model, "norm_moe_gen") + assert hasattr(causal_lm.model, "rotary_emb") + + layer = causal_lm.model.layers[0] + for name in _LAYER_UND_MODULES: + assert not hasattr(layer, name) + for name in _ATTN_UND_MODULES: + assert not hasattr(layer.self_attn, name) + + parameter_names = [name for name, _ in causal_lm.named_parameters()] + assert parameter_names + assert all("moe_gen" in name for name in parameter_names) + + +def test_for_causal_lm_builds_generator_only_structure() -> None: + causal_lm = Nemotron3DenseVLTextForCausalLM(_tiny_mot_config(include_und_pathway=False)) + + _assert_generator_only_structure(causal_lm) + with pytest.raises(RuntimeError, match="include_und_pathway=False"): + causal_lm.model.reasoner_forward(torch.ones((1, 1), dtype=torch.long), cache=None) + with pytest.raises(RuntimeError, match="include_und_pathway=False"): + causal_lm.generate_reasoner_text(torch.ones((1, 1), dtype=torch.long), max_new_tokens=0) + + +def test_prune_und_pathway_is_idempotent_and_preserves_gen_fqns() -> None: + with torch.device("meta"): + causal_lm = Nemotron3DenseVLTextForCausalLM(_tiny_mot_config()) + causal_lm.visual = nn.Linear(2, 2) + gen_parameter_names = {name for name, _ in causal_lm.named_parameters() if "moe_gen" in name} + + assert prune_und_pathway_(causal_lm) is causal_lm + assert prune_und_pathway_(causal_lm) is causal_lm + + _assert_generator_only_structure(causal_lm) + assert {name for name, _ in causal_lm.named_parameters()} == gen_parameter_names + assert all(parameter.is_meta for parameter in causal_lm.parameters()) + assert causal_lm.model.include_und_pathway is False + assert causal_lm.model.layers[0].include_und_pathway is False + assert causal_lm.model.layers[0].self_attn.include_und_pathway is False + assert not hasattr(causal_lm, "visual") + + +class _GenOnlyMemory(MemoryState): + def init(self, hidden_states: dict, device: torch.device) -> None: + del hidden_states, device + + def read_for_layer(self, layer_idx: int) -> MemoryValue: + del layer_idx + return MemoryValue() + + def write_for_layer(self, layer_idx: int, kv_to_store: KVToStore) -> None: + del layer_idx, kv_to_store + + def is_gen_only(self) -> bool: + return True + + +def _passthrough_gen_attention( + packed_query_states: SequencePack, + packed_key_states: SequencePack, + packed_value_states: SequencePack, + attention_mask: Any, + natten_metadata: dict | None = None, + memory_value: MemoryValue | None = None, + packed_key_states_normalized: SequencePack | None = None, +) -> tuple[SequencePack, None]: + del packed_key_states, packed_value_states, attention_mask, natten_metadata, packed_key_states_normalized + assert memory_value is not None + gen = get_gen_seq(packed_query_states).flatten(-2, -1) + empty_und = gen.new_empty((0, gen.shape[-1])) + return from_und_gen_splits(empty_und, gen, packed_query_states), None + + +@pytest.mark.parametrize("prune_after_build", [False, True]) +def test_generator_only_joint_forward_requires_and_accepts_gen_only_memory(prune_after_build: bool) -> None: + causal_lm = Nemotron3DenseVLTextForCausalLM(_tiny_mot_config(include_und_pathway=prune_after_build)) + if prune_after_build: + prune_und_pathway_(causal_lm) + for layer in causal_lm.model.layers: + layer.self_attn.dispatch_attention_fn = _passthrough_gen_attention + + hidden = torch.randn(4, NUM_Q_HEADS * HEAD_DIM) + pack, attention_mask, _ = build_packed_sequence( + "two_way", + packed_sequence=hidden, + attn_modes=["full"], + split_lens=[hidden.shape[0]], + sample_lens=[hidden.shape[0]], + packed_und_token_indexes=torch.empty(0, dtype=torch.long), + packed_gen_token_indexes=torch.arange(hidden.shape[0]), + num_heads=NUM_Q_HEADS, + head_dim=HEAD_DIM, + num_layers=1, + ) + + with pytest.raises(RuntimeError, match="requires a MemoryState"): + causal_lm(pack, attention_mask, torch.arange(hidden.shape[0])) + + output, metadata = causal_lm( + pack, + attention_mask, + torch.arange(hidden.shape[0]), + memory=_GenOnlyMemory(), + ) + assert metadata == {} + assert get_und_seq(output).shape == (0, hidden.shape[-1]) + assert get_gen_seq(output).shape == hidden.shape diff --git a/cosmos_framework/model/generator/mot/unified_mot.py b/cosmos_framework/model/generator/mot/unified_mot.py index 2c95692e0..dcf25b61e 100644 --- a/cosmos_framework/model/generator/mot/unified_mot.py +++ b/cosmos_framework/model/generator/mot/unified_mot.py @@ -13,9 +13,20 @@ from torch import nn from torch.distributed import ProcessGroup +from cosmos_framework.data.generator.sequence_packing.runtime import ( + SequencePack, + from_all_seq, + from_und_gen_splits, + get_gen_seq, + get_num_real_tokens, + get_und_seq, + has_pad_segment, + set_gen_seq, + set_und_seq, + zeros_like, +) from cosmos_framework.model.attention import attention as imaginaire_attention from cosmos_framework.model.attention.masks import CausalType -from cosmos_framework.utils import log from cosmos_framework.model.generator.mot.attention import ( AttentionMaskType, dispatch_attention, @@ -77,18 +88,7 @@ LBLMetadata, ) from cosmos_framework.model.generator.utils.memory import KVToStore, MemoryState, MemoryValue -from cosmos_framework.data.generator.sequence_packing.runtime import ( - SequencePack, - from_all_seq, - from_und_gen_splits, - get_gen_seq, - get_num_real_tokens, - get_und_seq, - has_pad_segment, - set_gen_seq, - set_und_seq, - zeros_like, -) +from cosmos_framework.utils import log # Torch optimization settings torch._dynamo.config.cache_size_limit = 512 @@ -278,6 +278,7 @@ def __init__( qk_norm_for_diffusion: bool = True, include_visual: bool = False, include_gen_pathway: bool = True, + include_und_pathway: bool = True, gen_noisy_gating: bool = False, gen_cosine_router_config: CosineRouterConfig | None = None, gen_aux_loss_free_load_balancing_config: AuxLossFreeLoadBalancingConfig | None = None, @@ -295,6 +296,10 @@ def __init__( # Build the MoT generation tower (the ``*_moe_gen`` duplicates). Reasoner-only # inference disables this; every other caller keeps the default and is unchanged. self.include_gen_pathway = include_gen_pathway + # Build the understanding/reasoner tower. Generator-only training can disable + # this after supplying the per-layer UND K/V from an external feature provider. + # The default deliberately preserves the historical full dual-pathway model. + self.include_und_pathway = include_und_pathway # Noisy top-k gating on the generation-tower MoE blocks (Shazeer 2017). # Gen-tower only; the understanding tower never receives this flag. self.gen_noisy_gating = gen_noisy_gating @@ -542,10 +547,12 @@ def __init__( qk_norm_for_diffusion: bool, use_und_k_norm_for_gen: bool = False, include_gen_pathway: bool = True, + include_und_pathway: bool = True, ): super().__init__() self.config = config self.include_gen_pathway = include_gen_pathway + self.include_und_pathway = include_und_pathway self.layer_idx = layer_idx self.head_dim = getattr(config, "head_dim", config.hidden_size // config.num_attention_heads) self.hidden_size = config.hidden_size @@ -557,19 +564,29 @@ def __init__( eps = config.rms_norm_eps - # Understanding pathway projections - self.q_proj = nn.Linear(self.hidden_size, self.num_attention_heads * self.head_dim, bias=config.attention_bias) - self.k_proj = nn.Linear(self.hidden_size, self.num_key_value_heads * self.head_dim, bias=config.attention_bias) - self.v_proj = nn.Linear(self.hidden_size, self.num_key_value_heads * self.head_dim, bias=config.attention_bias) - self.o_proj = nn.Linear(self.num_attention_heads * self.head_dim, self.hidden_size, bias=config.attention_bias) + # Understanding pathway projections and QK norm. These modules are + # intentionally absent (rather than frozen) in a generator-only model, + # so FSDP and EMA never see or materialize their parameters. + if include_und_pathway: + self.q_proj = nn.Linear( + self.hidden_size, self.num_attention_heads * self.head_dim, bias=config.attention_bias + ) + self.k_proj = nn.Linear( + self.hidden_size, self.num_key_value_heads * self.head_dim, bias=config.attention_bias + ) + self.v_proj = nn.Linear( + self.hidden_size, self.num_key_value_heads * self.head_dim, bias=config.attention_bias + ) + self.o_proj = nn.Linear( + self.num_attention_heads * self.head_dim, self.hidden_size, bias=config.attention_bias + ) - # Understanding pathway QK norm - if qk_norm_for_text: - self.q_norm = layer_types.rms_norm(self.head_dim, eps=eps) - self.k_norm = layer_types.rms_norm(self.head_dim, eps=eps) - else: - self.q_norm = nn.Identity() - self.k_norm = nn.Identity() + if qk_norm_for_text: + self.q_norm = layer_types.rms_norm(self.head_dim, eps=eps) + self.k_norm = layer_types.rms_norm(self.head_dim, eps=eps) + else: + self.q_norm = nn.Identity() + self.k_norm = nn.Identity() # Generation pathway QK norm. Everything below this point belongs to the # generation tower and is skipped wholesale when it is not built. @@ -590,7 +607,13 @@ def __init__( # When both pathways share the same QK norm (or neither has one) k_norm_und_for_gen # is None and the standard packed K tensor is used for all paths unchanged. # It serves the generation pathway only, so it is None whenever that tower is absent. - if include_gen_pathway and use_und_k_norm_for_gen and qk_norm_for_diffusion and not qk_norm_for_text: + if ( + include_gen_pathway + and include_und_pathway + and use_und_k_norm_for_gen + and qk_norm_for_diffusion + and not qk_norm_for_text + ): self.k_norm_und_for_gen: nn.Module | None = layer_types.rms_norm(self.head_dim, eps=eps) else: self.k_norm_und_for_gen = None @@ -732,23 +755,34 @@ def forward( memory_value, ) - q_und_in = self.q_proj(get_und_seq(pack)) # [N_und,num_heads*head_dim] - q_gen_in = self.q_proj_moe_gen(get_gen_seq(pack)) # [N_gen,num_heads*head_dim] + if not self.include_und_pathway and memory_value is None: + raise RuntimeError( + "PackedAttentionMoT was built without the understanding pathway; " + "its forward requires generator-only memory containing external UND K/V." + ) - k_und_in = self.k_proj(get_und_seq(pack)) # [N_und,num_kv_heads*head_dim] + q_gen_in = self.q_proj_moe_gen(get_gen_seq(pack)) # [N_gen,num_heads*head_dim] k_gen_in = self.k_proj_moe_gen(get_gen_seq(pack)) # [N_gen,num_kv_heads*head_dim] - - v_und_in = self.v_proj(get_und_seq(pack)) # [N_und,num_kv_heads*head_dim] v_gen_in = self.v_proj_moe_gen(get_gen_seq(pack)) # [N_gen,num_kv_heads*head_dim] - q_und = q_und_in.view(-1, self.num_attention_heads, self.head_dim) # [N_und,num_heads,head_dim] - k_und = k_und_in.view(-1, self.num_key_value_heads, self.head_dim) # [N_und,num_kv_heads,head_dim] - v_und = v_und_in.view(-1, self.num_key_value_heads, self.head_dim) # [N_und,num_kv_heads,head_dim] - q_gen = q_gen_in.view(-1, self.num_attention_heads, self.head_dim) # [N_gen,num_heads,head_dim] k_gen = k_gen_in.view(-1, self.num_key_value_heads, self.head_dim) # [N_gen,num_kv_heads,head_dim] v_gen = v_gen_in.view(-1, self.num_key_value_heads, self.head_dim) # [N_gen,num_kv_heads,head_dim] + if self.include_und_pathway: + q_und_in = self.q_proj(get_und_seq(pack)) # [N_und,num_heads*head_dim] + k_und_in = self.k_proj(get_und_seq(pack)) # [N_und,num_kv_heads*head_dim] + v_und_in = self.v_proj(get_und_seq(pack)) # [N_und,num_kv_heads*head_dim] + q_und = q_und_in.view(-1, self.num_attention_heads, self.head_dim) # [N_und,num_heads,head_dim] + k_und = k_und_in.view(-1, self.num_key_value_heads, self.head_dim) # [N_und,num_kv_heads,head_dim] + v_und = v_und_in.view(-1, self.num_key_value_heads, self.head_dim) # [N_und,num_kv_heads,head_dim] + else: + # External-memory attention supplies the per-layer UND K/V. Keep an + # empty UND split in the live pack so only generator projections run. + q_und = q_gen.new_empty((0, self.num_attention_heads, self.head_dim)) + k_und = k_gen.new_empty((0, self.num_key_value_heads, self.head_dim)) + v_und = v_gen.new_empty((0, self.num_key_value_heads, self.head_dim)) + # The sequence length is the only size that varies between steps, but Dynamo lifts the int # attributes of a module into SymInts, so the head counts reach the views above as symbols. # That breaks FlexAttention's Inductor lowering, which tests the Q:KV head ratio for a power @@ -761,22 +795,24 @@ def forward( for head_split in (q_und, k_und, v_und, q_gen, k_gen, v_gen): torch._dynamo.mark_static(head_split, 1) - q_und = self.q_norm(q_und) # [N_und,num_heads,head_dim] - k_und = self.k_norm(k_und) # [N_und,num_kv_heads,head_dim] - q_gen = self.q_norm_moe_gen(q_gen) # [N_gen,num_heads,head_dim] k_gen = self.k_norm_moe_gen(k_gen) # [N_gen,num_kv_heads,head_dim] packed_cos = packed_position_embeddings[0] packed_sin = packed_position_embeddings[1] - q_und_, k_und_ = self._apply_rotary_pos_emb( - q_und, - k_und, - get_und_seq(packed_cos), - get_und_seq(packed_sin), - unsqueeze_dim=1, - ) # q_und_: [N_und,num_heads,head_dim], k_und_: [N_und,num_kv_heads,head_dim] + if self.include_und_pathway: + q_und = self.q_norm(q_und) # [N_und,num_heads,head_dim] + k_und = self.k_norm(k_und) # [N_und,num_kv_heads,head_dim] + q_und_, k_und_ = self._apply_rotary_pos_emb( + q_und, + k_und, + get_und_seq(packed_cos), + get_und_seq(packed_sin), + unsqueeze_dim=1, + ) # q_und_: [N_und,num_heads,head_dim], k_und_: [N_und,num_kv_heads,head_dim] + else: + q_und_, k_und_ = q_und, k_und q_gen_, k_gen_ = self._apply_rotary_pos_emb( q_gen, k_gen, @@ -829,7 +865,10 @@ def forward( and kv_to_store is None and not bool(getattr(memory_value, "target_only_no_text", False)) ): - und_len = pack["_num_causal_tokens"] + # A generator-only structural model carries the original prompt + # layout in metadata for the external K/V provider, but it has no + # live UND projections to write back from this call. + und_len = pack["_num_causal_tokens"] if self.include_und_pathway else 0 gen_len = pack["_num_full_tokens"] # When und K-norm is active, AR frame 1+ gen→und cross-attention uses # the normalised K, so cache k_und_for_gen_ (RMSNorm+RoPE applied) instead @@ -890,7 +929,10 @@ def forward( cp_group, ) # [N_gen,hidden_size] else: - und_seq = self.o_proj(get_und_seq(packed_attn_output)) # [N_und,hidden_size] + if self.include_und_pathway: + und_seq = self.o_proj(get_und_seq(packed_attn_output)) # [N_und,hidden_size] + else: + und_seq = get_gen_seq(packed_attn_output).new_empty((0, self.hidden_size)) gen_seq = self.o_proj_moe_gen(get_gen_seq(packed_attn_output)) # [N_gen,hidden_size] return from_und_gen_splits(und_seq, gen_seq, pack), kv_to_store # [N_und+N_gen,hidden_size] @@ -924,6 +966,8 @@ def reasoner_forward( Multi-token incremental prefill on top of an existing cache is not supported here — ``_impl_generate_reasoner_text`` never triggers it. """ + if not self.include_und_pathway: + raise RuntimeError("reasoner_forward is unavailable because include_und_pathway=False.") B, T, _ = hidden_states.shape H = self.num_attention_heads H_kv = self.num_key_value_heads @@ -983,6 +1027,7 @@ def _impl_init( gen_moe_shared_expert_intermediate_scale: int = 1, gen_moe_top_k: int | None = None, include_gen_pathway: bool = True, + include_und_pathway: bool = True, ) -> None: """Shared ``__init__`` body for the three MoT text-model variants. @@ -995,8 +1040,10 @@ def _impl_init( # Read back by ``_impl_forward``'s guard: the joint generation forward cannot # run on a model built without the generation tower. self.include_gen_pathway = include_gen_pathway + self.include_und_pathway = include_und_pathway - self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx) + if include_und_pathway: + self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx) self.layers = nn.ModuleList() for layer_idx in range(config.num_hidden_layers): @@ -1016,11 +1063,13 @@ def _impl_init( gen_moe_shared_expert_intermediate_scale=gen_moe_shared_expert_intermediate_scale, gen_moe_top_k=gen_moe_top_k, include_gen_pathway=include_gen_pathway, + include_und_pathway=include_und_pathway, ) ) # Reasoner-pathway final norm. - self.norm = layer_types.rms_norm(config.hidden_size, eps=config.rms_norm_eps) + if include_und_pathway: + self.norm = layer_types.rms_norm(config.hidden_size, eps=config.rms_norm_eps) if include_gen_pathway: # Generation-pathway final norm (parallel to ``self.norm``). self.norm_moe_gen = layer_types.rms_norm(config.hidden_size, eps=config.rms_norm_eps) @@ -1142,6 +1191,11 @@ def _impl_forward( # Derive gen_only once (outside compile) if using MemoryState memory_gen_only = memory.is_gen_only() if memory is not None else False + if not getattr(self, "include_und_pathway", True) and not memory_gen_only: + raise RuntimeError( + "The joint generation forward was built with include_und_pathway=False and therefore requires " + "a MemoryState whose is_gen_only() is True and whose per-layer values provide external UND K/V." + ) for i, decoder_layer in enumerate(self.layers): # MemoryState: produce read-only MemoryValue for this layer (outside compile) @@ -1166,9 +1220,14 @@ def _impl_forward( # Dense models produce no metadata. MoE models produce one stacked entry per pathway. final_lbl_metadata = _stack_lbl_metadata(lbl_metadata_all) - hidden_states_out = zeros_like(hidden_states) - set_und_seq(hidden_states_out, self.norm(get_und_seq(hidden_states))) # [N_und,hidden_size] - set_gen_seq(hidden_states_out, self.norm_moe_gen(get_gen_seq(hidden_states))) # [N_gen,hidden_size] + if getattr(self, "include_und_pathway", True): + hidden_states_out = zeros_like(hidden_states) + set_und_seq(hidden_states_out, self.norm(get_und_seq(hidden_states))) # [N_und,hidden_size] + set_gen_seq(hidden_states_out, self.norm_moe_gen(get_gen_seq(hidden_states))) # [N_gen,hidden_size] + else: + gen_seq = self.norm_moe_gen(get_gen_seq(hidden_states)) # [N_gen,hidden_size] + empty_und = gen_seq.new_empty((0, gen_seq.shape[-1])) + hidden_states_out = from_und_gen_splits(empty_und, gen_seq, hidden_states) return hidden_states_out, final_lbl_metadata @@ -1303,10 +1362,12 @@ def __init__( gen_moe_shared_expert_intermediate_scale: int = 1, gen_moe_top_k: int | None = None, include_gen_pathway: bool = True, + include_und_pathway: bool = True, ) -> None: super().__init__() self.hidden_size = config.hidden_size self.include_gen_pathway = include_gen_pathway + self.include_und_pathway = include_und_pathway self.self_attn = PackedAttentionMoT( config, layer_types=layer_types, @@ -1315,6 +1376,7 @@ def __init__( qk_norm_for_diffusion=qk_norm_for_diffusion, use_und_k_norm_for_gen=use_und_k_norm_for_gen, include_gen_pathway=include_gen_pathway, + include_und_pathway=include_und_pathway, ) if ( @@ -1322,7 +1384,8 @@ def __init__( and (layer_idx not in config.mlp_only_layers) and (config.num_experts > 0 and (layer_idx + 1) % config.decoder_sparse_step == 0) ): - self.mlp = Qwen3VLMoeTextSparseMoeBlock(config) + if include_und_pathway: + self.mlp = Qwen3VLMoeTextSparseMoeBlock(config) if include_gen_pathway: # Noisy gating, the cosine router, aux-loss-free load balancing, # the shared expert, and the top-k override are gen-tower only. @@ -1336,17 +1399,20 @@ def __init__( top_k=gen_moe_top_k, ) else: - self.mlp = layer_types.mlp(config) + if include_und_pathway: + self.mlp = layer_types.mlp(config) if include_gen_pathway: self.mlp_moe_gen = layer_types.mlp(config) # Each ``*_moe_gen`` norm stays registered next to its und counterpart so the # module (and therefore state-dict) order is byte-for-byte the previous one # whenever the generation pathway is built. - self.input_layernorm = layer_types.rms_norm(config.hidden_size, eps=config.rms_norm_eps) + if include_und_pathway: + self.input_layernorm = layer_types.rms_norm(config.hidden_size, eps=config.rms_norm_eps) if include_gen_pathway: self.input_layernorm_moe_gen = layer_types.rms_norm(config.hidden_size, eps=config.rms_norm_eps) - self.post_attention_layernorm = layer_types.rms_norm(config.hidden_size, eps=config.rms_norm_eps) + if include_und_pathway: + self.post_attention_layernorm = layer_types.rms_norm(config.hidden_size, eps=config.rms_norm_eps) if include_gen_pathway: self.post_attention_layernorm_moe_gen = layer_types.rms_norm(config.hidden_size, eps=config.rms_norm_eps) self.lbl_config: LBLConfig = lbl_config or LBLConfig() @@ -1388,6 +1454,11 @@ def forward( Target-only teacher forcing also slices away control GEN rows before every decoder operation and restores the full layout on return. """ + if not self.include_und_pathway and not gen_only: + raise RuntimeError( + "MoTDecoderLayer was built with include_und_pathway=False; forward requires gen_only=True " + "with external UND K/V supplied through memory_value." + ) target_only_no_text = gen_only and bool(getattr(memory_value, "target_only_no_text", False)) layer_input = input layer_position_embeddings = packed_position_embeddings @@ -1398,9 +1469,8 @@ def forward( raise ValueError("Target-only teacher forcing does not support context-parallel NATTEN metadata.") if self._sample_lbl_und or self._sample_lbl_gen: raise ValueError("Target-only teacher forcing does not support sample load-balancing metadata.") - if isinstance(self.mlp, Qwen3VLMoeTextSparseMoeBlock) or isinstance( - self.mlp_moe_gen, Qwen3VLMoeTextSparseMoeBlock - ): + und_mlp_is_sparse = self.include_und_pathway and isinstance(self.mlp, Qwen3VLMoeTextSparseMoeBlock) + if und_mlp_is_sparse or isinstance(self.mlp_moe_gen, Qwen3VLMoeTextSparseMoeBlock): raise ValueError("Target-only teacher forcing currently supports dense decoder MLPs only.") target_start = int(getattr(memory_value, "target_gen_start", 0)) target_length = int(getattr(memory_value, "target_gen_length", 0)) @@ -1440,8 +1510,8 @@ def forward( ) # Pre-Attention layernorm - if target_only_no_text: - norm_und = get_und_seq(layer_input) # [0,hidden_size] + if gen_only: + norm_und = get_und_seq(layer_input)[:0] # [0,hidden_size] else: norm_und = self.input_layernorm(get_und_seq(layer_input)) # [N_und,hidden_size] pack_norm_out = from_und_gen_splits( @@ -1621,6 +1691,8 @@ def reasoner_forward( single-rank settings the registration is a no-op and this method just runs the und pathway directly. """ + if not self.include_und_pathway: + raise RuntimeError("reasoner_forward is unavailable because include_und_pathway=False.") residual = hidden_states h = self.input_layernorm(hidden_states) attn_out = self.self_attn.reasoner_forward(h, cos, sin, cache, layer_idx) @@ -1655,6 +1727,7 @@ def __init__( qk_norm_for_diffusion: bool, use_und_k_norm_for_gen: bool, include_gen_pathway: bool = True, + include_und_pathway: bool = True, ): super().__init__(config) _impl_init( @@ -1665,6 +1738,7 @@ def __init__( qk_norm_for_diffusion=qk_norm_for_diffusion, use_und_k_norm_for_gen=use_und_k_norm_for_gen, include_gen_pathway=include_gen_pathway, + include_und_pathway=include_und_pathway, ) def forward(self, *args, **kwargs): @@ -1696,6 +1770,7 @@ def __init__( gen_moe_shared_expert_intermediate_scale: int = 1, gen_moe_top_k: int | None = None, include_gen_pathway: bool = True, + include_und_pathway: bool = True, ) -> None: super().__init__(config) _impl_init( @@ -1713,6 +1788,7 @@ def __init__( gen_moe_shared_expert_intermediate_scale=gen_moe_shared_expert_intermediate_scale, gen_moe_top_k=gen_moe_top_k, include_gen_pathway=include_gen_pathway, + include_und_pathway=include_und_pathway, ) def forward(self, *args, **kwargs): @@ -1737,6 +1813,7 @@ def __init__( qk_norm_for_diffusion: bool, use_und_k_norm_for_gen: bool, include_gen_pathway: bool = True, + include_und_pathway: bool = True, ): super().__init__(config) _impl_init( @@ -1747,6 +1824,7 @@ def __init__( qk_norm_for_diffusion=qk_norm_for_diffusion, use_und_k_norm_for_gen=use_und_k_norm_for_gen, include_gen_pathway=include_gen_pathway, + include_und_pathway=include_und_pathway, ) def forward(self, *args, **kwargs): @@ -1873,6 +1951,8 @@ def _impl_reasoner_forward( ``inputs_embeds.dtype``; the canonical producer ``prepare_multimodal_reasoner_inputs`` aligns both. """ + if not getattr(self, "include_und_pathway", True): + raise RuntimeError("reasoner_forward is unavailable because include_und_pathway=False.") if (input_ids is None) == (inputs_embeds is None): raise ValueError("Specify exactly one of input_ids or inputs_embeds.") @@ -2177,6 +2257,8 @@ def _impl_generate_reasoner_text( ``return_only_new_tokens=True``). ``T_new <= max_new_tokens``; early termination only occurs when every sample emits EOS. """ + if not getattr(causal_lm.model, "include_und_pathway", True): + raise RuntimeError("generate_reasoner_text is unavailable because include_und_pathway=False.") if input_ids.dim() != 2 or input_ids.shape[1] < 1: raise ValueError(f"input_ids must have shape [B, T_prompt>=1], got {tuple(input_ids.shape)}") if max_new_tokens < 0: @@ -2419,15 +2501,20 @@ def __init__(self, config: Qwen3VLMoTConfig): super().__init__(config.full_config) text_config = config.text_config + include_und_pathway = getattr(config, "include_und_pathway", True) self.model = Qwen3VLTextModel( text_config, qk_norm_for_text=config.qk_norm_for_text, qk_norm_for_diffusion=config.qk_norm_for_diffusion, use_und_k_norm_for_gen=getattr(config, "use_und_k_norm_for_gen", False), include_gen_pathway=getattr(config, "include_gen_pathway", True), + include_und_pathway=include_und_pathway, ) self.vocab_size = text_config.vocab_size - self.lm_head = nn.Linear(text_config.hidden_size, text_config.vocab_size, bias=False) + if include_und_pathway: + self.lm_head = nn.Linear(text_config.hidden_size, text_config.vocab_size, bias=False) + else: + self._tied_weights_keys = [] # The wrapper's ``vision_config`` property gates on # ``include_visual`` and materializes the HF vision config from @@ -2451,6 +2538,8 @@ def init_moe(self) -> None: keep their default ``ones`` init and we just skip the copy rather than raising. """ + if not self.model.include_und_pathway: + raise RuntimeError("init_moe requires the understanding pathway, but include_und_pathway=False.") state_dict = self.state_dict() for name, param in self.named_parameters(): if "moe_gen" not in name: @@ -2467,9 +2556,13 @@ def get_input_embeddings(self) -> nn.Embedding: # via `base_model_prefix="model"`, but defining the method here is # the canonical HF idiom and removes a hidden dependency on # `base_model_prefix` being correctly set. + if not self.model.include_und_pathway: + raise RuntimeError("Input embeddings are unavailable because include_und_pathway=False.") return self.model.embed_tokens def set_input_embeddings(self, value: nn.Embedding) -> None: + if not self.model.include_und_pathway: + raise RuntimeError("Cannot set input embeddings because include_und_pathway=False.") self.model.embed_tokens = value def forward( @@ -2569,12 +2662,14 @@ def __init__( super().__init__(config.full_config) text_config = config.text_config + include_und_pathway = getattr(config, "include_und_pathway", True) self.model = Qwen3VLMoeTextModel( text_config, qk_norm_for_text=config.qk_norm_for_text, qk_norm_for_diffusion=config.qk_norm_for_diffusion, use_und_k_norm_for_gen=getattr(config, "use_und_k_norm_for_gen", False), include_gen_pathway=getattr(config, "include_gen_pathway", True), + include_und_pathway=include_und_pathway, gen_noisy_gating=config.gen_noisy_gating, gen_cosine_router_config=getattr(config, "gen_cosine_router_config", None), gen_aux_loss_free_load_balancing_config=config.gen_aux_loss_free_load_balancing_config, @@ -2588,7 +2683,10 @@ def __init__( gen_moe_top_k=getattr(config, "gen_moe_top_k", None), ) self.vocab_size = text_config.vocab_size - self.lm_head = nn.Linear(text_config.hidden_size, text_config.vocab_size, bias=False) + if include_und_pathway: + self.lm_head = nn.Linear(text_config.hidden_size, text_config.vocab_size, bias=False) + else: + self._tied_weights_keys = [] # The wrapper's ``vision_config`` property gates on # ``include_visual`` and materializes the HF vision config from @@ -2606,6 +2704,8 @@ def init_moe(self) -> None: See :meth:`Qwen3VLTextForCausalLM.init_moe` for the q_norm/k_norm Identity-tower handling shared with the dense variant. """ + if not self.model.include_und_pathway: + raise RuntimeError("init_moe requires the understanding pathway, but include_und_pathway=False.") state_dict = self.state_dict() for name, param in self.named_parameters(): if "moe_gen" not in name: @@ -2634,9 +2734,13 @@ def init_moe(self) -> None: def get_input_embeddings(self) -> nn.Embedding: # See note on `Qwen3VLTextForCausalLM.get_input_embeddings`. + if not self.model.include_und_pathway: + raise RuntimeError("Input embeddings are unavailable because include_und_pathway=False.") return self.model.embed_tokens def set_input_embeddings(self, value: nn.Embedding) -> None: + if not self.model.include_und_pathway: + raise RuntimeError("Cannot set input embeddings because include_und_pathway=False.") self.model.embed_tokens = value def forward( @@ -2773,15 +2877,18 @@ def __init__(self, config: Nemotron3DenseVLMoTConfig) -> None: super().__init__(config.full_config) text_config = config.text_config + include_und_pathway = getattr(config, "include_und_pathway", True) self.model = Nemotron3DenseVLTextModel( text_config, qk_norm_for_text=config.qk_norm_for_text, qk_norm_for_diffusion=config.qk_norm_for_diffusion, use_und_k_norm_for_gen=getattr(config, "use_und_k_norm_for_gen", False), include_gen_pathway=getattr(config, "include_gen_pathway", True), + include_und_pathway=include_und_pathway, ) self.vocab_size = text_config.vocab_size - self.lm_head = nn.Linear(text_config.hidden_size, text_config.vocab_size, bias=False) + if include_und_pathway: + self.lm_head = nn.Linear(text_config.hidden_size, text_config.vocab_size, bias=False) assert config.vision_config is None, "Nemotron 3 Dense VL has no vision config" @@ -2789,6 +2896,8 @@ def __init__(self, config: Nemotron3DenseVLMoTConfig) -> None: def init_moe(self) -> None: """Copy understanding-pathway weights into the generation-pathway parameters.""" + if not self.model.include_und_pathway: + raise RuntimeError("init_moe requires the understanding pathway, but include_und_pathway=False.") state_dict = self.state_dict() for name, param in self.named_parameters(): if "moe_gen" not in name: @@ -2805,9 +2914,13 @@ def init_moe(self) -> None: def get_input_embeddings(self) -> nn.Embedding: # See note on `Qwen3VLTextForCausalLM.get_input_embeddings`. + if not self.model.include_und_pathway: + raise RuntimeError("Input embeddings are unavailable because include_und_pathway=False.") return self.model.embed_tokens def set_input_embeddings(self, value: nn.Embedding) -> None: + if not self.model.include_und_pathway: + raise RuntimeError("Cannot set input embeddings because include_und_pathway=False.") self.model.embed_tokens = value def forward( @@ -3002,3 +3115,57 @@ def _hub_error(filename: str) -> RuntimeError: self.config.video_token_id = top_cfg["video_token_id"] self.config.vision_start_token_id = top_cfg["vision_start_token_id"] self.config.vision_config = SimpleNamespace(spatial_merge_size=pc["spatial_merge_size"]) + + +def prune_und_pathway_(causal_lm: nn.Module) -> nn.Module: + """Remove the frozen understanding tower before materialization/FSDP. + + The operation is in-place and idempotent. It is intended for a complete + ``*TextForCausalLM`` wrapper constructed on the meta device: pruning there + prevents the deleted parameters from ever being allocated, sharded, or + copied into EMA. Generation-tower module names are not rewritten, so their + checkpoint FQNs stay compatible with the full model. + + After pruning, joint forward is valid only with a generator-only + :class:`MemoryState` that supplies external per-layer UND K/V. + """ + model = getattr(causal_lm, "model", None) + layers = getattr(model, "layers", None) + if model is None or layers is None: + raise TypeError("prune_und_pathway_ expects a *TextForCausalLM wrapper with model.layers.") + if not getattr(model, "include_gen_pathway", True): + raise ValueError("Cannot prune the understanding pathway from a model without a generation pathway.") + + def _delete(module: nn.Module, *names: str) -> None: + for name in names: + if hasattr(module, name): + delattr(module, name) + + # ``visual`` is the optional Reasoner-side image/video encoder. External + # text K/V makes it as unnecessary as the token embedding and LM head. + _delete(causal_lm, "lm_head", "visual") + _delete(model, "embed_tokens", "norm") + model.include_und_pathway = False + if hasattr(model, "config"): + model.config.include_und_pathway = False + + for layer in layers: + _delete(layer, "mlp", "input_layernorm", "post_attention_layernorm") + layer.include_und_pathway = False + + self_attn = layer.self_attn + _delete(self_attn, "q_proj", "k_proj", "v_proj", "o_proj", "q_norm", "k_norm") + # This optional normalizer acts only on newly projected UND K. External + # features already contain the generator-facing normalized/RoPE K. + self_attn.k_norm_und_for_gen = None + self_attn.include_und_pathway = False + + if hasattr(causal_lm, "config"): + causal_lm.config.include_und_pathway = False + if hasattr(causal_lm.config, "include_visual"): + causal_lm.config.include_visual = False + # Qwen wrappers declare the reasoner embedding/head tie at class level. + # Shadow it on this generator-only instance so save/load utilities do not + # look for a pair that was deliberately removed. + causal_lm._tied_weights_keys = [] + return causal_lm diff --git a/cosmos_framework/model/generator/omni_mot_causal_model.py b/cosmos_framework/model/generator/omni_mot_causal_model.py index 9b1f81501..45c364e7c 100644 --- a/cosmos_framework/model/generator/omni_mot_causal_model.py +++ b/cosmos_framework/model/generator/omni_mot_causal_model.py @@ -27,19 +27,21 @@ from typing_extensions import override import cosmos_framework.model.generator.omni_mot_model as omni_mot_model_module -from cosmos_framework.configs.base.defaults.model_config import OmniMoTModelConfig -from cosmos_framework.data.generator.augmentors.text_tokenizer import TEXT_SYSTEM_PROMPT_KEY -from cosmos_framework.model.generator.omni_mot_model import OmniMoTModel, _broadcast_seed, _per_view_caption_groups -from cosmos_framework.model.generator.utils.data_and_condition import GenerationDataClean -from cosmos_framework.model.generator.utils.memory import MemoryState -from cosmos_framework.data.generator.sequence_packing import PackedSequence, build_sequence_plans_from_data_batch -from cosmos_framework.data.generator.sequence_packing.modality import compute_text_split_length -from cosmos_framework.data.generator.sequence_packing.runtime import to_device_nonblocking from cosmos_framework.configs.base.defaults.causal_flex_attention import CausalFlexAttentionConfig +from cosmos_framework.configs.base.defaults.model_config import OmniMoTModelConfig from cosmos_framework.configs.base.defaults.replay_attention import ( TeacherForcingKVImplementation, TeacherForcingReplayPolicyConfig, ) +from cosmos_framework.data.generator.augmentors.text_tokenizer import TEXT_SYSTEM_PROMPT_KEY +from cosmos_framework.data.generator.sequence_packing import PackedSequence, build_sequence_plans_from_data_batch +from cosmos_framework.data.generator.sequence_packing.autoregressive import ( + pack_input_sequence_autoregressive, + pack_input_sequence_autoregressive_batch, + resolve_text_system_prompt, +) +from cosmos_framework.data.generator.sequence_packing.modality import compute_text_split_length +from cosmos_framework.data.generator.sequence_packing.runtime import to_device_nonblocking from cosmos_framework.model.generator.attention_io_layout import AttentionIOLayout from cosmos_framework.model.generator.joint_transfer_ar import sample_joint_transfer_ar from cosmos_framework.model.generator.mot.causal_attention import dispatch_attention_with_memory @@ -61,9 +63,11 @@ validate_ar_static_und_cache_lengths, ) from cosmos_framework.model.generator.multiview_transfer_ar import MultiviewTransferARBackend +from cosmos_framework.model.generator.omni_mot_model import OmniMoTModel, _broadcast_seed, _per_view_caption_groups from cosmos_framework.model.generator.teacher_forcing import ( make_teacher_forcing_clean_pack, ) +from cosmos_framework.model.generator.utils.data_and_condition import GenerationDataClean from cosmos_framework.model.generator.utils.kv_cache import ( ARMemoryState, DualKVCache, @@ -72,12 +76,8 @@ TeacherForcingMemoryState, ) from cosmos_framework.model.generator.utils.kv_storage_backend import validate_kv_cache_dtype +from cosmos_framework.model.generator.utils.memory import MemoryState from cosmos_framework.model.generator.utils.nvfp4 import resolve_legacy_nvfp4_mode -from cosmos_framework.data.generator.sequence_packing.autoregressive import ( - pack_input_sequence_autoregressive, - pack_input_sequence_autoregressive_batch, - resolve_text_system_prompt, -) from cosmos_framework.utils.generator.data_batch import condition_frame_indexes_vision_from_batch _ARBranch = Literal["conditional", "unconditional"] @@ -552,6 +552,13 @@ class OmniMoTCausalModel(OmniMoTModel): _teacher_forcing_replay_policy_runtime: TeacherForcingReplayPolicyConfig def __init__(self, config: OmniMoTCausalModelConfig): + reasoner_backend = omni_mot_model_module._reasoner_conditioning_backend(config) + if reasoner_backend != "joint": + raise ValueError( + "OmniMoTCausalModel currently supports only reasoner_conditioning.backend='joint'; " + f"got {reasoner_backend!r}. Its AR/teacher-forcing memory dispatcher is not yet composable " + "with inline or external Reasoner K/V conditioning." + ) # LazyCall deliberately keeps nested config values as DictConfig. Hold # the validated attrs object separately: assigning it back into the # structured DictConfig would immediately coerce it back to DictConfig. diff --git a/cosmos_framework/model/generator/omni_mot_model.py b/cosmos_framework/model/generator/omni_mot_model.py index a70431234..f4b77d092 100644 --- a/cosmos_framework/model/generator/omni_mot_model.py +++ b/cosmos_framework/model/generator/omni_mot_model.py @@ -8,6 +8,7 @@ import inspect import json import time +from concurrent.futures import Future from contextlib import contextmanager from typing import Any, Callable, Dict, Mapping, Optional, Tuple, get_args @@ -20,20 +21,6 @@ from torch.distributed.fsdp import MixedPrecisionPolicy from torch.nn.modules.module import _IncompatibleKeys -from cosmos_framework.utils.flags import DEVICE, Device -from cosmos_framework.utils.lazy_config import LazyDict -from cosmos_framework.utils.lazy_config import instantiate as lazy_instantiate -from cosmos_framework.utils.lazy_config.registry import locate -from cosmos_framework.model._base import ImaginaireModel -from cosmos_framework.utils import log, misc -from cosmos_framework.utils.count_params import count_params -from cosmos_framework.model.generator.algorithm.loss.flow_matching import ( - ACTION_SLOT_SAMPLE_COUNT_KEY, - ACTION_SLOT_SAMPLE_LOSS_KEY, - ActionSlotLossStats, - compute_flow_matching_loss, -) -from cosmos_framework.model.generator.algorithm.loss.load_balancing import compute_load_balancing_loss from cosmos_framework.configs.base.defaults.joint_attention import JointAttnImplementation from cosmos_framework.configs.base.defaults.model_config import OmniMoTModelConfig from cosmos_framework.configs.base.defaults.parallelism import PRECISION_TO_TORCH_DTYPE @@ -42,7 +29,23 @@ get_action_processing_records, ) from cosmos_framework.data.generator.action.utils.unified_action_schema import UNIFIED_ACTION_SLOT_GROUPS +from cosmos_framework.data.generator.sequence_packing import ( + PackedSequence, + SequencePlan, + build_sequence_plans_from_data_batch, + pack_input_sequence, +) +from cosmos_framework.data.generator.sequence_packing.modality import add_special_tokens +from cosmos_framework.data.generator.sequence_packing.packers import is_item_generated, uses_single_timestep from cosmos_framework.data.generator.utils import IMAGE_RES_SIZE_INFO, VIDEO_RES_SIZE_INFO +from cosmos_framework.model._base import ImaginaireModel +from cosmos_framework.model.generator.algorithm.loss.flow_matching import ( + ACTION_SLOT_SAMPLE_COUNT_KEY, + ACTION_SLOT_SAMPLE_LOSS_KEY, + ActionSlotLossStats, + compute_flow_matching_loss, +) +from cosmos_framework.model.generator.algorithm.loss.load_balancing import compute_load_balancing_loss from cosmos_framework.model.generator.diffusion.rectified_flow import RectifiedFlow from cosmos_framework.model.generator.diffusion.samplers.edm import EDMSampler from cosmos_framework.model.generator.diffusion.samplers.fixed_step import FixedStepSampler @@ -62,7 +65,22 @@ from cosmos_framework.model.generator.mot.modeling_utils import has_noisy_tokens from cosmos_framework.model.generator.mot.parallelize_unified_mot import materialize_non_offloaded_state from cosmos_framework.model.generator.mot.parallelize_vfm_network import parallelize_vfm_network +from cosmos_framework.model.generator.mot.unified_mot import prune_und_pathway_ from cosmos_framework.model.generator.reasoner.qwen3_vl.utils import tokenize_caption +from cosmos_framework.model.generator.reasoner_feature_cache import ( + OfflineReasonerFeatureProvider, + ReasonerFeatureCacheIdentity, + build_reasoner_feature_requests, +) +from cosmos_framework.model.generator.reasoner_features import ( + CapturingReasonerKVMemoryState, + ReasonerFeatureBatch, + ReasonerFeatureProvider, + StaticReasonerKVMemoryState, + install_reasoner_feature_attention_dispatch, +) +from cosmos_framework.model.generator.tokenizers.interface import VideoTokenizerInterface +from cosmos_framework.model.generator.upsampler.prompts import build_messages, clean_response from cosmos_framework.model.generator.utils.data_and_condition import ( GenerationDataClean, GenerationDataNoised, @@ -93,21 +111,17 @@ get_vae_pixel_shapes, normalize_uint8_item, ) -from cosmos_framework.data.generator.sequence_packing import ( - PackedSequence, - SequencePlan, - build_sequence_plans_from_data_batch, - pack_input_sequence, -) -from cosmos_framework.data.generator.sequence_packing.modality import add_special_tokens -from cosmos_framework.data.generator.sequence_packing.packers import is_item_generated, uses_single_timestep -from cosmos_framework.model.generator.tokenizers.interface import VideoTokenizerInterface -from cosmos_framework.model.generator.upsampler.prompts import build_messages, clean_response +from cosmos_framework.utils import log, misc +from cosmos_framework.utils.count_params import count_params +from cosmos_framework.utils.flags import DEVICE, Device from cosmos_framework.utils.generator.data_utils import get_vision_data_resolution, read_positive_int_metadata from cosmos_framework.utils.generator.dtensor_helper import DTensorFastEmaModelUpdater from cosmos_framework.utils.generator.model_weights_stats import WeightTrainingStat from cosmos_framework.utils.generator.parallelism import ParallelDims from cosmos_framework.utils.generator.quantization import swap_modelopt_fp8_linears_on_meta +from cosmos_framework.utils.lazy_config import LazyDict +from cosmos_framework.utils.lazy_config import instantiate as lazy_instantiate +from cosmos_framework.utils.lazy_config.registry import locate def _all_group_ranks_allow( @@ -235,6 +249,100 @@ def _densify_action_family( INFERENCE_RAW_VISION_RETAINED_ITEMS_KEY = "_inference_raw_vision_retained_items" +REASONER_FEATURE_BATCH_KEY = "reasoner_feature_batch" +REASONER_FEATURE_FUTURE_KEY = "reasoner_feature_future" +REASONER_SAMPLE_KEYS_KEY = "reasoner_sample_keys" + +_EXTERNAL_REASONER_BACKENDS = frozenset({"offline", "remote", "read_through"}) +_REASONER_FEATURE_BACKENDS = frozenset({"inline", *_EXTERNAL_REASONER_BACKENDS}) + + +def _reasoner_conditioning_backend(config: OmniMoTModelConfig) -> str: + conditioning = getattr(config, "reasoner_conditioning", None) + if conditioning is None: + return "joint" + if isinstance(conditioning, Mapping): + return str(conditioning.get("backend", "joint")) + return str(getattr(conditioning, "backend", "joint")) + + +def _conditioning_value(config: OmniMoTModelConfig, name: str, default: Any = None) -> Any: + conditioning = getattr(config, "reasoner_conditioning", None) + if conditioning is None: + return default + if isinstance(conditioning, Mapping): + return conditioning.get(name, default) + return getattr(conditioning, name, default) + + +def _reasoner_cache_identity(config: OmniMoTModelConfig) -> ReasonerFeatureCacheIdentity | None: + if _reasoner_conditioning_backend(config) not in _EXTERNAL_REASONER_BACKENDS: + return None + values = { + "reasoner": _conditioning_value(config, "reasoner_fingerprint"), + "tokenizer": _conditioning_value(config, "tokenizer_fingerprint"), + "framing": _conditioning_value(config, "framing_fingerprint"), + } + if all(value is None for value in values.values()): + return None + missing = [name for name, value in values.items() if not value] + if missing: + raise ValueError(f"Incomplete Reasoner cache identity; missing {missing}") + return ReasonerFeatureCacheIdentity(**values) + + +def _validate_reasoner_conditioning(config: OmniMoTModelConfig) -> str: + """Validate the deliberately narrow first implementation of external K/V.""" + backend = _reasoner_conditioning_backend(config) + supported = {"joint", *_REASONER_FEATURE_BACKENDS} + if backend not in supported: + raise ValueError(f"Unsupported reasoner_conditioning.backend={backend!r}; expected one of {sorted(supported)}") + if backend == "joint": + return backend + + if config.joint_attn_implementation != "two_way": + raise ValueError( + f"reasoner_conditioning.backend={backend!r} currently requires joint_attn_implementation='two_way'" + ) + if config.parallelism.context_parallel_shard_degree != 1: + raise ValueError(f"reasoner_conditioning.backend={backend!r} currently requires context parallel degree 1") + if config.video_temporal_causal: + raise ValueError(f"reasoner_conditioning.backend={backend!r} does not yet support temporal-causal training") + if getattr(config, "causal_training_strategy", "none") != "none": + raise ValueError( + f"reasoner_conditioning.backend={backend!r} currently requires causal_training_strategy='none'" + ) + + cache_root = _conditioning_value(config, "cache_root") + endpoint = _conditioning_value(config, "endpoint") + if backend in {"offline", "read_through"} and not cache_root: + raise ValueError(f"reasoner_conditioning.backend={backend!r} requires cache_root") + if backend in {"remote", "read_through"} and not endpoint: + raise ValueError(f"reasoner_conditioning.backend={backend!r} requires endpoint") + if backend in _EXTERNAL_REASONER_BACKENDS and not bool(_conditioning_value(config, "strict_fingerprint", True)): + raise ValueError(f"reasoner_conditioning.backend={backend!r} requires strict_fingerprint=true") + if backend in _EXTERNAL_REASONER_BACKENDS: + if _reasoner_cache_identity(config) is None: + raise ValueError( + f"reasoner_conditioning.backend={backend!r} with strict_fingerprint=true requires " + "reasoner_fingerprint, tokenizer_fingerprint, and framing_fingerprint" + ) + if bool(_conditioning_value(config, "layerwise_h2d", False)): + raise ValueError("reasoner_conditioning.layerwise_h2d is reserved for a later implementation") + return backend + + +def _reasoner_sample_keys(data_batch: Mapping[str, Any], batch_size: int) -> tuple[str, ...]: + raw_keys = data_batch.get("__key__") + if isinstance(raw_keys, str): + keys = (raw_keys,) + elif isinstance(raw_keys, (list, tuple)): + keys = tuple(str(key) for key in raw_keys) + else: + raise ValueError("External Reasoner conditioning requires a per-sample '__key__' field in the data batch") + if len(keys) != batch_size or any(not key for key in keys): + raise ValueError(f"Expected {batch_size} non-empty Reasoner sample keys, got {keys!r}") + return keys @dataclasses.dataclass(frozen=True) @@ -269,6 +377,19 @@ def __init__(self, config: OmniMoTModelConfig): # each successful training step and returns to 0 for the next window. self._cp_window_slot: int = 0 self.config = config + self.reasoner_conditioning_backend = _validate_reasoner_conditioning(config) + if self.reasoner_conditioning_backend != "joint" and type(self) is not OmniMoTModel: + raise ValueError( + f"{type(self).__name__} does not yet compose its overridden lifecycle with " + f"reasoner_conditioning.backend={self.reasoner_conditioning_backend!r}; " + "use the base OmniMoTModel Nano SFT path." + ) + self.reasoner_cache_identity = _reasoner_cache_identity(config) + self.reasoner_feature_provider = self._create_reasoner_feature_provider() + if self.reasoner_cache_identity is None and isinstance( + self.reasoner_feature_provider, OfflineReasonerFeatureProvider + ): + self.reasoner_cache_identity = self.reasoner_feature_provider.manifest.identity log.info(f"OmniMoTModel: config {self.config}") # 0. Set up precision @@ -289,6 +410,44 @@ def __init__(self, config: OmniMoTModelConfig): # 5. Set up training time scheduler and inference time sampler self.set_up_scheduler_and_sampler() + def _create_reasoner_feature_provider(self) -> ReasonerFeatureProvider | None: + if self.reasoner_conditioning_backend in {"joint", "inline"}: + return None + if self.reasoner_conditioning_backend == "offline": + cache_root = _conditioning_value(self.config, "cache_root") + assert cache_root + return OfflineReasonerFeatureProvider( + cache_root, + expected_identity=self.reasoner_cache_identity, + expected_dtype=PRECISION_TO_TORCH_DTYPE[self.config.precision], + strict_fingerprint=bool(_conditioning_value(self.config, "strict_fingerprint", True)), + ) + raise NotImplementedError( + f"reasoner_conditioning.backend={self.reasoner_conditioning_backend!r} is reserved by the common " + "feature contract, but its service client has not been implemented yet; use 'offline' or 'inline'." + ) + + def _validate_reasoner_feature_signature(self, net: torch.nn.Module) -> None: + """Reject a cache built for a different decoder before FSDP/materialization.""" + if not isinstance(self.reasoner_feature_provider, OfflineReasonerFeatureProvider): + return + manifest = self.reasoner_feature_provider.manifest + expected = ( + int(net.num_hidden_layers), + int(net.num_kv_heads), + int(net.head_dim), + ) + actual = ( + manifest.num_layers, + manifest.num_kv_heads, + manifest.head_dim, + ) + if actual != expected: + raise ValueError( + "Reasoner cache/model architecture mismatch: " + f"cache(layers,kv_heads,head_dim)={actual}, model={expected}" + ) + def set_precision(self) -> None: self.precision = PRECISION_TO_TORCH_DTYPE[self.config.precision] self.tensor_kwargs = {"device": DEVICE, "dtype": self.precision} @@ -407,7 +566,6 @@ def set_up_tokenizers(self) -> None: else: self.tokenizer_sound_gen = None - def build_net( self, dtype: torch.dtype, @@ -496,6 +654,16 @@ def build_net( ) net.pad_for_cuda_graphs = self.config.compile.use_cuda_graphs + self._validate_reasoner_feature_signature(net) + + # External providers make the Reasoner structurally unnecessary on + # training ranks. Prune while every parameter is still on meta, so + # the deleted weights are never materialized, FSDP-sharded, or + # duplicated into EMA. ``inline`` intentionally retains both towers + # as a numerical-reference/capture mode. + if self.reasoner_conditioning_backend in _EXTERNAL_REASONER_BACKENDS: + prune_und_pathway_(net.language_model) + # Inject LoRA BEFORE FSDP wrap, while still on meta device. The # injector must see unsharded Linear shapes; injecting post-FSDP causes # lora_B to be created at the per-rank shard size and crashes at @@ -517,6 +685,8 @@ def build_net( swap_modelopt_fp8_linears_on_meta(net, self.config.quantization.modelopt_fp8_target_fqns) self.install_attention_dispatch(net) + if self.reasoner_conditioning_backend in _REASONER_FEATURE_BACKENDS: + install_reasoner_feature_attention_dispatch(net) # Cast while still on meta (free -- no data to convert) and BEFORE sharding: each # FSDP2 unit records its parameters' dtype as it is built, so a cast afterwards @@ -561,6 +731,9 @@ def load_pretrained_model_if_needed( *, has_resumable_checkpoint: bool, has_load_path: bool, + warm_start_ema_skipped: bool = False, + warm_start_ema_partially_skipped: bool = False, + warm_start_strict_resume: bool = True, ) -> None: """Conditionally seed pretrained understanding/reasoner weights at startup. @@ -584,6 +757,16 @@ def load_pretrained_model_if_needed( still re-seeded from HF (e.g. to swap Qwen3-VL -> Cosmos-Reason), but the understanding->generation copy is skipped because the generation pathway was already populated from ``load_path``. + warm_start_ema_skipped: The warm-start DCP planner deliberately skipped + the complete ``net_ema`` subtree. External generator-only models + initialize EMA from the loaded regular generator only in this case. + warm_start_ema_partially_skipped: The warm-start DCP planner skipped + only some ``net_ema`` leaves. External generator-only models reject + this state because it would mix restored and random EMA weights. + warm_start_strict_resume: Whether DCP rejects missing warm-start keys. + With non-strict loading and no complete EMA skip, the caller cannot + prove whether EMA was loaded or left randomly initialized, so the + external path fails closed. The gates combine into three startup scenarios: 1. Fresh init (neither gate set): seed understanding weights from HF and @@ -599,6 +782,53 @@ def load_pretrained_model_if_needed( # generation copy further below must be skipped. has_checkpoint = has_resumable_checkpoint or has_load_path + if self.reasoner_conditioning_backend in _EXTERNAL_REASONER_BACKENDS: + if self.config.diffusion_expert_config.load_weights_from_pretrained and not has_checkpoint: + raise ValueError( + f"reasoner_conditioning.backend={self.reasoner_conditioning_backend!r} removed the " + "understanding pathway, so it cannot initialize generator weights by copying from the " + "Reasoner. Set checkpoint.load_path to a generator checkpoint (or explicitly set " + "diffusion_expert_config.load_weights_from_pretrained=false for random initialization)." + ) + # External backends build a structurally generator-only network, so + # the regular warm-start checkpoint is the source of truth for every + # parameter that remains. Nano warm starts intentionally skip + # ``net_ema.*``; without this copy, EMA would retain the random state + # captured in ``set_up_model`` before DCP loaded ``net.*``. Do not do + # this on same-job resume: that path restores the real EMA trajectory + # and must not overwrite it with regular weights. + if ( + has_load_path + and not has_resumable_checkpoint + and self.config.ema.enabled + and warm_start_ema_partially_skipped + ): + raise ValueError( + "External Reasoner conditioning with EMA cannot use a warm start whose " + "checkpoint.keys_to_skip_loading matches only part of the 'net_ema.' subtree. " + "Load all EMA leaves or skip the complete EMA subtree." + ) + if ( + has_load_path + and not has_resumable_checkpoint + and self.config.ema.enabled + and not warm_start_ema_skipped + and not warm_start_strict_resume + ): + raise ValueError( + "External Reasoner conditioning with EMA cannot safely use a non-strict warm start unless " + "checkpoint.keys_to_skip_loading covers the complete 'net_ema.' subtree. Otherwise it is " + "ambiguous whether EMA was loaded or left randomly initialized." + ) + if has_load_path and not has_resumable_checkpoint and warm_start_ema_skipped and self.config.ema.enabled: + self.net_ema_worker.copy_to(src_model=self.net, tgt_model=self.net_ema) + log.info("Initialized external-backend EMA from warm-started generator weights.") + log.info( + "Skipping Reasoner checkpoint loading for structurally generator-only backend " + f"{self.reasoner_conditioning_backend}." + ) + return + pretrained_weights = self.vlm_config.pretrained_weights if self.config.exclude_reasoner_weights_from_checkpoint and not pretrained_weights.enabled: @@ -1057,10 +1287,13 @@ def memory_init_training( ``(gen_data_clean, memory_info)`` where *memory_info* is a dict with keys: ``skip_text``, ``initial_temporal_offset`` """ - return gen_data_clean, { + memory_info: dict[str, Any] = { "skip_text": False, "initial_temporal_offset": 0, } + if self.reasoner_conditioning_backend in _EXTERNAL_REASONER_BACKENDS: + memory_info[REASONER_SAMPLE_KEYS_KEY] = _reasoner_sample_keys(data_batch, gen_data_clean.batch_size) + return gen_data_clean, memory_info def build_memory_state( self, @@ -1081,7 +1314,91 @@ def build_memory_state( (for the training path) or constructed by the AR inference caller. See ``memory_init_training()`` for the base keys. """ - return None + del packed_seq + if self.reasoner_conditioning_backend == "joint": + return None + if self.reasoner_conditioning_backend == "inline": + return CapturingReasonerKVMemoryState(len(self.net.language_model.model.layers)) + + feature_batch = memory_info.get(REASONER_FEATURE_BATCH_KEY) + if not isinstance(feature_batch, ReasonerFeatureBatch): + raise RuntimeError( + f"reasoner_conditioning.backend={self.reasoner_conditioning_backend!r} requires " + f"memory_info[{REASONER_FEATURE_BATCH_KEY!r}] to contain a ReasonerFeatureBatch. " + "The offline/remote provider must populate it before denoise()." + ) + return StaticReasonerKVMemoryState(feature_batch) + + def _denoise_training_with_reasoner_conditioning( + self, + packed_sequence: PackedSequence, + memory_info: dict, + ) -> dict: + """Run the training denoiser with the configured Reasoner feature source.""" + if self.reasoner_conditioning_backend in _EXTERNAL_REASONER_BACKENDS: + self._resolve_reasoner_feature_batch(memory_info) + memory = self.build_memory_state(packed_sequence, memory_info) + if self.reasoner_conditioning_backend == "inline": + if not isinstance(memory, CapturingReasonerKVMemoryState): + raise TypeError("The inline Reasoner backend requires CapturingReasonerKVMemoryState") + # The first pass is a frozen feature extraction reference. Its + # generator outputs are discarded and captured UND K/V are + # detached, so retaining an autograd graph only wastes memory. + with torch.no_grad(), misc.timer("inline reasoner K/V capture"): + capture_output = self.denoise( + data_batch_packed=packed_sequence, + memory=memory, + ) + del capture_output + if not memory.is_gen_only(): + raise RuntimeError("Inline Reasoner K/V capture did not populate every decoder layer") + + return self.denoise( + data_batch_packed=packed_sequence, + memory=memory, + ) + + def _resolve_reasoner_feature_batch(self, memory_info: dict[str, Any]) -> ReasonerFeatureBatch: + existing = memory_info.get(REASONER_FEATURE_BATCH_KEY) + if isinstance(existing, ReasonerFeatureBatch): + return existing + future = memory_info.get(REASONER_FEATURE_FUTURE_KEY) + if not isinstance(future, Future): + raise RuntimeError( + f"memory_info[{REASONER_FEATURE_FUTURE_KEY!r}] is missing; " + "pre_noise_memory_hook() must submit the provider request before denoising" + ) + + result: ReasonerFeatureBatch | None = None + local_error: Exception | None = None + try: + timeout = float(_conditioning_value(self.config, "request_timeout_s", 300.0)) + with misc.timer("reasoner feature provider wait", debug=True): + candidate = future.result(timeout=timeout) + if not isinstance(candidate, ReasonerFeatureBatch): + raise TypeError(f"Reasoner feature provider returned {type(candidate).__name__}") + result = candidate + except Exception as error: + local_error = error + + local_message = None if local_error is None else f"{type(local_error).__name__}: {local_error}" + all_messages = [local_message] + if dist.is_initialized(): + all_messages = [None] * dist.get_world_size() + dist.all_gather_object(all_messages, local_message) + failures = [f"rank {rank}: {message}" for rank, message in enumerate(all_messages) if message is not None] + if failures: + synchronized_error = RuntimeError( + "Reasoner feature provider failed before FSDP forward: " + "; ".join(failures) + ) + if local_error is not None: + raise synchronized_error from local_error + raise synchronized_error + if result is None: + raise RuntimeError("Reasoner feature provider returned no result without reporting an error") + memory_info[REASONER_FEATURE_BATCH_KEY] = result + memory_info.pop(REASONER_FEATURE_FUTURE_KEY, None) + return result def pre_noise_memory_hook( self, @@ -1094,6 +1411,22 @@ def pre_noise_memory_hook( The packed sequence still contains clean tokens at this point. Override in subclasses to run a clean forward pass (e.g. for teacher forcing). """ + if self.reasoner_conditioning_backend in _EXTERNAL_REASONER_BACKENDS: + if REASONER_FEATURE_BATCH_KEY in memory_info or REASONER_FEATURE_FUTURE_KEY in memory_info: + return memory_info + if self.reasoner_feature_provider is None or self.reasoner_cache_identity is None: + raise RuntimeError( + f"reasoner_conditioning.backend={self.reasoner_conditioning_backend!r} has no feature provider" + ) + sample_keys = memory_info.get(REASONER_SAMPLE_KEYS_KEY) + if not isinstance(sample_keys, (list, tuple)): + raise RuntimeError(f"memory_info[{REASONER_SAMPLE_KEYS_KEY!r}] is missing") + requests = build_reasoner_feature_requests( + packed_sequence, + tuple(str(key) for key in sample_keys), + self.reasoner_cache_identity, + ) + memory_info[REASONER_FEATURE_FUTURE_KEY] = self.reasoner_feature_provider.submit(requests) return memory_info def _prepare_training_data( @@ -1498,11 +1831,7 @@ def training_step( packed_sequence.to_cuda() # Network forward pass - memory = self.build_memory_state(packed_sequence, memory_info) # pylint: disable=assignment-from-none - out_net = self.denoise( - data_batch_packed=packed_sequence, - memory=memory, - ) + out_net = self._denoise_training_with_reasoner_conditioning(packed_sequence, memory_info) loss, losses_dict = self._compute_losses( out_net=out_net, diff --git a/cosmos_framework/model/generator/omni_mot_reasoner_conditioning_test.py b/cosmos_framework/model/generator/omni_mot_reasoner_conditioning_test.py new file mode 100644 index 000000000..7c10e0f2b --- /dev/null +++ b/cosmos_framework/model/generator/omni_mot_reasoner_conditioning_test.py @@ -0,0 +1,390 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +from concurrent.futures import Future +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest +import torch + +from cosmos_framework.callbacks.load_pretrained import ( + _warm_start_partially_skips_ema, + _warm_start_skips_complete_ema, +) +from cosmos_framework.data.generator.sequence_packing import PackedSequence +from cosmos_framework.model.generator.omni_mot_model import ( + REASONER_FEATURE_BATCH_KEY, + REASONER_FEATURE_FUTURE_KEY, + OmniMoTModel, + _reasoner_cache_identity, + _validate_reasoner_conditioning, +) +from cosmos_framework.model.generator.reasoner_feature_cache import ( + OfflineReasonerFeatureProvider, + ReasonerFeatureCacheIdentity, + build_reasoner_feature_requests, +) +from cosmos_framework.model.generator.reasoner_features import ( + CapturingReasonerKVMemoryState, + ReasonerFeatureBatch, + ReasonerLayerKV, + StaticReasonerKVMemoryState, +) + + +def _config(backend: str, **conditioning_overrides: object) -> SimpleNamespace: + conditioning = dict( + backend=backend, + cache_root=None, + endpoint=None, + reasoner_fingerprint=None, + tokenizer_fingerprint=None, + framing_fingerprint=None, + strict_fingerprint=True, + layerwise_h2d=False, + ) + conditioning.update(conditioning_overrides) + return SimpleNamespace( + reasoner_conditioning=conditioning, + joint_attn_implementation="two_way", + parallelism=SimpleNamespace(context_parallel_shard_degree=1), + video_temporal_causal=False, + causal_training_strategy="none", + ) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_reasoner_conditioning_validation_fails_closed() -> None: + assert _validate_reasoner_conditioning(_config("joint")) == "joint" + assert _validate_reasoner_conditioning(_config("inline")) == "inline" + + with pytest.raises(ValueError, match="requires cache_root"): + _validate_reasoner_conditioning(_config("offline")) + with pytest.raises(ValueError, match="requires endpoint"): + _validate_reasoner_conditioning(_config("remote")) + + with pytest.raises(ValueError, match="strict_fingerprint=true"): + _validate_reasoner_conditioning(_config("offline", cache_root="/features")) + with pytest.raises(ValueError, match="requires strict_fingerprint=true"): + _validate_reasoner_conditioning( + _config( + "offline", + cache_root="/features", + reasoner_fingerprint="reasoner", + tokenizer_fingerprint="tokenizer", + framing_fingerprint="framing", + strict_fingerprint=False, + ) + ) + assert ( + _validate_reasoner_conditioning( + _config( + "offline", + cache_root="/features", + reasoner_fingerprint="reasoner", + tokenizer_fingerprint="tokenizer", + framing_fingerprint="framing", + ) + ) + == "offline" + ) + + config = _config("inline") + config.parallelism.context_parallel_shard_degree = 2 + with pytest.raises(ValueError, match="context parallel degree 1"): + _validate_reasoner_conditioning(config) + + config = _config("inline") + config.causal_training_strategy = "teacher_forcing" + with pytest.raises(ValueError, match="causal_training_strategy='none'"): + _validate_reasoner_conditioning(config) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_non_external_backend_ignores_partial_cache_identity() -> None: + assert _reasoner_cache_identity(_config("joint", reasoner_fingerprint="bookkeeping-only")) is None + assert _reasoner_cache_identity(_config("inline", reasoner_fingerprint="bookkeeping-only")) is None + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_causal_model_rejects_uncomposed_reasoner_feature_dispatch() -> None: + from cosmos_framework.model.generator.omni_mot_causal_model import OmniMoTCausalModel + + config = SimpleNamespace(reasoner_conditioning={"backend": "offline"}) + with pytest.raises(ValueError, match="OmniMoTCausalModel currently supports only"): + OmniMoTCausalModel(config) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_non_base_model_subclass_rejects_uncomposed_reasoner_lifecycle() -> None: + class UnsupportedModel(OmniMoTModel): + pass + + config = _config( + "offline", + cache_root="/features", + reasoner_fingerprint="reasoner", + tokenizer_fingerprint="tokenizer", + framing_fingerprint="framing", + ) + with pytest.raises(ValueError, match="does not yet compose its overridden lifecycle"): + UnsupportedModel(config) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_offline_cache_signature_is_checked_before_model_materialization() -> None: + model = object.__new__(OmniMoTModel) + provider = MagicMock(spec=OfflineReasonerFeatureProvider) + provider.manifest = SimpleNamespace(num_layers=3, num_kv_heads=2, head_dim=8) + model.reasoner_feature_provider = provider + net = SimpleNamespace(num_hidden_layers=3, num_kv_heads=2, head_dim=8) + + model._validate_reasoner_feature_signature(net) + provider.manifest.num_layers = 2 + with pytest.raises(ValueError, match="architecture mismatch"): + model._validate_reasoner_feature_signature(net) + + +class _MemoryBuilder: + build_memory_state = OmniMoTModel.build_memory_state + + def __init__(self, backend: str) -> None: + self.reasoner_conditioning_backend = backend + self.net = SimpleNamespace( + language_model=SimpleNamespace(model=SimpleNamespace(layers=[object(), object()])), + ) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_build_memory_state_selects_joint_inline_and_external_modes() -> None: + packed = SimpleNamespace() + + assert _MemoryBuilder("joint").build_memory_state(packed, {}) is None + inline = _MemoryBuilder("inline").build_memory_state(packed, {}) + assert isinstance(inline, CapturingReasonerKVMemoryState) + + external = _MemoryBuilder("offline") + with pytest.raises(RuntimeError, match=REASONER_FEATURE_BATCH_KEY): + external.build_memory_state(packed, {}) + + features = ReasonerFeatureBatch( + cross_k=(torch.zeros(3, 1, 2), torch.zeros(3, 1, 2)), + cross_v=(torch.zeros(3, 1, 2), torch.zeros(3, 1, 2)), + causal_offsets=torch.tensor([0, 3], dtype=torch.int32), + ) + state = external.build_memory_state(packed, {REASONER_FEATURE_BATCH_KEY: features}) + assert isinstance(state, StaticReasonerKVMemoryState) + + +class _InlineDenoiser(_MemoryBuilder): + _denoise_training_with_reasoner_conditioning = OmniMoTModel._denoise_training_with_reasoner_conditioning + + def __init__(self) -> None: + super().__init__("inline") + self.calls: list[tuple[bool, bool]] = [] + + def denoise(self, *, data_batch_packed: object, memory: object) -> dict[str, str]: + del data_batch_packed + assert isinstance(memory, CapturingReasonerKVMemoryState) + self.calls.append((memory.is_gen_only(), torch.is_grad_enabled())) + if not memory.is_gen_only(): + memory._causal_offsets = torch.tensor([0, 2], dtype=torch.int32) + layer = ReasonerLayerKV(torch.zeros(2, 1, 2), torch.zeros(2, 1, 2)) + memory._layers = [layer, layer] + return {"phase": "capture"} + return {"phase": "train"} + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_inline_training_runs_capture_then_generator_only() -> None: + model = _InlineDenoiser() + + output = model._denoise_training_with_reasoner_conditioning(SimpleNamespace(), {}) + + assert model.calls == [(False, False), (True, True)] + assert output == {"phase": "train"} + + +class _ExternalCheckpointLoader: + load_pretrained_model_if_needed = OmniMoTModel.load_pretrained_model_if_needed + + def __init__(self, *, copy_from_reasoner: bool) -> None: + self.reasoner_conditioning_backend = "offline" + self.net = object() + self.net_ema = object() + self.net_ema_worker = MagicMock() + self.config = SimpleNamespace( + diffusion_expert_config=SimpleNamespace(load_weights_from_pretrained=copy_from_reasoner), + ema=SimpleNamespace(enabled=True), + ) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_external_backend_requires_generator_checkpoint_for_reasoner_copy() -> None: + loader = _ExternalCheckpointLoader(copy_from_reasoner=True) + + with pytest.raises(ValueError, match="checkpoint.load_path"): + loader.load_pretrained_model_if_needed(has_resumable_checkpoint=False, has_load_path=False) + + loader.load_pretrained_model_if_needed( + has_resumable_checkpoint=False, + has_load_path=True, + warm_start_ema_skipped=True, + ) + + loader.net_ema_worker.copy_to.assert_called_once_with(src_model=loader.net, tgt_model=loader.net_ema) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_external_backend_preserves_loaded_warm_start_ema() -> None: + loader = _ExternalCheckpointLoader(copy_from_reasoner=True) + + loader.load_pretrained_model_if_needed( + has_resumable_checkpoint=False, + has_load_path=True, + warm_start_ema_skipped=False, + ) + + loader.net_ema_worker.copy_to.assert_not_called() + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_warm_start_ema_skip_detection_requires_a_root_wide_pattern() -> None: + assert _warm_start_skips_complete_ema(["net_ema."]) + assert _warm_start_skips_complete_ema(["ema"]) + assert not _warm_start_skips_complete_ema([]) + assert not _warm_start_skips_complete_ema(["net_ema.language_model.model.layers.0"]) + assert not _warm_start_skips_complete_ema(["model.net_ema."]) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_warm_start_partial_ema_skip_detection_uses_actual_state_leaves() -> None: + ema_fqns = [ + "net_ema.language_model.layer.0.weight", + "net_ema.language_model.layer.1.weight", + "net_ema.vfm.weight", + ] + + assert _warm_start_partially_skips_ema(["layer.0"], ema_fqns) + assert _warm_start_partially_skips_ema(["language_model"], ema_fqns) + assert not _warm_start_partially_skips_ema(["net_ema."], ema_fqns) + assert _warm_start_skips_complete_ema(["language_model", "vfm"], ema_fqns) + assert not _warm_start_partially_skips_ema(["missing"], ema_fqns) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_external_backend_rejects_partial_ema_warm_start() -> None: + loader = _ExternalCheckpointLoader(copy_from_reasoner=True) + + with pytest.raises(ValueError, match="matches only part of the 'net_ema.' subtree"): + loader.load_pretrained_model_if_needed( + has_resumable_checkpoint=False, + has_load_path=True, + warm_start_ema_partially_skipped=True, + warm_start_strict_resume=True, + ) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_external_backend_rejects_ambiguous_non_strict_ema_warm_start() -> None: + loader = _ExternalCheckpointLoader(copy_from_reasoner=True) + + with pytest.raises(ValueError, match="cannot safely use a non-strict warm start"): + loader.load_pretrained_model_if_needed( + has_resumable_checkpoint=False, + has_load_path=True, + warm_start_ema_skipped=False, + warm_start_strict_resume=False, + ) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_external_backend_allows_explicit_random_generator_initialization() -> None: + loader = _ExternalCheckpointLoader(copy_from_reasoner=False) + + loader.load_pretrained_model_if_needed(has_resumable_checkpoint=False, has_load_path=False) + + loader.net_ema_worker.copy_to.assert_not_called() + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_external_backend_resume_preserves_checkpoint_ema() -> None: + loader = _ExternalCheckpointLoader(copy_from_reasoner=True) + + # ``load_path`` may remain configured while a latest same-job checkpoint + # takes precedence. That resume restores EMA and must not reset it from net. + loader.load_pretrained_model_if_needed(has_resumable_checkpoint=True, has_load_path=True) + + loader.net_ema_worker.copy_to.assert_not_called() + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_build_reasoner_requests_uses_framed_text_positions_and_sample_offsets() -> None: + packed = PackedSequence( + sample_lens=[5, 5], + split_lens=[2, 3, 3, 2], + attn_modes=["causal", "full", "causal", "full"], + sequence_length=10, + text_ids=torch.tensor([11, 12, 21, 22, 23]), + text_indexes=torch.tensor([0, 1, 5, 6, 7]), + position_ids=torch.arange(30, dtype=torch.float32).reshape(3, 10), + ) + identity = ReasonerFeatureCacheIdentity("reasoner", "tokenizer", "framing") + + requests = build_reasoner_feature_requests(packed, ("first", "second"), identity) + + assert [request.sample_key for request in requests] == ["first", "second"] + assert requests[0].token_ids.tolist() == [11, 12] + assert requests[1].token_ids.tolist() == [21, 22, 23] + assert requests[0].position_ids.shape == (3, 2) + assert requests[1].position_ids.shape == (3, 3) + assert requests[0].causal_offsets.tolist() == [0, 2] + assert requests[1].causal_offsets.tolist() == [0, 3] + assert requests[0].fingerprint != requests[1].fingerprint + + +class _FeatureResolver: + _resolve_reasoner_feature_batch = OmniMoTModel._resolve_reasoner_feature_batch + + def __init__(self) -> None: + self.config = SimpleNamespace(reasoner_conditioning={"request_timeout_s": 1.0}) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_resolve_reasoner_features_consumes_future_and_surfaces_failure() -> None: + features = ReasonerFeatureBatch( + cross_k=(torch.zeros(2, 1, 2),), + cross_v=(torch.zeros(2, 1, 2),), + causal_offsets=torch.tensor([0, 2]), + fingerprints=("fingerprint",), + ) + successful: Future[ReasonerFeatureBatch] = Future() + successful.set_result(features) + memory_info = {REASONER_FEATURE_FUTURE_KEY: successful} + + assert _FeatureResolver()._resolve_reasoner_feature_batch(memory_info) is features + assert memory_info[REASONER_FEATURE_BATCH_KEY] is features + assert REASONER_FEATURE_FUTURE_KEY not in memory_info + + failed: Future[ReasonerFeatureBatch] = Future() + failed.set_exception(KeyError("missing")) + with pytest.raises(RuntimeError, match="failed before FSDP forward"): + _FeatureResolver()._resolve_reasoner_feature_batch({REASONER_FEATURE_FUTURE_KEY: failed}) diff --git a/cosmos_framework/model/generator/reasoner_feature_cache.py b/cosmos_framework/model/generator/reasoner_feature_cache.py new file mode 100644 index 000000000..29b346933 --- /dev/null +++ b/cosmos_framework/model/generator/reasoner_feature_cache.py @@ -0,0 +1,1663 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +"""Immutable, sharded on-disk cache for frozen Reasoner K/V features. + +The cache is deliberately provider-neutral and training-safe: + +* tensors are stored layer-major in multi-record safetensors shards; +* the manifest is published only after every shard is complete; +* model/tokenizer/framing identities and per-record fingerprints fail closed; +* shard checksums are verified before the first read; and +* :class:`OfflineReasonerFeatureProvider` implements the same completed-Future + interface used by future asynchronous/remote providers. + +The incremental writer/finalizer supports rank-local extraction onto a shared +POSIX filesystem without owning dataset enumeration or process launch. Those +orchestration concerns belong in the extraction CLI. +""" + +from __future__ import annotations + +import hashlib +import json +import os +import shutil +import tempfile +from collections import defaultdict +from concurrent.futures import Future +from contextlib import suppress +from dataclasses import asdict, dataclass +from errno import EEXIST, ENOTEMPTY +from pathlib import Path +from typing import Any, Mapping, Sequence + +import torch +from safetensors import safe_open +from safetensors.torch import save_file + +from cosmos_framework.data.generator.sequence_packing import PackedSequence +from cosmos_framework.data.generator.sequence_packing.sequence import PackedSequenceBuilder +from cosmos_framework.model.generator.reasoner_features import ( + ReasonerFeatureBatch, + ReasonerFeatureRequest, +) + +REASONER_FEATURE_CACHE_SCHEMA_VERSION = 1 +REASONER_FEATURE_CACHE_MANIFEST = "manifest.json" +_FORMAT_NAME = "cosmos3-reasoner-kv" +_OFFSETS_KEY = "record_offsets" +_INCREMENTAL_SCHEMA_VERSION = 1 +_INCREMENTAL_STAGING_FILE = "staging.json" +_INCREMENTAL_RANK_COMPLETE_FILE = "rank.complete.json" +_INCREMENTAL_SIDECAR_SUFFIX = ".index.json" + +_CacheSignature = tuple[int, int, int, str] + + +@dataclass(frozen=True) +class ReasonerFeatureCacheIdentity: + """Immutable global identity of every record in one cache.""" + + reasoner: str + tokenizer: str + framing: str + + def __post_init__(self) -> None: + for field_name, value in asdict(self).items(): + if not isinstance(value, str) or not value: + raise ValueError(f"Cache identity field {field_name!r} must be a non-empty string") + + +@dataclass(frozen=True) +class ReasonerFeatureCacheEntry: + """One cache record; ``features`` must describe exactly one sample.""" + + sample_key: str + features: ReasonerFeatureBatch + + def __post_init__(self) -> None: + if not self.sample_key: + raise ValueError("sample_key must be non-empty") + if self.features.num_samples != 1: + raise ValueError( + f"Cache entries must contain exactly one sample, got {self.features.num_samples} for {self.sample_key!r}" + ) + if len(self.features.fingerprints) != 1: + raise ValueError(f"Cache entry {self.sample_key!r} must carry exactly one non-empty fingerprint") + + +@dataclass(frozen=True) +class _Record: + sample_key: str + start: int + end: int + fingerprint: str + + @property + def num_tokens(self) -> int: + return self.end - self.start + + +@dataclass(frozen=True) +class _Shard: + path: str + sha256: str + num_tokens: int + records: tuple[_Record, ...] + + +@dataclass(frozen=True) +class _Manifest: + cache_fingerprint: str + identity: ReasonerFeatureCacheIdentity + dtype: str + num_layers: int + num_kv_heads: int + head_dim: int + shards: tuple[_Shard, ...] + + +@dataclass(frozen=True) +class _IncrementalShard: + """One committed rank-local shard and its durable index sidecar.""" + + rank: int + world_size: int + shard_index: int + signature: _CacheSignature + cache_fingerprint: str + shard: _Shard + shard_path: Path + sidecar_path: Path + + +def _layer_key(kind: str, layer_idx: int) -> str: + return f"{kind}.{layer_idx:03d}" + + +def _dtype_name(dtype: torch.dtype) -> str: + name = str(dtype).removeprefix("torch.") + if name not in {"bfloat16", "float16", "float32"}: + raise TypeError(f"Unsupported Reasoner cache dtype: {dtype}") + return name + + +def _manifest_fingerprint( + identity: ReasonerFeatureCacheIdentity, + *, + dtype: str, + num_layers: int, + num_kv_heads: int, + head_dim: int, +) -> str: + payload = { + "schema_version": REASONER_FEATURE_CACHE_SCHEMA_VERSION, + "format": _FORMAT_NAME, + "identity": asdict(identity), + "dtype": dtype, + "num_layers": num_layers, + "num_kv_heads": num_kv_heads, + "head_dim": head_dim, + } + encoded = json.dumps(payload, sort_keys=True, separators=(",", ":")).encode("utf-8") + return hashlib.sha256(encoded).hexdigest() + + +def _sha256_file(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as handle: + while chunk := handle.read(8 * 1024 * 1024): + digest.update(chunk) + return digest.hexdigest() + + +def _fsync_directory(path: Path) -> None: + """Persist directory-entry changes on the shared POSIX filesystem.""" + flags = os.O_RDONLY | getattr(os, "O_DIRECTORY", 0) + descriptor = os.open(path, flags) + try: + os.fsync(descriptor) + finally: + os.close(descriptor) + + +def _write_json_file(path: Path, payload: Mapping[str, Any]) -> None: + with path.open("w", encoding="utf-8") as handle: + json.dump(payload, handle, indent=2, sort_keys=True) + handle.write("\n") + handle.flush() + os.fsync(handle.fileno()) + + +def _atomic_write_json(path: Path, payload: Mapping[str, Any]) -> None: + """Write a JSON commit marker through a same-directory atomic rename.""" + path.parent.mkdir(parents=True, exist_ok=True) + descriptor, temporary_name = tempfile.mkstemp(prefix=f".{path.name}.tmp-", dir=path.parent) + os.close(descriptor) + temporary_path = Path(temporary_name) + try: + _write_json_file(temporary_path, payload) + os.replace(temporary_path, path) + _fsync_directory(path.parent) + except BaseException: + with suppress(FileNotFoundError): + temporary_path.unlink() + raise + + +def _atomic_create_json(path: Path, payload: Mapping[str, Any]) -> bool: + """Atomically publish a JSON file only if no peer has published it first. + + The hard-link commit is the shared-POSIX equivalent of ``O_EXCL`` while + retaining the requested temp-file + atomic-publish discipline. + """ + path.parent.mkdir(parents=True, exist_ok=True) + descriptor, temporary_name = tempfile.mkstemp(prefix=f".{path.name}.tmp-", dir=path.parent) + os.close(descriptor) + temporary_path = Path(temporary_name) + try: + _write_json_file(temporary_path, payload) + try: + os.link(temporary_path, path) + created = True + except FileExistsError: + created = False + if created: + _fsync_directory(path.parent) + return created + finally: + with suppress(FileNotFoundError): + temporary_path.unlink() + + +def _incremental_staging_root(cache_root: Path) -> Path: + return cache_root.parent / f".{cache_root.name}.reasoner-kv-staging" + + +def _signature_to_dict(signature: _CacheSignature) -> dict[str, int | str]: + num_layers, num_kv_heads, head_dim, dtype = signature + return { + "dtype": dtype, + "num_layers": num_layers, + "num_kv_heads": num_kv_heads, + "head_dim": head_dim, + } + + +def _signature_from_dict(value: Any, *, field: str) -> _CacheSignature: + raw = _require_dict(value, field=field) + dtype = raw.get("dtype") + if dtype not in {"bfloat16", "float16", "float32"}: + raise ValueError(f"Incremental cache field {field!r} has unsupported dtype {dtype!r}") + return ( + _require_positive_int(raw.get("num_layers"), field=f"{field}.num_layers"), + _require_positive_int(raw.get("num_kv_heads"), field=f"{field}.num_kv_heads"), + _require_positive_int(raw.get("head_dim"), field=f"{field}.head_dim"), + dtype, + ) + + +def compute_reasoner_feature_fingerprint( + token_ids: torch.Tensor, + position_ids: torch.Tensor, + causal_offsets: torch.Tensor, + *, + identity: ReasonerFeatureCacheIdentity, +) -> str: + """Hash the exact framed Reasoner input plus its immutable global identity.""" + + digest = hashlib.sha256() + digest.update(f"{_FORMAT_NAME}:{REASONER_FEATURE_CACHE_SCHEMA_VERSION}\n".encode()) + digest.update(json.dumps(asdict(identity), sort_keys=True, separators=(",", ":")).encode()) + for name, tensor in ( + ("token_ids", token_ids), + ("position_ids", position_ids), + ("causal_offsets", causal_offsets), + ): + if not isinstance(tensor, torch.Tensor) or tensor.device.type == "meta": + raise TypeError(f"{name} must be a materialized torch.Tensor") + value = tensor.detach().to(device="cpu").contiguous() + digest.update(name.encode()) + digest.update(str(value.dtype).encode()) + digest.update(json.dumps(list(value.shape), separators=(",", ":")).encode()) + digest.update(value.view(torch.uint8).numpy().tobytes()) + return digest.hexdigest() + + +def build_reasoner_feature_requests( + packed_sequence: PackedSequence, + sample_keys: Sequence[str], + identity: ReasonerFeatureCacheIdentity, +) -> tuple[ReasonerFeatureRequest, ...]: + """Build exact per-sample requests from the finalized CPU training pack. + + Producers must call the same sequence packer first and use this function + instead of reproducing BOS/EOS/start-of-generation framing or mRoPE + positions independently. + """ + packed_sequence.prepare_sequence_pack_metadata() + metadata = packed_sequence.get_sequence_pack_metadata() + if metadata is None: + raise RuntimeError("PackedSequence failed to prepare Reasoner attention metadata") + offsets = metadata.causal_seq_offsets.detach().to(device="cpu", dtype=torch.int64) + if len(sample_keys) != offsets.numel() - 1: + raise ValueError( + f"Reasoner sample keys ({len(sample_keys)}) do not match packed causal segments ({offsets.numel() - 1})" + ) + + text_ids = packed_sequence.text_ids.detach().to(device="cpu") + text_indexes = packed_sequence.text_indexes.detach().to(device="cpu", dtype=torch.int64) + position_ids = packed_sequence.position_ids.detach().to(device="cpu").index_select(-1, text_indexes) + if text_ids.numel() != int(offsets[-1]) or position_ids.shape[-1] != text_ids.numel(): + raise ValueError( + "External Reasoner conditioning requires every causal token to be a framed text token; " + f"text={text_ids.numel()} positions={position_ids.shape[-1]} causal={int(offsets[-1])}" + ) + + requests: list[ReasonerFeatureRequest] = [] + for sample_idx, sample_key in enumerate(sample_keys): + start = int(offsets[sample_idx]) + end = int(offsets[sample_idx + 1]) + token_slice = text_ids[start:end] + position_slice = position_ids[..., start:end] + sample_offsets = torch.tensor([0, end - start], dtype=torch.int64) + fingerprint = compute_reasoner_feature_fingerprint( + token_slice, + position_slice, + sample_offsets, + identity=identity, + ) + requests.append( + ReasonerFeatureRequest( + sample_key=str(sample_key), + token_ids=token_slice, + position_ids=position_slice, + causal_offsets=sample_offsets, + fingerprint=fingerprint, + ) + ) + return tuple(requests) + + +def build_reasoner_feature_request_from_text_tokens( + *, + sample_key: str, + text_ids: Sequence[int] | torch.Tensor, + special_tokens: Mapping[str, int], + use_float_positions: bool, + identity: ReasonerFeatureCacheIdentity, +) -> ReasonerFeatureRequest: + """Frame one ordinary SFT caption through the canonical sequence builder. + + The helper is for deterministic offline producers. It deliberately covers + only the currently supported single-caption, non-AR SFT layout and always + includes the start-of-generation token used by a video sample. + """ + raw_text_ids = ( + text_ids.detach().to(device="cpu", dtype=torch.int64).tolist() + if isinstance(text_ids, torch.Tensor) + else list(text_ids) + ) + builder = PackedSequenceBuilder() + builder.begin_sample(initial_mrope_temporal_offset=0) + split_len = builder.pack_text_tokens( + raw_text_ids, + dict(special_tokens), + has_generation=True, + use_float_positions=use_float_positions, + ) + framed_ids = torch.tensor(builder.text_ids, dtype=torch.int64) + if len(builder.position_ids) != 1: + raise RuntimeError(f"Expected one framed text position block, got {len(builder.position_ids)}") + position_ids = builder.position_ids[0] + causal_offsets = torch.tensor([0, split_len], dtype=torch.int64) + fingerprint = compute_reasoner_feature_fingerprint( + framed_ids, + position_ids, + causal_offsets, + identity=identity, + ) + return ReasonerFeatureRequest( + sample_key=sample_key, + token_ids=framed_ids, + position_ids=position_ids, + causal_offsets=causal_offsets, + fingerprint=fingerprint, + ) + + +def _entry_bytes(entry: ReasonerFeatureCacheEntry) -> int: + return sum(tensor.numel() * tensor.element_size() for tensor in (*entry.features.cross_k, *entry.features.cross_v)) + + +def _partition_entries( + entries: Sequence[ReasonerFeatureCacheEntry], + max_shard_bytes: int, +) -> list[list[ReasonerFeatureCacheEntry]]: + if max_shard_bytes <= 0: + raise ValueError(f"max_shard_bytes must be positive, got {max_shard_bytes}") + shards: list[list[ReasonerFeatureCacheEntry]] = [] + current: list[ReasonerFeatureCacheEntry] = [] + current_bytes = 0 + for entry in entries: + size = _entry_bytes(entry) + if current and current_bytes + size > max_shard_bytes: + shards.append(current) + current = [] + current_bytes = 0 + current.append(entry) + current_bytes += size + if current: + shards.append(current) + return shards + + +def _validate_entries(entries: Sequence[ReasonerFeatureCacheEntry]) -> _CacheSignature: + if not entries: + raise ValueError("At least one Reasoner feature cache entry is required") + identities = [(entry.sample_key, entry.features.fingerprints[0]) for entry in entries] + duplicates = sorted(identity for identity in set(identities) if identities.count(identity) > 1) + if duplicates: + raise ValueError(f"Duplicate Reasoner feature sample/fingerprint pairs: {duplicates}") + + first = entries[0].features + first_layer = first.layer(0) + signature = ( + first.num_layers, + first_layer.num_kv_heads, + first_layer.head_dim, + _dtype_name(first_layer.cross_k.dtype), + ) + for entry in entries[1:]: + features = entry.features + layer = features.layer(0) + candidate = ( + features.num_layers, + layer.num_kv_heads, + layer.head_dim, + _dtype_name(layer.cross_k.dtype), + ) + if candidate != signature: + raise ValueError( + f"Reasoner cache entry {entry.sample_key!r} has signature {candidate}, expected {signature}" + ) + return signature + + +def _write_shard( + path: Path, + entries: Sequence[ReasonerFeatureCacheEntry], + *, + cache_fingerprint: str, +) -> _Shard: + offsets = [0] + records: list[_Record] = [] + for entry in entries: + start = offsets[-1] + end = start + entry.features.sequence_length + offsets.append(end) + records.append( + _Record( + sample_key=entry.sample_key, + start=start, + end=end, + fingerprint=entry.features.fingerprints[0], + ) + ) + + num_layers = entries[0].features.num_layers + tensors: dict[str, torch.Tensor] = { + _OFFSETS_KEY: torch.tensor(offsets, dtype=torch.int64), + } + for layer_idx in range(num_layers): + tensors[_layer_key("cross_k", layer_idx)] = torch.cat( + [entry.features.cross_k[layer_idx].detach().to(device="cpu").contiguous() for entry in entries], + dim=0, + ) + tensors[_layer_key("cross_v", layer_idx)] = torch.cat( + [entry.features.cross_v[layer_idx].detach().to(device="cpu").contiguous() for entry in entries], + dim=0, + ) + + save_file( + tensors, + str(path), + metadata={ + "format": _FORMAT_NAME, + "schema_version": str(REASONER_FEATURE_CACHE_SCHEMA_VERSION), + "cache_fingerprint": cache_fingerprint, + }, + ) + return _Shard( + path=path.name, + sha256=_sha256_file(path), + num_tokens=offsets[-1], + records=tuple(records), + ) + + +def write_reasoner_feature_cache( + cache_root: str | Path, + entries: Sequence[ReasonerFeatureCacheEntry], + *, + identity: ReasonerFeatureCacheIdentity, + max_shard_bytes: int = 4 * 1024**3, +) -> Path: + """Atomically publish a new immutable Reasoner feature cache. + + ``cache_root`` must not already exist. The complete cache is first written + to a sibling temporary directory and becomes visible through one atomic + rename, so readers never observe a partial manifest or shard set. + """ + root = Path(cache_root) + if root.exists(): + raise FileExistsError(f"Refusing to overwrite existing Reasoner feature cache: {root}") + num_layers, num_kv_heads, head_dim, dtype = _validate_entries(entries) + cache_fingerprint = _manifest_fingerprint( + identity, + dtype=dtype, + num_layers=num_layers, + num_kv_heads=num_kv_heads, + head_dim=head_dim, + ) + entry_shards = _partition_entries(entries, max_shard_bytes) + + root.parent.mkdir(parents=True, exist_ok=True) + temporary_root = Path(tempfile.mkdtemp(prefix=f".{root.name}.tmp-", dir=root.parent)) + try: + shards: list[_Shard] = [] + for shard_idx, shard_entries in enumerate(entry_shards): + shard_name = f"reasoner-kv-{shard_idx:05d}-of-{len(entry_shards):05d}.safetensors" + shards.append( + _write_shard( + temporary_root / shard_name, + shard_entries, + cache_fingerprint=cache_fingerprint, + ) + ) + + manifest = { + "schema_version": REASONER_FEATURE_CACHE_SCHEMA_VERSION, + "format": _FORMAT_NAME, + "cache_fingerprint": cache_fingerprint, + "identity": asdict(identity), + "dtype": dtype, + "num_layers": num_layers, + "num_kv_heads": num_kv_heads, + "head_dim": head_dim, + "shards": [ + { + "path": shard.path, + "sha256": shard.sha256, + "num_tokens": shard.num_tokens, + "records": [asdict(record) for record in shard.records], + } + for shard in shards + ], + } + manifest_tmp = temporary_root / f"{REASONER_FEATURE_CACHE_MANIFEST}.tmp" + manifest_tmp.write_text(json.dumps(manifest, indent=2, sort_keys=True) + "\n", encoding="utf-8") + os.replace(manifest_tmp, temporary_root / REASONER_FEATURE_CACHE_MANIFEST) + os.replace(temporary_root, root) + except BaseException: + shutil.rmtree(temporary_root, ignore_errors=True) + raise + return root / REASONER_FEATURE_CACHE_MANIFEST + + +def _require_dict(value: Any, *, field: str) -> dict[str, Any]: + if not isinstance(value, dict): + raise ValueError(f"Manifest field {field!r} must be an object") + return value + + +def _require_positive_int(value: Any, *, field: str) -> int: + if not isinstance(value, int) or isinstance(value, bool) or value <= 0: + raise ValueError(f"Manifest field {field!r} must be a positive integer") + return value + + +def _safe_relative_path(value: Any, *, field: str) -> str: + if not isinstance(value, str) or not value: + raise ValueError(f"Manifest field {field!r} must be a non-empty string") + path = Path(value) + if path.is_absolute() or ".." in path.parts or len(path.parts) != 1: + raise ValueError(f"Manifest field {field!r} must be a simple relative filename, got {value!r}") + return value + + +def _load_manifest(cache_root: Path) -> _Manifest: + path = cache_root / REASONER_FEATURE_CACHE_MANIFEST + try: + raw = json.loads(path.read_text(encoding="utf-8")) + except FileNotFoundError as error: + raise FileNotFoundError(f"Reasoner feature cache manifest not found: {path}") from error + except json.JSONDecodeError as error: + raise ValueError(f"Invalid Reasoner feature cache manifest JSON: {path}") from error + raw = _require_dict(raw, field="root") + if raw.get("schema_version") != REASONER_FEATURE_CACHE_SCHEMA_VERSION: + raise ValueError( + f"Unsupported Reasoner cache schema_version={raw.get('schema_version')!r}; " + f"expected {REASONER_FEATURE_CACHE_SCHEMA_VERSION}" + ) + if raw.get("format") != _FORMAT_NAME: + raise ValueError(f"Unexpected Reasoner cache format={raw.get('format')!r}") + + identity_dict = _require_dict(raw.get("identity"), field="identity") + try: + identity = ReasonerFeatureCacheIdentity(**identity_dict) + except TypeError as error: + raise ValueError(f"Invalid Reasoner cache identity fields: {sorted(identity_dict)}") from error + dtype = raw.get("dtype") + if dtype not in {"bfloat16", "float16", "float32"}: + raise ValueError(f"Manifest field 'dtype' is unsupported: {dtype!r}") + num_layers = _require_positive_int(raw.get("num_layers"), field="num_layers") + num_kv_heads = _require_positive_int(raw.get("num_kv_heads"), field="num_kv_heads") + head_dim = _require_positive_int(raw.get("head_dim"), field="head_dim") + expected_cache_fingerprint = _manifest_fingerprint( + identity, + dtype=dtype, + num_layers=num_layers, + num_kv_heads=num_kv_heads, + head_dim=head_dim, + ) + if raw.get("cache_fingerprint") != expected_cache_fingerprint: + raise ValueError("Reasoner cache manifest fingerprint does not match its identity/shape fields") + + raw_shards = raw.get("shards") + if not isinstance(raw_shards, list) or not raw_shards: + raise ValueError("Manifest field 'shards' must be a non-empty list") + shards: list[_Shard] = [] + seen_records: set[tuple[str, str]] = set() + seen_paths: set[str] = set() + for shard_idx, raw_shard_value in enumerate(raw_shards): + raw_shard = _require_dict(raw_shard_value, field=f"shards[{shard_idx}]") + shard_path = _safe_relative_path(raw_shard.get("path"), field=f"shards[{shard_idx}].path") + if shard_path in seen_paths: + raise ValueError(f"Duplicate shard path in manifest: {shard_path!r}") + seen_paths.add(shard_path) + checksum = raw_shard.get("sha256") + if not isinstance(checksum, str) or len(checksum) != 64 or any(c not in "0123456789abcdef" for c in checksum): + raise ValueError(f"Invalid SHA256 for shard {shard_path!r}") + num_tokens = _require_positive_int(raw_shard.get("num_tokens"), field=f"{shard_path}.num_tokens") + raw_records = raw_shard.get("records") + if not isinstance(raw_records, list) or not raw_records: + raise ValueError(f"Shard {shard_path!r} must contain a non-empty records list") + records: list[_Record] = [] + expected_start = 0 + for record_idx, raw_record_value in enumerate(raw_records): + raw_record = _require_dict(raw_record_value, field=f"{shard_path}.records[{record_idx}]") + sample_key = raw_record.get("sample_key") + fingerprint = raw_record.get("fingerprint") + start = raw_record.get("start") + end = raw_record.get("end") + if not isinstance(sample_key, str) or not sample_key: + raise ValueError(f"Shard {shard_path!r} contains an invalid sample key") + if not isinstance(fingerprint, str) or not fingerprint: + raise ValueError(f"Cache record {sample_key!r} has an invalid fingerprint") + record_identity = (sample_key, fingerprint) + if record_identity in seen_records: + raise ValueError(f"Duplicate Reasoner feature record in manifest: {record_identity!r}") + if not isinstance(start, int) or not isinstance(end, int) or start != expected_start or end <= start: + raise ValueError( + f"Cache record {sample_key!r} has invalid/non-contiguous range [{start}, {end}); " + f"expected start {expected_start}" + ) + record = _Record(sample_key, start, end, fingerprint) + records.append(record) + seen_records.add(record_identity) + expected_start = end + if expected_start != num_tokens: + raise ValueError( + f"Shard {shard_path!r} record ranges end at {expected_start}, expected num_tokens={num_tokens}" + ) + shards.append(_Shard(shard_path, checksum, num_tokens, tuple(records))) + + return _Manifest( + cache_fingerprint=expected_cache_fingerprint, + identity=identity, + dtype=dtype, + num_layers=num_layers, + num_kv_heads=num_kv_heads, + head_dim=head_dim, + shards=tuple(shards), + ) + + +def _read_json_object(path: Path, *, description: str) -> dict[str, Any]: + try: + raw = json.loads(path.read_text(encoding="utf-8")) + except FileNotFoundError as error: + raise FileNotFoundError(f"{description} not found: {path}") from error + except json.JSONDecodeError as error: + raise ValueError(f"Invalid {description} JSON: {path}") from error + return _require_dict(raw, field=description) + + +def _incremental_shard_name(rank: int, shard_index: int) -> str: + return f"reasoner-kv-r{rank:05d}-s{shard_index:05d}.safetensors" + + +def _incremental_sidecar_name(rank: int, shard_index: int) -> str: + return f"{_incremental_shard_name(rank, shard_index)}{_INCREMENTAL_SIDECAR_SUFFIX}" + + +def _incremental_staging_payload( + identity: ReasonerFeatureCacheIdentity, + *, + world_size: int, +) -> dict[str, Any]: + return { + "schema_version": _INCREMENTAL_SCHEMA_VERSION, + "format": _FORMAT_NAME, + "kind": "distributed-staging", + "identity": asdict(identity), + "world_size": world_size, + } + + +def _validate_incremental_staging( + staging_root: Path, + *, + expected_identity: ReasonerFeatureCacheIdentity, + expected_world_size: int, +) -> None: + raw = _read_json_object(staging_root / _INCREMENTAL_STAGING_FILE, description="incremental staging metadata") + if raw.get("schema_version") != _INCREMENTAL_SCHEMA_VERSION: + raise ValueError(f"Unsupported incremental cache schema_version={raw.get('schema_version')!r}") + if raw.get("format") != _FORMAT_NAME or raw.get("kind") != "distributed-staging": + raise ValueError("Incremental staging metadata has an unexpected format or kind") + identity_dict = _require_dict(raw.get("identity"), field="incremental identity") + try: + identity = ReasonerFeatureCacheIdentity(**identity_dict) + except TypeError as error: + raise ValueError(f"Invalid incremental cache identity fields: {sorted(identity_dict)}") from error + if identity != expected_identity: + raise ValueError(f"Incremental cache identity mismatch: staging={identity!r}, expected={expected_identity!r}") + world_size = _require_positive_int(raw.get("world_size"), field="incremental world_size") + if world_size != expected_world_size: + raise ValueError(f"Incremental cache world-size mismatch: staging={world_size}, expected={expected_world_size}") + + +def _ensure_incremental_staging( + cache_root: Path, + *, + identity: ReasonerFeatureCacheIdentity, + world_size: int, +) -> Path: + if cache_root.exists(): + raise FileExistsError(f"Reasoner feature cache is already published and immutable: {cache_root}") + staging_root = _incremental_staging_root(cache_root) + staging_root.mkdir(parents=True, exist_ok=True) + _fsync_directory(staging_root.parent) + payload = _incremental_staging_payload(identity, world_size=world_size) + _atomic_create_json(staging_root / _INCREMENTAL_STAGING_FILE, payload) + _validate_incremental_staging( + staging_root, + expected_identity=identity, + expected_world_size=world_size, + ) + return staging_root + + +def _parse_incremental_records( + value: Any, + *, + shard_name: str, + num_tokens: int, +) -> tuple[_Record, ...]: + if not isinstance(value, list) or not value: + raise ValueError(f"Incremental shard {shard_name!r} must contain a non-empty records list") + records: list[_Record] = [] + expected_start = 0 + seen: set[tuple[str, str]] = set() + for record_idx, raw_record_value in enumerate(value): + raw_record = _require_dict(raw_record_value, field=f"{shard_name}.records[{record_idx}]") + sample_key = raw_record.get("sample_key") + fingerprint = raw_record.get("fingerprint") + start = raw_record.get("start") + end = raw_record.get("end") + if not isinstance(sample_key, str) or not sample_key: + raise ValueError(f"Incremental shard {shard_name!r} contains an invalid sample key") + if not isinstance(fingerprint, str) or not fingerprint: + raise ValueError(f"Incremental cache record {sample_key!r} has an invalid fingerprint") + identity = (sample_key, fingerprint) + if identity in seen: + raise ValueError(f"Duplicate Reasoner feature record inside shard {shard_name!r}: {identity!r}") + if not isinstance(start, int) or not isinstance(end, int) or start != expected_start or end <= start: + raise ValueError( + f"Incremental record {sample_key!r} has invalid/non-contiguous range [{start}, {end}); " + f"expected start {expected_start}" + ) + records.append(_Record(sample_key, start, end, fingerprint)) + seen.add(identity) + expected_start = end + if expected_start != num_tokens: + raise ValueError( + f"Incremental shard {shard_name!r} record ranges end at {expected_start}, expected num_tokens={num_tokens}" + ) + return tuple(records) + + +def _validate_incremental_shard_payload( + path: Path, + *, + shard: _Shard, + signature: _CacheSignature, + cache_fingerprint: str, +) -> None: + num_layers, num_kv_heads, head_dim, dtype = signature + safe_dtype = {"bfloat16": "BF16", "float16": "F16", "float32": "F32"}[dtype] + with safe_open(str(path), framework="pt", device="cpu") as handle: + metadata = handle.metadata() or {} + if metadata.get("format") != _FORMAT_NAME: + raise ValueError(f"Incremental cache shard {path} has invalid format metadata") + if metadata.get("schema_version") != str(REASONER_FEATURE_CACHE_SCHEMA_VERSION): + raise ValueError(f"Incremental cache shard {path} has invalid schema metadata") + if metadata.get("cache_fingerprint") != cache_fingerprint: + raise ValueError(f"Incremental cache shard {path} belongs to a different cache identity/signature") + + expected_keys = {_OFFSETS_KEY} + for layer_idx in range(num_layers): + expected_keys.add(_layer_key("cross_k", layer_idx)) + expected_keys.add(_layer_key("cross_v", layer_idx)) + actual_keys = set(handle.keys()) + if actual_keys != expected_keys: + raise ValueError( + f"Incremental cache shard {path} tensor keys disagree with its sidecar: " + f"missing={sorted(expected_keys - actual_keys)}, extra={sorted(actual_keys - expected_keys)}" + ) + + offsets = handle.get_tensor(_OFFSETS_KEY) + expected_offsets = torch.tensor([0, *(record.end for record in shard.records)], dtype=torch.int64) + if not torch.equal(offsets, expected_offsets): + raise ValueError(f"Incremental cache shard {path} offsets disagree with its sidecar") + expected_shape = [shard.num_tokens, num_kv_heads, head_dim] + for layer_idx in range(num_layers): + for kind in ("cross_k", "cross_v"): + tensor_slice = handle.get_slice(_layer_key(kind, layer_idx)) + if tensor_slice.get_shape() != expected_shape: + raise ValueError( + f"Incremental cache shard {path} tensor {_layer_key(kind, layer_idx)!r} has shape " + f"{tensor_slice.get_shape()}, expected {expected_shape}" + ) + if tensor_slice.get_dtype() != safe_dtype: + raise ValueError( + f"Incremental cache shard {path} tensor {_layer_key(kind, layer_idx)!r} has dtype " + f"{tensor_slice.get_dtype()}, expected {safe_dtype}" + ) + + +def _load_incremental_shard( + sidecar_path: Path, + *, + identity: ReasonerFeatureCacheIdentity, + expected_rank: int, + expected_world_size: int, + verify_checksum: bool, +) -> _IncrementalShard: + raw = _read_json_object(sidecar_path, description="incremental shard sidecar") + if raw.get("schema_version") != _INCREMENTAL_SCHEMA_VERSION: + raise ValueError(f"Unsupported incremental sidecar schema_version={raw.get('schema_version')!r}") + if raw.get("format") != _FORMAT_NAME or raw.get("kind") != "rank-shard": + raise ValueError(f"Incremental sidecar {sidecar_path} has an unexpected format or kind") + rank = raw.get("rank") + world_size = raw.get("world_size") + shard_index = raw.get("shard_index") + if not isinstance(rank, int) or isinstance(rank, bool) or rank != expected_rank: + raise ValueError(f"Incremental sidecar {sidecar_path} has rank={rank!r}, expected {expected_rank}") + if not isinstance(world_size, int) or isinstance(world_size, bool) or world_size != expected_world_size: + raise ValueError( + f"Incremental sidecar {sidecar_path} has world_size={world_size!r}, expected {expected_world_size}" + ) + if not isinstance(shard_index, int) or isinstance(shard_index, bool) or shard_index < 0: + raise ValueError(f"Incremental sidecar {sidecar_path} has invalid shard_index={shard_index!r}") + expected_sidecar_name = _incremental_sidecar_name(rank, shard_index) + if sidecar_path.name != expected_sidecar_name: + raise ValueError( + f"Incremental sidecar name {sidecar_path.name!r} does not match rank/index {expected_sidecar_name!r}" + ) + + signature = _signature_from_dict(raw.get("signature"), field=f"{sidecar_path.name}.signature") + num_layers, num_kv_heads, head_dim, dtype = signature + expected_cache_fingerprint = _manifest_fingerprint( + identity, + dtype=dtype, + num_layers=num_layers, + num_kv_heads=num_kv_heads, + head_dim=head_dim, + ) + if raw.get("cache_fingerprint") != expected_cache_fingerprint: + raise ValueError(f"Incremental sidecar {sidecar_path} has a stale cache fingerprint") + + raw_shard = _require_dict(raw.get("shard"), field=f"{sidecar_path.name}.shard") + shard_name = _safe_relative_path(raw_shard.get("path"), field=f"{sidecar_path.name}.shard.path") + expected_shard_name = _incremental_shard_name(rank, shard_index) + if shard_name != expected_shard_name: + raise ValueError(f"Incremental sidecar points to {shard_name!r}, expected {expected_shard_name!r}") + checksum = raw_shard.get("sha256") + if not isinstance(checksum, str) or len(checksum) != 64 or any(c not in "0123456789abcdef" for c in checksum): + raise ValueError(f"Invalid SHA256 for incremental shard {shard_name!r}") + num_tokens = _require_positive_int(raw_shard.get("num_tokens"), field=f"{shard_name}.num_tokens") + records = _parse_incremental_records( + raw_shard.get("records"), + shard_name=shard_name, + num_tokens=num_tokens, + ) + shard = _Shard(shard_name, checksum, num_tokens, records) + shard_path = sidecar_path.parent / shard_name + if not shard_path.is_file(): + raise FileNotFoundError(f"Committed incremental cache shard not found: {shard_path}") + if verify_checksum: + actual_checksum = _sha256_file(shard_path) + if actual_checksum != checksum: + raise ValueError( + f"Incremental cache shard checksum mismatch for {shard_path}: " + f"expected {checksum}, got {actual_checksum}" + ) + _validate_incremental_shard_payload( + shard_path, + shard=shard, + signature=signature, + cache_fingerprint=expected_cache_fingerprint, + ) + return _IncrementalShard( + rank=rank, + world_size=world_size, + shard_index=shard_index, + signature=signature, + cache_fingerprint=expected_cache_fingerprint, + shard=shard, + shard_path=shard_path, + sidecar_path=sidecar_path, + ) + + +def _scan_incremental_rank( + rank_dir: Path, + *, + identity: ReasonerFeatureCacheIdentity, + rank: int, + world_size: int, + verify_checksums: bool, +) -> list[_IncrementalShard]: + sidecars = sorted(rank_dir.glob(f"*{_INCREMENTAL_SIDECAR_SUFFIX}")) + shards = [ + _load_incremental_shard( + sidecar, + identity=identity, + expected_rank=rank, + expected_world_size=world_size, + verify_checksum=verify_checksums, + ) + for sidecar in sidecars + ] + shards.sort(key=lambda item: item.shard_index) + indices = [item.shard_index for item in shards] + if indices != list(range(len(shards))): + raise ValueError(f"Rank {rank} incremental shard indices must be contiguous from zero, got {indices}") + signatures = {item.signature for item in shards} + if len(signatures) > 1: + raise ValueError(f"Rank {rank} incremental shards have inconsistent signatures: {sorted(signatures)}") + seen: set[tuple[str, str]] = set() + for item in shards: + for record in item.shard.records: + record_identity = (record.sample_key, record.fingerprint) + if record_identity in seen: + raise ValueError(f"Rank {rank} contains duplicate incremental record {record_identity!r}") + seen.add(record_identity) + return shards + + +def _rank_completion_payload( + *, + rank: int, + world_size: int, + identity: ReasonerFeatureCacheIdentity, + signature: _CacheSignature | None, + shards: Sequence[_IncrementalShard], +) -> dict[str, Any]: + return { + "schema_version": _INCREMENTAL_SCHEMA_VERSION, + "format": _FORMAT_NAME, + "kind": "rank-complete", + "rank": rank, + "world_size": world_size, + "identity": asdict(identity), + "signature": None if signature is None else _signature_to_dict(signature), + "num_records": sum(len(item.shard.records) for item in shards), + "num_tokens": sum(item.shard.num_tokens for item in shards), + "sidecars": [ + { + "path": item.sidecar_path.name, + "sha256": _sha256_file(item.sidecar_path), + } + for item in shards + ], + } + + +def _validate_rank_completion( + rank_dir: Path, + *, + identity: ReasonerFeatureCacheIdentity, + rank: int, + world_size: int, + shards: Sequence[_IncrementalShard], +) -> None: + completion_path = rank_dir / _INCREMENTAL_RANK_COMPLETE_FILE + raw = _read_json_object(completion_path, description=f"rank {rank} completion marker") + if raw.get("schema_version") != _INCREMENTAL_SCHEMA_VERSION: + raise ValueError(f"Unsupported rank completion schema_version={raw.get('schema_version')!r}") + if raw.get("format") != _FORMAT_NAME or raw.get("kind") != "rank-complete": + raise ValueError(f"Rank {rank} completion marker has an unexpected format or kind") + completed_rank = raw.get("rank") + completed_world_size = raw.get("world_size") + if ( + not isinstance(completed_rank, int) + or isinstance(completed_rank, bool) + or completed_rank != rank + or not isinstance(completed_world_size, int) + or isinstance(completed_world_size, bool) + or completed_world_size != world_size + ): + raise ValueError( + f"Rank completion ownership mismatch: rank/world={completed_rank!r}/{completed_world_size!r}, " + f"expected {rank}/{world_size}" + ) + identity_dict = _require_dict(raw.get("identity"), field=f"rank {rank} completion identity") + try: + completed_identity = ReasonerFeatureCacheIdentity(**identity_dict) + except TypeError as error: + raise ValueError(f"Invalid rank {rank} completion identity fields: {sorted(identity_dict)}") from error + if completed_identity != identity: + raise ValueError(f"Rank {rank} completion identity does not match the staging identity") + + signatures = {item.signature for item in shards} + expected_signature = next(iter(signatures)) if signatures else None + raw_signature = raw.get("signature") + actual_signature = ( + None if raw_signature is None else _signature_from_dict(raw_signature, field=f"rank {rank} signature") + ) + if actual_signature != expected_signature: + raise ValueError( + f"Rank {rank} completion signature {actual_signature!r} does not match shards {expected_signature!r}" + ) + + expected_num_records = sum(len(item.shard.records) for item in shards) + expected_num_tokens = sum(item.shard.num_tokens for item in shards) + if raw.get("num_records") != expected_num_records or raw.get("num_tokens") != expected_num_tokens: + raise ValueError(f"Rank {rank} completion counts do not match its committed shards") + raw_sidecars = raw.get("sidecars") + if not isinstance(raw_sidecars, list): + raise ValueError(f"Rank {rank} completion sidecars must be a list") + expected_sidecars = [{"path": item.sidecar_path.name, "sha256": _sha256_file(item.sidecar_path)} for item in shards] + if raw_sidecars != expected_sidecars: + raise ValueError(f"Rank {rank} completion marker does not match its committed sidecar set/checksums") + + +def _entry_on_cpu(entry: ReasonerFeatureCacheEntry) -> ReasonerFeatureCacheEntry: + features = entry.features + return ReasonerFeatureCacheEntry( + sample_key=entry.sample_key, + features=ReasonerFeatureBatch( + cross_k=tuple(tensor.detach().to(device="cpu").contiguous() for tensor in features.cross_k), + cross_v=tuple(tensor.detach().to(device="cpu").contiguous() for tensor in features.cross_v), + causal_offsets=features.causal_offsets.detach().to(device="cpu", dtype=torch.int64).contiguous(), + fingerprints=features.fingerprints, + ), + ) + + +class IncrementalReasonerFeatureCacheWriter: + """Bounded-memory, resumable writer for one rank on a shared POSIX filesystem. + + A sidecar is the commit record for exactly one rank-local safetensors shard. + Construction scans committed sidecars, validates their shards, and rebuilds + the ``(sample_key, fingerprint)`` skip index. Uncommitted temporary files or + orphan shard files are ignored and can be atomically replaced on retry. + + Only one live process may own a given ``rank``. Different ranks never write + the same pathname and may append/flush concurrently. + """ + + def __init__( + self, + cache_root: str | Path, + *, + identity: ReasonerFeatureCacheIdentity, + rank: int, + world_size: int, + max_shard_bytes: int = 4 * 1024**3, + verify_checksums_on_resume: bool = True, + ) -> None: + if not isinstance(world_size, int) or isinstance(world_size, bool) or world_size <= 0: + raise ValueError(f"world_size must be positive, got {world_size}") + if not isinstance(rank, int) or isinstance(rank, bool) or not 0 <= rank < world_size: + raise ValueError(f"rank must satisfy 0 <= rank < world_size, got rank={rank}, world_size={world_size}") + if not isinstance(max_shard_bytes, int) or isinstance(max_shard_bytes, bool) or max_shard_bytes <= 0: + raise ValueError(f"max_shard_bytes must be positive, got {max_shard_bytes}") + self.cache_root = Path(cache_root) + self.identity = identity + self.rank = rank + self.world_size = world_size + self.max_shard_bytes = max_shard_bytes + self.staging_root = _ensure_incremental_staging( + self.cache_root, + identity=identity, + world_size=world_size, + ) + self.rank_dir = self.staging_root / f"rank-{rank:05d}" + self.rank_dir.mkdir(parents=True, exist_ok=True) + _fsync_directory(self.staging_root) + self._shards = _scan_incremental_rank( + self.rank_dir, + identity=identity, + rank=rank, + world_size=world_size, + verify_checksums=verify_checksums_on_resume, + ) + signatures = {item.signature for item in self._shards} + self._signature: _CacheSignature | None = next(iter(signatures)) if signatures else None + self._seen: set[tuple[str, str]] = { + (record.sample_key, record.fingerprint) for item in self._shards for record in item.shard.records + } + self._buffer: list[ReasonerFeatureCacheEntry] = [] + self._buffered_bytes = 0 + completion_path = self.rank_dir / _INCREMENTAL_RANK_COMPLETE_FILE + self._finalized = completion_path.exists() + if self._finalized: + _validate_rank_completion( + self.rank_dir, + identity=self.identity, + rank=self.rank, + world_size=self.world_size, + shards=self._shards, + ) + + @classmethod + def resume( + cls, + cache_root: str | Path, + *, + identity: ReasonerFeatureCacheIdentity, + rank: int, + world_size: int, + max_shard_bytes: int = 4 * 1024**3, + verify_checksums: bool = True, + ) -> IncrementalReasonerFeatureCacheWriter: + """Open a new or interrupted rank-local writer and validate its commits.""" + return cls( + cache_root, + identity=identity, + rank=rank, + world_size=world_size, + max_shard_bytes=max_shard_bytes, + verify_checksums_on_resume=verify_checksums, + ) + + @property + def buffered_bytes(self) -> int: + return self._buffered_bytes + + @property + def committed_records(self) -> int: + return len(self._seen) - len(self._buffer) + + @property + def is_finalized(self) -> bool: + return self._finalized + + def contains(self, sample_key: str, fingerprint: str) -> bool: + """Return whether an exact record is committed or buffered. + + Extraction loops should call this before running the frozen Reasoner so + a resumed job does not recompute already committed features. + """ + return (sample_key, fingerprint) in self._seen + + def append(self, entry: ReasonerFeatureCacheEntry) -> bool: + """Buffer one entry; return ``False`` when its exact identity was committed/buffered.""" + record_identity = (entry.sample_key, entry.features.fingerprints[0]) + if record_identity in self._seen: + return False + if self._finalized: + raise RuntimeError(f"Rank {self.rank} Reasoner cache writer has already been finalized") + + signature = _validate_entries([entry]) + if self._signature is not None and signature != self._signature: + raise ValueError( + f"Rank {self.rank} Reasoner cache entry has signature {signature}, expected {self._signature}" + ) + entry = _entry_on_cpu(entry) + if self._signature is None: + self._signature = signature + entry_bytes = _entry_bytes(entry) + if self._buffer and self._buffered_bytes + entry_bytes > self.max_shard_bytes: + self.flush() + self._buffer.append(entry) + self._buffered_bytes += entry_bytes + self._seen.add(record_identity) + # A single record may exceed the target; commit it immediately so the + # retained feature payload never grows beyond that unavoidable record. + if self._buffered_bytes >= self.max_shard_bytes: + self.flush() + return True + + def flush(self) -> Path | None: + """Atomically commit the buffered entries as one shard + index sidecar.""" + if not self._buffer: + return None + if self._finalized: + raise RuntimeError(f"Rank {self.rank} Reasoner cache writer has already been finalized") + assert self._signature is not None + num_layers, num_kv_heads, head_dim, dtype = self._signature + cache_fingerprint = _manifest_fingerprint( + self.identity, + dtype=dtype, + num_layers=num_layers, + num_kv_heads=num_kv_heads, + head_dim=head_dim, + ) + shard_index = len(self._shards) + shard_name = _incremental_shard_name(self.rank, shard_index) + sidecar_name = _incremental_sidecar_name(self.rank, shard_index) + shard_path = self.rank_dir / shard_name + sidecar_path = self.rank_dir / sidecar_name + + descriptor, temporary_name = tempfile.mkstemp(prefix=f".{shard_name}.tmp-", dir=self.rank_dir) + os.close(descriptor) + temporary_path = Path(temporary_name) + try: + temporary_shard = _write_shard( + temporary_path, + self._buffer, + cache_fingerprint=cache_fingerprint, + ) + with temporary_path.open("rb") as handle: + os.fsync(handle.fileno()) + os.replace(temporary_path, shard_path) + _fsync_directory(self.rank_dir) + except BaseException: + with suppress(FileNotFoundError): + temporary_path.unlink() + raise + + shard = _Shard( + path=shard_name, + sha256=temporary_shard.sha256, + num_tokens=temporary_shard.num_tokens, + records=temporary_shard.records, + ) + sidecar_payload = { + "schema_version": _INCREMENTAL_SCHEMA_VERSION, + "format": _FORMAT_NAME, + "kind": "rank-shard", + "rank": self.rank, + "world_size": self.world_size, + "shard_index": shard_index, + "signature": _signature_to_dict(self._signature), + "cache_fingerprint": cache_fingerprint, + "shard": { + "path": shard.path, + "sha256": shard.sha256, + "num_tokens": shard.num_tokens, + "records": [asdict(record) for record in shard.records], + }, + } + _atomic_write_json(sidecar_path, sidecar_payload) + committed = _IncrementalShard( + rank=self.rank, + world_size=self.world_size, + shard_index=shard_index, + signature=self._signature, + cache_fingerprint=cache_fingerprint, + shard=shard, + shard_path=shard_path, + sidecar_path=sidecar_path, + ) + self._shards.append(committed) + self._buffer.clear() + self._buffered_bytes = 0 + return shard_path + + def finalize(self) -> Path: + """Flush and atomically mark this rank complete; safe to call repeatedly.""" + if self._finalized: + _validate_rank_completion( + self.rank_dir, + identity=self.identity, + rank=self.rank, + world_size=self.world_size, + shards=self._shards, + ) + return self.rank_dir / _INCREMENTAL_RANK_COMPLETE_FILE + self.flush() + payload = _rank_completion_payload( + rank=self.rank, + world_size=self.world_size, + identity=self.identity, + signature=self._signature, + shards=self._shards, + ) + completion_path = self.rank_dir / _INCREMENTAL_RANK_COMPLETE_FILE + _atomic_write_json(completion_path, payload) + self._finalized = True + return completion_path + + +def _manifest_from_incremental_shards( + shards: Sequence[_IncrementalShard], + *, + identity: ReasonerFeatureCacheIdentity, +) -> _Manifest: + if not shards: + raise ValueError("Cannot finalize an empty distributed Reasoner feature cache") + signatures = {item.signature for item in shards} + if len(signatures) != 1: + raise ValueError(f"Distributed Reasoner cache ranks have inconsistent signatures: {sorted(signatures)}") + num_layers, num_kv_heads, head_dim, dtype = next(iter(signatures)) + cache_fingerprint = _manifest_fingerprint( + identity, + dtype=dtype, + num_layers=num_layers, + num_kv_heads=num_kv_heads, + head_dim=head_dim, + ) + seen: dict[tuple[str, str], tuple[int, int]] = {} + for item in shards: + if item.cache_fingerprint != cache_fingerprint: + raise ValueError(f"Rank {item.rank} shard {item.shard_index} has an inconsistent cache fingerprint") + for record in item.shard.records: + record_identity = (record.sample_key, record.fingerprint) + previous = seen.get(record_identity) + if previous is not None: + raise ValueError( + f"Duplicate Reasoner feature record across ranks: {record_identity!r} appears in " + f"rank/shard {previous[0]}/{previous[1]} and {item.rank}/{item.shard_index}" + ) + seen[record_identity] = (item.rank, item.shard_index) + return _Manifest( + cache_fingerprint=cache_fingerprint, + identity=identity, + dtype=dtype, + num_layers=num_layers, + num_kv_heads=num_kv_heads, + head_dim=head_dim, + shards=tuple(item.shard for item in shards), + ) + + +def _manifest_payload(manifest: _Manifest) -> dict[str, Any]: + return { + "schema_version": REASONER_FEATURE_CACHE_SCHEMA_VERSION, + "format": _FORMAT_NAME, + "cache_fingerprint": manifest.cache_fingerprint, + "identity": asdict(manifest.identity), + "dtype": manifest.dtype, + "num_layers": manifest.num_layers, + "num_kv_heads": manifest.num_kv_heads, + "head_dim": manifest.head_dim, + "shards": [ + { + "path": shard.path, + "sha256": shard.sha256, + "num_tokens": shard.num_tokens, + "records": [asdict(record) for record in shard.records], + } + for shard in manifest.shards + ], + } + + +def finalize_incremental_reasoner_feature_cache( + cache_root: str | Path, + *, + identity: ReasonerFeatureCacheIdentity, + world_size: int, + verify_checksums: bool = True, +) -> Path: + """Publish all completed rank-local shards as one immutable provider cache. + + This MVP assumes every rank writes to the same POSIX filesystem. The caller + is responsible for invoking this function on one coordinator after rank + writers have called :meth:`IncrementalReasonerFeatureCacheWriter.finalize`. + Concurrent coordinators are harmless: a single directory rename wins and a + loser accepts the winner only when the complete manifest is identical. + """ + if not isinstance(world_size, int) or isinstance(world_size, bool) or world_size <= 0: + raise ValueError(f"world_size must be positive, got {world_size}") + root = Path(cache_root) + staging_root = _incremental_staging_root(root) + _validate_incremental_staging( + staging_root, + expected_identity=identity, + expected_world_size=world_size, + ) + + incremental_shards: list[_IncrementalShard] = [] + for rank in range(world_size): + rank_dir = staging_root / f"rank-{rank:05d}" + if not rank_dir.is_dir(): + raise FileNotFoundError(f"Incremental Reasoner cache rank directory not found: {rank_dir}") + shards = _scan_incremental_rank( + rank_dir, + identity=identity, + rank=rank, + world_size=world_size, + verify_checksums=verify_checksums, + ) + _validate_rank_completion( + rank_dir, + identity=identity, + rank=rank, + world_size=world_size, + shards=shards, + ) + incremental_shards.extend(shards) + incremental_shards.sort(key=lambda item: (item.rank, item.shard_index)) + expected_manifest = _manifest_from_incremental_shards(incremental_shards, identity=identity) + + if root.exists(): + published = _load_manifest(root) + if published != expected_manifest: + raise FileExistsError(f"Published cache {root} does not match the completed incremental staging data") + return root / REASONER_FEATURE_CACHE_MANIFEST + + root.parent.mkdir(parents=True, exist_ok=True) + temporary_root = Path(tempfile.mkdtemp(prefix=f".{root.name}.publish-", dir=root.parent)) + try: + for item in incremental_shards: + destination = temporary_root / item.shard.path + # Do not hard-link staging into the immutable publication: staging + # is intentionally retained for audit/retry, and a later accidental + # write through that pathname must not mutate the published cache. + shutil.copy2(item.shard_path, destination) + with destination.open("rb") as handle: + os.fsync(handle.fileno()) + if _sha256_file(destination) != item.shard.sha256: + raise ValueError(f"Published shard copy checksum mismatch: {destination}") + _atomic_write_json( + temporary_root / REASONER_FEATURE_CACHE_MANIFEST, + _manifest_payload(expected_manifest), + ) + _fsync_directory(temporary_root) + try: + os.replace(temporary_root, root) + except OSError as error: + if error.errno not in {EEXIST, ENOTEMPTY} or not root.exists(): + raise + published = _load_manifest(root) + if published != expected_manifest: + raise FileExistsError( + f"Concurrent finalizer published a different Reasoner feature cache at {root}" + ) from error + _fsync_directory(root.parent) + except BaseException: + if temporary_root.exists(): + shutil.rmtree(temporary_root, ignore_errors=True) + raise + if temporary_root.exists(): + shutil.rmtree(temporary_root, ignore_errors=True) + return root / REASONER_FEATURE_CACHE_MANIFEST + + +class OfflineReasonerFeatureProvider: + """Read immutable local shards and return completed feature futures.""" + + def __init__( + self, + cache_root: str | Path, + *, + expected_identity: ReasonerFeatureCacheIdentity | None, + expected_dtype: torch.dtype | None = None, + strict_fingerprint: bool = True, + verify_checksums: bool = True, + ) -> None: + self.cache_root = Path(cache_root) + self.manifest = _load_manifest(self.cache_root) + if strict_fingerprint and expected_identity is None: + raise ValueError("strict_fingerprint=True requires an expected cache identity") + if expected_identity is not None and expected_identity != self.manifest.identity: + raise ValueError( + f"Reasoner cache identity mismatch: manifest={self.manifest.identity!r}, expected={expected_identity!r}" + ) + if expected_dtype is not None: + dtype_name = _dtype_name(expected_dtype) + if dtype_name != self.manifest.dtype: + raise ValueError( + f"Reasoner cache dtype mismatch: manifest={self.manifest.dtype!r}, expected={dtype_name!r}" + ) + self.strict_fingerprint = strict_fingerprint + self.verify_checksums = verify_checksums + self._verified_shards: set[str] = set() + self._records: dict[tuple[str, str], tuple[_Shard, _Record]] = { + (record.sample_key, record.fingerprint): (shard, record) + for shard in self.manifest.shards + for record in shard.records + } + self._records_by_key: dict[str, list[tuple[_Shard, _Record]]] = defaultdict(list) + self._records_by_fingerprint: dict[str, list[tuple[_Shard, _Record]]] = defaultdict(list) + for shard in self.manifest.shards: + for record in shard.records: + self._records_by_key[record.sample_key].append((shard, record)) + self._records_by_fingerprint[record.fingerprint].append((shard, record)) + + @property + def cache_fingerprint(self) -> str: + return self.manifest.cache_fingerprint + + def submit(self, requests: Sequence[ReasonerFeatureRequest]) -> Future[ReasonerFeatureBatch]: + future: Future[ReasonerFeatureBatch] = Future() + try: + future.set_result(self.load(requests)) + except BaseException as error: + future.set_exception(error) + return future + + def _verify_shard(self, shard: _Shard) -> Path: + path = self.cache_root / shard.path + if not path.is_file(): + raise FileNotFoundError(f"Reasoner cache shard not found: {path}") + if self.verify_checksums and shard.path not in self._verified_shards: + actual = _sha256_file(path) + if actual != shard.sha256: + raise ValueError( + f"Reasoner cache shard checksum mismatch for {path}: expected {shard.sha256}, got {actual}" + ) + self._verified_shards.add(shard.path) + return path + + def verify_all_shards(self) -> None: + """Stream-validate every published shard without loading feature payloads.""" + signature: _CacheSignature = ( + self.manifest.num_layers, + self.manifest.num_kv_heads, + self.manifest.head_dim, + self.manifest.dtype, + ) + for shard in self.manifest.shards: + path = self._verify_shard(shard) + _validate_incremental_shard_payload( + path, + shard=shard, + signature=signature, + cache_fingerprint=self.manifest.cache_fingerprint, + ) + + def _validate_request(self, request: ReasonerFeatureRequest, record: _Record) -> None: + if request.causal_offsets.numel() != 2: + raise ValueError( + f"Offline cache request {request.sample_key!r} must describe one causal document, " + f"got {request.causal_offsets.numel() - 1}" + ) + if request.token_ids.numel() != record.num_tokens: + raise ValueError( + f"Offline cache token count mismatch for {request.sample_key!r}: " + f"request={request.token_ids.numel()} cache={record.num_tokens}" + ) + if self.strict_fingerprint and request.fingerprint != record.fingerprint: + raise ValueError( + f"Offline cache fingerprint mismatch for {request.sample_key!r}: " + f"request={request.fingerprint!r} cache={record.fingerprint!r}" + ) + + def load(self, requests: Sequence[ReasonerFeatureRequest]) -> ReasonerFeatureBatch: + if not requests: + raise ValueError("At least one Reasoner feature request is required") + + resolved: list[tuple[ReasonerFeatureRequest, _Shard, _Record]] = [] + for request in requests: + if self.strict_fingerprint: + match = self._records.get((request.sample_key, request.fingerprint)) + # Content-addressed fallback deduplicates shared prompts such as + # CFG's null caption across otherwise unrelated sample keys. + if match is None: + fingerprint_matches = self._records_by_fingerprint.get(request.fingerprint, []) + match = fingerprint_matches[0] if fingerprint_matches else None + if match is None: + if request.sample_key in self._records_by_key: + available = sorted( + record.fingerprint for _shard, record in self._records_by_key[request.sample_key] + ) + raise ValueError( + f"Offline cache fingerprint mismatch for {request.sample_key!r}: " + f"request={request.fingerprint!r}, available={available}" + ) + raise KeyError(f"Reasoner feature cache miss for sample {request.sample_key!r}") + shard, record = match + else: + matches = self._records_by_key.get(request.sample_key, []) + if not matches: + raise KeyError(f"Reasoner feature cache miss for sample {request.sample_key!r}") + if len(matches) != 1: + raise ValueError( + f"Non-strict cache lookup for {request.sample_key!r} is ambiguous across {len(matches)} records" + ) + shard, record = matches[0] + self._validate_request(request, record) + resolved.append((request, shard, record)) + + by_shard: dict[str, list[tuple[int, _Record]]] = defaultdict(list) + shard_by_path: dict[str, _Shard] = {} + for output_idx, (_request, shard, record) in enumerate(resolved): + by_shard[shard.path].append((output_idx, record)) + shard_by_path[shard.path] = shard + + layer_k: list[list[torch.Tensor | None]] = [[None] * len(requests) for _ in range(self.manifest.num_layers)] + layer_v: list[list[torch.Tensor | None]] = [[None] * len(requests) for _ in range(self.manifest.num_layers)] + for shard_path, indexed_records in by_shard.items(): + shard = shard_by_path[shard_path] + path = self._verify_shard(shard) + with safe_open(str(path), framework="pt", device="cpu") as handle: + metadata = handle.metadata() or {} + if metadata.get("format") != _FORMAT_NAME: + raise ValueError(f"Reasoner cache shard {path} has invalid format metadata") + if metadata.get("cache_fingerprint") != self.manifest.cache_fingerprint: + raise ValueError(f"Reasoner cache shard {path} belongs to a different manifest") + offsets = handle.get_tensor(_OFFSETS_KEY) + expected_offsets = torch.tensor( + [0, *(record.end for record in shard.records)], + dtype=torch.int64, + ) + if not torch.equal(offsets, expected_offsets): + raise ValueError(f"Reasoner cache shard {path} offsets disagree with the manifest") + for layer_idx in range(self.manifest.num_layers): + k_slice = handle.get_slice(_layer_key("cross_k", layer_idx)) + v_slice = handle.get_slice(_layer_key("cross_v", layer_idx)) + for output_idx, record in indexed_records: + layer_k[layer_idx][output_idx] = k_slice[record.start : record.end].contiguous() + layer_v[layer_idx][output_idx] = v_slice[record.start : record.end].contiguous() + + def _concatenate(parts: list[torch.Tensor | None], *, name: str) -> torch.Tensor: + if any(part is None for part in parts): + raise RuntimeError(f"Internal Reasoner cache read left an unresolved {name} slice") + tensors = [part for part in parts if part is not None] + return torch.cat(tensors, dim=0) + + cross_k = tuple(_concatenate(parts, name=f"K layer {idx}") for idx, parts in enumerate(layer_k)) + cross_v = tuple(_concatenate(parts, name=f"V layer {idx}") for idx, parts in enumerate(layer_v)) + lengths = [record.num_tokens for _request, _shard, record in resolved] + causal_offsets = torch.tensor([0, *torch.tensor(lengths).cumsum(0).tolist()], dtype=torch.int64) + return ReasonerFeatureBatch( + cross_k=cross_k, + cross_v=cross_v, + causal_offsets=causal_offsets, + fingerprints=tuple(request.fingerprint for request in requests), + ) + + +__all__ = [ + "IncrementalReasonerFeatureCacheWriter", + "OfflineReasonerFeatureProvider", + "REASONER_FEATURE_CACHE_MANIFEST", + "REASONER_FEATURE_CACHE_SCHEMA_VERSION", + "ReasonerFeatureCacheEntry", + "ReasonerFeatureCacheIdentity", + "build_reasoner_feature_request_from_text_tokens", + "build_reasoner_feature_requests", + "compute_reasoner_feature_fingerprint", + "finalize_incremental_reasoner_feature_cache", + "write_reasoner_feature_cache", +] diff --git a/cosmos_framework/model/generator/reasoner_feature_cache_test.py b/cosmos_framework/model/generator/reasoner_feature_cache_test.py new file mode 100644 index 000000000..0147bd943 --- /dev/null +++ b/cosmos_framework/model/generator/reasoner_feature_cache_test.py @@ -0,0 +1,596 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +from __future__ import annotations + +import json +from pathlib import Path + +import pytest +import torch + +from cosmos_framework.data.generator.sequence_packing import PackedSequence +from cosmos_framework.model.generator.reasoner_feature_cache import ( + REASONER_FEATURE_CACHE_MANIFEST, + IncrementalReasonerFeatureCacheWriter, + OfflineReasonerFeatureProvider, + ReasonerFeatureCacheEntry, + ReasonerFeatureCacheIdentity, + build_reasoner_feature_request_from_text_tokens, + build_reasoner_feature_requests, + compute_reasoner_feature_fingerprint, + finalize_incremental_reasoner_feature_cache, + write_reasoner_feature_cache, +) +from cosmos_framework.model.generator.reasoner_features import ReasonerFeatureBatch, ReasonerFeatureRequest + + +def _identity(**overrides: str) -> ReasonerFeatureCacheIdentity: + fields = { + "reasoner": "reasoner-checkpoint-sha256", + "tokenizer": "tokenizer-sha256", + "framing": "framing-v1-sha256", + } + fields.update(overrides) + return ReasonerFeatureCacheIdentity(**fields) + + +def _features( + sample_key: str, + length: int, + *, + value_offset: int = 0, + fingerprint: str | None = None, +) -> ReasonerFeatureBatch: + layers_k: list[torch.Tensor] = [] + layers_v: list[torch.Tensor] = [] + for layer_idx in range(2): + values = torch.arange(length * 2 * 4, dtype=torch.float32).reshape(length, 2, 4) + values = (values + value_offset + 100 * layer_idx).to(torch.bfloat16) + layers_k.append(values) + layers_v.append(values + 0.5) + return ReasonerFeatureBatch( + cross_k=tuple(layers_k), + cross_v=tuple(layers_v), + causal_offsets=torch.tensor([0, length], dtype=torch.int64), + fingerprints=(fingerprint or f"fingerprint-{sample_key}",), + ) + + +def _request(sample_key: str, length: int, *, fingerprint: str | None = None) -> ReasonerFeatureRequest: + return ReasonerFeatureRequest( + sample_key=sample_key, + token_ids=torch.arange(length, dtype=torch.int64), + position_ids=torch.arange(length, dtype=torch.int64), + causal_offsets=torch.tensor([0, length], dtype=torch.int64), + fingerprint=fingerprint or f"fingerprint-{sample_key}", + ) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_content_fingerprint_covers_positions_offsets_and_identity() -> None: + token_ids = torch.tensor([10, 20], dtype=torch.int64) + position_ids = torch.tensor([[0.0, 1.0], [0.0, 1.0], [0.0, 1.0]], dtype=torch.float32) + offsets = torch.tensor([0, 2], dtype=torch.int64) + baseline = compute_reasoner_feature_fingerprint( + token_ids, + position_ids, + offsets, + identity=_identity(), + ) + assert baseline == compute_reasoner_feature_fingerprint(token_ids, position_ids, offsets, identity=_identity()) + assert baseline != compute_reasoner_feature_fingerprint( + token_ids, + position_ids + 1, + offsets, + identity=_identity(), + ) + assert baseline != compute_reasoner_feature_fingerprint( + token_ids, + position_ids, + offsets, + identity=_identity(reasoner="other"), + ) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_text_only_producer_framing_matches_finalized_training_pack() -> None: + direct = build_reasoner_feature_request_from_text_tokens( + sample_key="sample", + text_ids=[10, 20], + special_tokens={"eos_token_id": 30, "start_of_generation": 40}, + use_float_positions=True, + identity=_identity(), + ) + assert direct.token_ids.tolist() == [10, 20, 30, 40] + assert direct.position_ids.dtype == torch.float32 + + packed = PackedSequence( + sample_lens=[5], + split_lens=[4, 1], + attn_modes=["causal", "full"], + sequence_length=5, + text_ids=direct.token_ids, + text_indexes=torch.arange(4), + position_ids=torch.cat((direct.position_ids, torch.zeros(3, 1)), dim=1), + ) + from_training = build_reasoner_feature_requests(packed, ("sample",), _identity())[0] + + torch.testing.assert_close(from_training.token_ids, direct.token_ids) + torch.testing.assert_close(from_training.position_ids, direct.position_ids) + torch.testing.assert_close(from_training.causal_offsets, direct.causal_offsets) + assert from_training.fingerprint == direct.fingerprint + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_sharded_cache_round_trip_preserves_request_order(tmp_path: Path) -> None: + entries = [ + ReasonerFeatureCacheEntry("first", _features("first", 3, value_offset=0)), + ReasonerFeatureCacheEntry("second", _features("second", 2, value_offset=1000)), + ] + cache_root = tmp_path / "cache" + manifest_path = write_reasoner_feature_cache( + cache_root, + entries, + identity=_identity(), + # The first entry is 192 bytes and the second is 128 bytes, forcing two shards. + max_shard_bytes=200, + ) + manifest = json.loads(manifest_path.read_text(encoding="utf-8")) + assert len(manifest["shards"]) == 2 + assert all(len(shard["records"]) == 1 for shard in manifest["shards"]) + + provider = OfflineReasonerFeatureProvider(cache_root, expected_identity=_identity()) + future = provider.submit([_request("second", 2), _request("first", 3)]) + assert future.done() + loaded = future.result() + + assert loaded.causal_offsets.tolist() == [0, 2, 5] + assert loaded.fingerprints == ("fingerprint-second", "fingerprint-first") + for layer_idx in range(2): + torch.testing.assert_close( + loaded.cross_k[layer_idx], + torch.cat((entries[1].features.cross_k[layer_idx], entries[0].features.cross_k[layer_idx])), + ) + torch.testing.assert_close( + loaded.cross_v[layer_idx], + torch.cat((entries[1].features.cross_v[layer_idx], entries[0].features.cross_v[layer_idx])), + ) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_writer_rejects_duplicate_keys_and_incompatible_shapes(tmp_path: Path) -> None: + duplicate_entries = [ + ReasonerFeatureCacheEntry("same", _features("same", 2)), + ReasonerFeatureCacheEntry("same", _features("same", 3)), + ] + with pytest.raises(ValueError, match="Duplicate"): + write_reasoner_feature_cache(tmp_path / "duplicate", duplicate_entries, identity=_identity()) + + incompatible = _features("bad", 2) + incompatible = ReasonerFeatureBatch( + cross_k=(torch.zeros(2, 1, 4), torch.zeros(2, 1, 4)), + cross_v=(torch.zeros(2, 1, 4), torch.zeros(2, 1, 4)), + causal_offsets=incompatible.causal_offsets, + fingerprints=incompatible.fingerprints, + ) + with pytest.raises(ValueError, match="signature"): + write_reasoner_feature_cache( + tmp_path / "shape", + [ + ReasonerFeatureCacheEntry("good", _features("good", 2)), + ReasonerFeatureCacheEntry("bad", incompatible), + ], + identity=_identity(), + ) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_writer_is_immutable_and_requires_single_sample_fingerprint(tmp_path: Path) -> None: + cache_root = tmp_path / "cache" + write_reasoner_feature_cache( + cache_root, + [ReasonerFeatureCacheEntry("one", _features("one", 2))], + identity=_identity(), + ) + with pytest.raises(FileExistsError, match="Refusing to overwrite"): + write_reasoner_feature_cache( + cache_root, + [ReasonerFeatureCacheEntry("two", _features("two", 2))], + identity=_identity(), + ) + + no_fingerprint = ReasonerFeatureBatch( + cross_k=(torch.zeros(2, 1, 4),), + cross_v=(torch.zeros(2, 1, 4),), + causal_offsets=torch.tensor([0, 2]), + ) + with pytest.raises(ValueError, match="fingerprint"): + ReasonerFeatureCacheEntry("missing", no_fingerprint) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_provider_fails_closed_on_identity_and_record_fingerprint(tmp_path: Path) -> None: + cache_root = tmp_path / "cache" + write_reasoner_feature_cache( + cache_root, + [ReasonerFeatureCacheEntry("one", _features("one", 2))], + identity=_identity(), + ) + + with pytest.raises(ValueError, match="requires an expected cache identity"): + OfflineReasonerFeatureProvider(cache_root, expected_identity=None) + with pytest.raises(ValueError, match="identity mismatch"): + OfflineReasonerFeatureProvider(cache_root, expected_identity=_identity(reasoner="wrong")) + with pytest.raises(ValueError, match="dtype mismatch"): + OfflineReasonerFeatureProvider( + cache_root, + expected_identity=_identity(), + expected_dtype=torch.float16, + ) + + provider = OfflineReasonerFeatureProvider( + cache_root, + expected_identity=_identity(), + expected_dtype=torch.bfloat16, + ) + future = provider.submit([_request("one", 2, fingerprint="stale")]) + with pytest.raises(ValueError, match="fingerprint mismatch"): + future.result() + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_provider_can_reuse_content_addressed_null_caption(tmp_path: Path) -> None: + shared = _features("null", 2) + cache_root = tmp_path / "cache" + write_reasoner_feature_cache( + cache_root, + [ReasonerFeatureCacheEntry("__null__", shared)], + identity=_identity(), + ) + provider = OfflineReasonerFeatureProvider(cache_root, expected_identity=_identity()) + + loaded = provider.submit([_request("unrelated-video", 2, fingerprint=shared.fingerprints[0])]).result() + torch.testing.assert_close(loaded.cross_k[0], shared.cross_k[0]) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_provider_supports_repeated_sample_keys_with_distinct_prompt_variants(tmp_path: Path) -> None: + first = _features("same", 2, value_offset=0, fingerprint="prompt-a") + second = _features("same", 3, value_offset=1000, fingerprint="prompt-b") + cache_root = tmp_path / "cache" + write_reasoner_feature_cache( + cache_root, + [ + ReasonerFeatureCacheEntry("same", first), + ReasonerFeatureCacheEntry("same", second), + ], + identity=_identity(), + ) + provider = OfflineReasonerFeatureProvider(cache_root, expected_identity=_identity()) + + loaded = provider.submit( + [ + _request("same", 3, fingerprint="prompt-b"), + _request("same", 2, fingerprint="prompt-a"), + _request("same", 2, fingerprint="prompt-a"), + ] + ).result() + + assert loaded.causal_offsets.tolist() == [0, 3, 5, 7] + torch.testing.assert_close( + loaded.cross_k[0], + torch.cat((second.cross_k[0], first.cross_k[0], first.cross_k[0])), + ) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_provider_rejects_cache_miss_token_count_and_multi_document_request(tmp_path: Path) -> None: + cache_root = tmp_path / "cache" + write_reasoner_feature_cache( + cache_root, + [ReasonerFeatureCacheEntry("one", _features("one", 2))], + identity=_identity(), + ) + provider = OfflineReasonerFeatureProvider(cache_root, expected_identity=_identity()) + + with pytest.raises(KeyError, match="cache miss"): + provider.submit([_request("missing", 2)]).result() + with pytest.raises(ValueError, match="token count mismatch"): + provider.submit([_request("one", 3)]).result() + multi_document = ReasonerFeatureRequest( + sample_key="one", + token_ids=torch.arange(2), + position_ids=torch.arange(2), + causal_offsets=torch.tensor([0, 1, 2]), + fingerprint="fingerprint-one", + ) + with pytest.raises(ValueError, match="one causal document"): + provider.submit([multi_document]).result() + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_provider_rejects_stale_manifest_and_corrupt_shard(tmp_path: Path) -> None: + stale_root = tmp_path / "stale" + write_reasoner_feature_cache( + stale_root, + [ReasonerFeatureCacheEntry("one", _features("one", 2))], + identity=_identity(), + ) + manifest_path = stale_root / REASONER_FEATURE_CACHE_MANIFEST + manifest = json.loads(manifest_path.read_text(encoding="utf-8")) + manifest["identity"]["reasoner"] = "silently-changed" + manifest_path.write_text(json.dumps(manifest), encoding="utf-8") + with pytest.raises(ValueError, match="manifest fingerprint"): + OfflineReasonerFeatureProvider(stale_root, expected_identity=_identity()) + + corrupt_root = tmp_path / "corrupt" + write_reasoner_feature_cache( + corrupt_root, + [ReasonerFeatureCacheEntry("one", _features("one", 2))], + identity=_identity(), + ) + corrupt_manifest = json.loads((corrupt_root / REASONER_FEATURE_CACHE_MANIFEST).read_text(encoding="utf-8")) + shard_path = corrupt_root / corrupt_manifest["shards"][0]["path"] + with shard_path.open("r+b") as handle: + handle.seek(-1, 2) + original = handle.read(1) + handle.seek(-1, 2) + handle.write(bytes([original[0] ^ 0xFF])) + corrupt_provider = OfflineReasonerFeatureProvider(corrupt_root, expected_identity=_identity()) + with pytest.raises(ValueError, match="checksum mismatch"): + corrupt_provider.verify_all_shards() + with pytest.raises(ValueError, match="checksum mismatch"): + corrupt_provider.submit([_request("one", 2)]).result() + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_incremental_writer_resumes_skips_and_publishes_provider_compatible_cache(tmp_path: Path) -> None: + cache_root = tmp_path / "incremental" + first = ReasonerFeatureCacheEntry("first", _features("first", 3, value_offset=0)) + second = ReasonerFeatureCacheEntry("second", _features("second", 2, value_offset=1000)) + + rank0 = IncrementalReasonerFeatureCacheWriter( + cache_root, + identity=_identity(), + rank=0, + world_size=2, + max_shard_bytes=1024, + ) + assert rank0.append(first) + assert rank0.buffered_bytes > 0 + committed_path = rank0.flush() + assert committed_path is not None and committed_path.is_file() + assert rank0.buffered_bytes == 0 + + resumed = IncrementalReasonerFeatureCacheWriter.resume( + cache_root, + identity=_identity(), + rank=0, + world_size=2, + max_shard_bytes=1024, + ) + assert resumed.committed_records == 1 + assert resumed.contains("first", first.features.fingerprints[0]) + assert not resumed.contains("missing", "missing-fingerprint") + assert not resumed.append(first) + completion = resumed.finalize() + assert completion.is_file() + assert resumed.finalize() == completion + + rank1 = IncrementalReasonerFeatureCacheWriter( + cache_root, + identity=_identity(), + rank=1, + world_size=2, + max_shard_bytes=1024, + ) + assert rank1.append(second) + rank1.finalize() + + manifest_path = finalize_incremental_reasoner_feature_cache( + cache_root, + identity=_identity(), + world_size=2, + ) + assert manifest_path.is_file() + assert ( + finalize_incremental_reasoner_feature_cache( + cache_root, + identity=_identity(), + world_size=2, + ) + == manifest_path + ) + + staging_root = cache_root.parent / f".{cache_root.name}.reasoner-kv-staging" + staged_shard = next(staging_root.glob("rank-00000/*.safetensors")) + assert not staged_shard.samefile(cache_root / staged_shard.name) + + provider = OfflineReasonerFeatureProvider( + cache_root, + expected_identity=_identity(), + expected_dtype=torch.bfloat16, + ) + loaded = provider.submit([_request("second", 2), _request("first", 3)]).result() + assert loaded.causal_offsets.tolist() == [0, 2, 5] + torch.testing.assert_close( + loaded.cross_k[0], + torch.cat((second.features.cross_k[0], first.features.cross_k[0])), + ) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_incremental_writer_bounds_payload_and_commits_oversized_records(tmp_path: Path) -> None: + cache_root = tmp_path / "bounded" + writer = IncrementalReasonerFeatureCacheWriter( + cache_root, + identity=_identity(), + rank=0, + world_size=1, + # A length-1 entry is 64 bytes; two entries must become separate shards. + max_shard_bytes=100, + ) + assert writer.append(ReasonerFeatureCacheEntry("one", _features("one", 1))) + assert writer.buffered_bytes == 64 + assert writer.append(ReasonerFeatureCacheEntry("two", _features("two", 1))) + assert writer.buffered_bytes == 64 + writer.finalize() + + staging_root = cache_root.parent / f".{cache_root.name}.reasoner-kv-staging" + assert len(list(staging_root.glob("rank-00000/*.safetensors"))) == 2 + + oversized_root = tmp_path / "oversized" + oversized = IncrementalReasonerFeatureCacheWriter( + oversized_root, + identity=_identity(), + rank=0, + world_size=1, + max_shard_bytes=32, + ) + assert oversized.append(ReasonerFeatureCacheEntry("large", _features("large", 2))) + assert oversized.buffered_bytes == 0 + assert oversized.committed_records == 1 + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_incremental_finalizer_rejects_missing_completion_duplicates_and_signature_mismatch(tmp_path: Path) -> None: + incomplete_root = tmp_path / "incomplete" + incomplete = IncrementalReasonerFeatureCacheWriter( + incomplete_root, + identity=_identity(), + rank=0, + world_size=1, + ) + incomplete.append(ReasonerFeatureCacheEntry("one", _features("one", 2))) + incomplete.flush() + with pytest.raises(FileNotFoundError, match="completion marker"): + finalize_incremental_reasoner_feature_cache( + incomplete_root, + identity=_identity(), + world_size=1, + ) + + duplicate_root = tmp_path / "duplicate-ranks" + duplicate_entry = ReasonerFeatureCacheEntry("same", _features("same", 2)) + for rank in range(2): + writer = IncrementalReasonerFeatureCacheWriter( + duplicate_root, + identity=_identity(), + rank=rank, + world_size=2, + ) + writer.append(duplicate_entry) + writer.finalize() + with pytest.raises(ValueError, match="Duplicate Reasoner feature record across ranks"): + finalize_incremental_reasoner_feature_cache( + duplicate_root, + identity=_identity(), + world_size=2, + ) + + signature_root = tmp_path / "signature-ranks" + rank0 = IncrementalReasonerFeatureCacheWriter( + signature_root, + identity=_identity(), + rank=0, + world_size=2, + ) + rank0.append(ReasonerFeatureCacheEntry("normal", _features("normal", 2))) + rank0.finalize() + incompatible = ReasonerFeatureBatch( + cross_k=(torch.zeros(2, 1, 4, dtype=torch.bfloat16),) * 2, + cross_v=(torch.zeros(2, 1, 4, dtype=torch.bfloat16),) * 2, + causal_offsets=torch.tensor([0, 2]), + fingerprints=("different-signature",), + ) + rank1 = IncrementalReasonerFeatureCacheWriter( + signature_root, + identity=_identity(), + rank=1, + world_size=2, + ) + rank1.append(ReasonerFeatureCacheEntry("incompatible", incompatible)) + rank1.finalize() + with pytest.raises(ValueError, match="inconsistent signatures"): + finalize_incremental_reasoner_feature_cache( + signature_root, + identity=_identity(), + world_size=2, + ) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_incremental_resume_and_finalizer_reject_corrupt_committed_shard(tmp_path: Path) -> None: + cache_root = tmp_path / "corrupt-incremental" + writer = IncrementalReasonerFeatureCacheWriter( + cache_root, + identity=_identity(), + rank=0, + world_size=1, + ) + writer.append(ReasonerFeatureCacheEntry("one", _features("one", 2))) + writer.finalize() + + staging_root = cache_root.parent / f".{cache_root.name}.reasoner-kv-staging" + shard_path = next(staging_root.glob("rank-00000/*.safetensors")) + with shard_path.open("r+b") as handle: + handle.seek(-1, 2) + original = handle.read(1) + handle.seek(-1, 2) + handle.write(bytes([original[0] ^ 0xFF])) + + with pytest.raises(ValueError, match="checksum mismatch"): + IncrementalReasonerFeatureCacheWriter.resume( + cache_root, + identity=_identity(), + rank=0, + world_size=1, + ) + with pytest.raises(ValueError, match="checksum mismatch"): + finalize_incremental_reasoner_feature_cache( + cache_root, + identity=_identity(), + world_size=1, + ) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_incremental_finalizer_allows_an_empty_completed_rank(tmp_path: Path) -> None: + cache_root = tmp_path / "empty-rank" + rank0 = IncrementalReasonerFeatureCacheWriter( + cache_root, + identity=_identity(), + rank=0, + world_size=2, + ) + rank0.append(ReasonerFeatureCacheEntry("one", _features("one", 2))) + rank0.finalize() + IncrementalReasonerFeatureCacheWriter( + cache_root, + identity=_identity(), + rank=1, + world_size=2, + ).finalize() + + manifest_path = finalize_incremental_reasoner_feature_cache( + cache_root, + identity=_identity(), + world_size=2, + ) + assert manifest_path.is_file() diff --git a/cosmos_framework/model/generator/reasoner_features.py b/cosmos_framework/model/generator/reasoner_features.py new file mode 100644 index 000000000..a41cbaf3c --- /dev/null +++ b/cosmos_framework/model/generator/reasoner_features.py @@ -0,0 +1,870 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +"""Provider-neutral Reasoner K/V features for generator-only execution. + +The canonical boundary is the exact key/value pair consumed by the generator's +cross-attention, after all Reasoner-side normalization and RoPE. Canonical +tensors are unpadded and use ``[sequence, kv_heads, head_dim]`` layout. This +module deliberately lives outside :mod:`cosmos_framework.inference` so inline, +offline, and remote providers can share one training-safe contract. + +Only ordinary two-way attention without context parallelism is supported here. +Unsupported attention layouts fail closed instead of silently applying a less +restrictive mask. +""" + +from __future__ import annotations + +from concurrent.futures import Future +from dataclasses import dataclass +from typing import Mapping, Protocol, Sequence + +import torch + +from cosmos_framework.data.generator.sequence_packing.runtime import ( + SequencePack, + drop_pad_segment, + from_und_gen_splits, + get_gen_seq, +) +from cosmos_framework.model.attention import attention +from cosmos_framework.model.generator.mot.attention import SplitInfo, dispatch_attention +from cosmos_framework.model.generator.mot.unified_mot import ReasonerKVCache +from cosmos_framework.model.generator.utils.memory import KVToStore, MemoryState, MemoryValue + +_CACHE_CROSS_K = "cross_k" +_CACHE_CROSS_V = "cross_v" +_CACHE_CAUSAL_OFFSETS = "causal_offsets" + + +def _detached_tensor(value: torch.Tensor, *, name: str) -> torch.Tensor: + if not isinstance(value, torch.Tensor): + raise TypeError(f"{name} must be a torch.Tensor, got {type(value).__name__}") + if value.device.type == "meta": + raise ValueError(f"{name} must contain materialized values, not a meta tensor") + return value.detach() + + +def _validate_offsets( + offsets: torch.Tensor, + *, + name: str, + expected_total: int | None = None, +) -> torch.Tensor: + offsets = _detached_tensor(offsets, name=name) + if offsets.ndim != 1: + raise ValueError(f"{name} must have shape [segments + 1], got {tuple(offsets.shape)}") + if offsets.numel() < 2: + raise ValueError(f"{name} must contain at least [0, end]") + if offsets.dtype not in (torch.int32, torch.int64): + raise TypeError(f"{name} must use int32 or int64, got {offsets.dtype}") + + # Feature/state construction is deliberately outside torch.compile. A host + # copy gives useful validation errors and avoids data-dependent guards in the + # compiled decoder layers. + host_offsets = offsets.to(device="cpu", dtype=torch.int64) + if int(host_offsets[0]) != 0: + raise ValueError(f"{name} must start at 0, got {int(host_offsets[0])}") + if bool(torch.any(host_offsets[1:] < host_offsets[:-1])): + raise ValueError(f"{name} must be monotonically non-decreasing") + total = int(host_offsets[-1]) + if expected_total is not None and total != expected_total: + raise ValueError(f"{name} ends at {total}, expected {expected_total}") + if total > torch.iinfo(torch.int32).max: + raise ValueError(f"{name} exceeds the int32 attention-metadata limit") + return offsets + + +@dataclass(frozen=True) +class ReasonerLayerKV: + """Canonical K/V for one decoder layer. + + ``cross_k`` and ``cross_v`` are unpadded tensors with shape + ``[S_und, num_kv_heads, head_dim]``. Construction detaches both tensors; + gradients must stop at the frozen Reasoner/generator boundary. + """ + + cross_k: torch.Tensor + cross_v: torch.Tensor + + def __post_init__(self) -> None: + cross_k = _detached_tensor(self.cross_k, name="cross_k") + cross_v = _detached_tensor(self.cross_v, name="cross_v") + if cross_k.ndim != 3 or cross_v.ndim != 3: + raise ValueError( + "Reasoner layer K/V must have shape [sequence, kv_heads, head_dim], " + f"got K={tuple(cross_k.shape)} V={tuple(cross_v.shape)}" + ) + if cross_k.shape != cross_v.shape: + raise ValueError( + f"Reasoner layer K/V shapes must match, got K={tuple(cross_k.shape)} V={tuple(cross_v.shape)}" + ) + if any(size <= 0 for size in cross_k.shape): + raise ValueError(f"Reasoner layer K/V dimensions must be positive, got {tuple(cross_k.shape)}") + if cross_k.dtype != cross_v.dtype: + raise TypeError(f"Reasoner layer K/V dtypes must match, got K={cross_k.dtype} V={cross_v.dtype}") + if cross_k.device != cross_v.device: + raise ValueError(f"Reasoner layer K/V devices must match, got K={cross_k.device} V={cross_v.device}") + if not cross_k.is_floating_point(): + raise TypeError(f"Reasoner layer K/V must be floating point, got {cross_k.dtype}") + object.__setattr__(self, "cross_k", cross_k) + object.__setattr__(self, "cross_v", cross_v) + + @property + def sequence_length(self) -> int: + return self.cross_k.shape[0] + + @property + def num_kv_heads(self) -> int: + return self.cross_k.shape[1] + + @property + def head_dim(self) -> int: + return self.cross_k.shape[2] + + def to( + self, + device: torch.device | str, + *, + dtype: torch.dtype | None = None, + non_blocking: bool = False, + ) -> ReasonerLayerKV: + """Return this layer on ``device`` while preserving the detached boundary.""" + return ReasonerLayerKV( + self.cross_k.to(device=device, dtype=dtype, non_blocking=non_blocking), + self.cross_v.to(device=device, dtype=dtype, non_blocking=non_blocking), + ) + + +@dataclass(frozen=True) +class ReasonerFeatureRequest: + """Storage/provider-independent identity and inputs for one feature request.""" + + sample_key: str + token_ids: torch.Tensor + position_ids: torch.Tensor + causal_offsets: torch.Tensor + fingerprint: str + + def __post_init__(self) -> None: + if not self.sample_key: + raise ValueError("sample_key must be non-empty") + if not self.fingerprint: + raise ValueError("fingerprint must be non-empty") + token_ids = _detached_tensor(self.token_ids, name="token_ids") + position_ids = _detached_tensor(self.position_ids, name="position_ids") + if token_ids.ndim != 1: + raise ValueError(f"token_ids must have shape [sequence], got {tuple(token_ids.shape)}") + if token_ids.dtype not in (torch.int32, torch.int64): + raise TypeError(f"token_ids must use int32 or int64, got {token_ids.dtype}") + if position_ids.ndim not in (1, 2) or position_ids.shape[-1] != token_ids.numel(): + raise ValueError( + "position_ids must have shape [sequence] or [axes, sequence] matching token_ids, " + f"got {tuple(position_ids.shape)} for {token_ids.numel()} tokens" + ) + if position_ids.dtype not in (torch.int32, torch.int64, torch.float32, torch.float64): + raise TypeError(f"position_ids must use int32, int64, float32, or float64, got {position_ids.dtype}") + causal_offsets = _validate_offsets( + self.causal_offsets, + name="causal_offsets", + expected_total=token_ids.numel(), + ) + object.__setattr__(self, "token_ids", token_ids) + object.__setattr__(self, "position_ids", position_ids) + object.__setattr__(self, "causal_offsets", causal_offsets) + + +@dataclass(frozen=True) +class ReasonerFeatureBatch: + """Canonical unpadded Reasoner K/V for a packed batch. + + ``cross_k[layer]`` and ``cross_v[layer]`` each use the canonical + ``[S_und, H_kv, D]`` layout. ``causal_offsets`` partitions that shared + sequence stream into samples. ``fingerprints`` may be empty for ephemeral + inline capture; provider-produced batches carry one fingerprint per sample. + """ + + cross_k: tuple[torch.Tensor, ...] + cross_v: tuple[torch.Tensor, ...] + causal_offsets: torch.Tensor + fingerprints: tuple[str, ...] = () + + def __post_init__(self) -> None: + cross_k = tuple(self.cross_k) + cross_v = tuple(self.cross_v) + fingerprints = tuple(self.fingerprints) + if not cross_k: + raise ValueError("ReasonerFeatureBatch must contain at least one layer") + if len(cross_k) != len(cross_v): + raise ValueError(f"ReasonerFeatureBatch layer counts disagree: K={len(cross_k)} V={len(cross_v)}") + + layers = tuple(ReasonerLayerKV(k, v) for k, v in zip(cross_k, cross_v)) + reference = layers[0] + for layer_idx, layer in enumerate(layers[1:], start=1): + if layer.cross_k.shape != reference.cross_k.shape: + raise ValueError( + "All Reasoner layers must share [sequence, kv_heads, head_dim], " + f"layer 0={tuple(reference.cross_k.shape)} layer {layer_idx}={tuple(layer.cross_k.shape)}" + ) + if layer.cross_k.dtype != reference.cross_k.dtype: + raise TypeError( + f"All Reasoner layers must share a dtype, layer 0={reference.cross_k.dtype} " + f"layer {layer_idx}={layer.cross_k.dtype}" + ) + if layer.cross_k.device != reference.cross_k.device: + raise ValueError( + f"All Reasoner layers must share a device, layer 0={reference.cross_k.device} " + f"layer {layer_idx}={layer.cross_k.device}" + ) + + causal_offsets = _validate_offsets( + self.causal_offsets, + name="causal_offsets", + expected_total=reference.sequence_length, + ) + num_samples = causal_offsets.numel() - 1 + if fingerprints and len(fingerprints) != num_samples: + raise ValueError( + f"fingerprints must be empty or contain one entry per sample, got {len(fingerprints)} for " + f"{num_samples} samples" + ) + if any(not fingerprint for fingerprint in fingerprints): + raise ValueError("fingerprints must not contain empty entries") + + object.__setattr__(self, "cross_k", tuple(layer.cross_k for layer in layers)) + object.__setattr__(self, "cross_v", tuple(layer.cross_v for layer in layers)) + object.__setattr__(self, "causal_offsets", causal_offsets) + object.__setattr__(self, "fingerprints", fingerprints) + + @property + def num_layers(self) -> int: + return len(self.cross_k) + + @property + def num_samples(self) -> int: + return self.causal_offsets.numel() - 1 + + @property + def sequence_length(self) -> int: + return self.cross_k[0].shape[0] + + def layer(self, layer_idx: int) -> ReasonerLayerKV: + """Return one validated canonical layer.""" + return ReasonerLayerKV(self.cross_k[layer_idx], self.cross_v[layer_idx]) + + @classmethod + def from_stacked( + cls, + cross_k: torch.Tensor, + cross_v: torch.Tensor, + causal_offsets: torch.Tensor, + fingerprints: Sequence[str] = (), + ) -> ReasonerFeatureBatch: + """Build from cache-friendly ``[layers, sequence, heads, dim]`` tensors.""" + cross_k = _detached_tensor(cross_k, name="cross_k") + cross_v = _detached_tensor(cross_v, name="cross_v") + if cross_k.ndim != 4 or cross_v.ndim != 4: + raise ValueError( + "Stacked Reasoner K/V must have shape [layers, sequence, kv_heads, head_dim], " + f"got K={tuple(cross_k.shape)} V={tuple(cross_v.shape)}" + ) + if cross_k.shape != cross_v.shape: + raise ValueError(f"Stacked Reasoner K/V shapes must match, got K={cross_k.shape} V={cross_v.shape}") + return cls(tuple(cross_k.unbind(0)), tuple(cross_v.unbind(0)), causal_offsets, tuple(fingerprints)) + + def to_stacked(self) -> tuple[torch.Tensor, torch.Tensor]: + """Return cache-friendly ``[layers, sequence, heads, dim]`` K/V tensors.""" + return torch.stack(self.cross_k, dim=0), torch.stack(self.cross_v, dim=0) + + @classmethod + def from_cache_tensors( + cls, + tensors: Mapping[str, torch.Tensor], + *, + fingerprints: Sequence[str] = (), + ) -> ReasonerFeatureBatch: + """Build from the tensor payload used by safetensors/cache providers.""" + missing = {_CACHE_CROSS_K, _CACHE_CROSS_V, _CACHE_CAUSAL_OFFSETS} - tensors.keys() + if missing: + raise KeyError(f"Reasoner feature cache payload is missing: {sorted(missing)}") + return cls.from_stacked( + tensors[_CACHE_CROSS_K], + tensors[_CACHE_CROSS_V], + tensors[_CACHE_CAUSAL_OFFSETS], + fingerprints, + ) + + def to_cache_tensors(self) -> dict[str, torch.Tensor]: + """Return a storage-friendly tensor mapping; fingerprints belong in the manifest.""" + cross_k, cross_v = self.to_stacked() + return { + _CACHE_CROSS_K: cross_k, + _CACHE_CROSS_V: cross_v, + _CACHE_CAUSAL_OFFSETS: self.causal_offsets, + } + + def to( + self, + device: torch.device | str, + *, + dtype: torch.dtype | None = None, + non_blocking: bool = False, + ) -> ReasonerFeatureBatch: + """Stage all layers and offsets on a device.""" + return ReasonerFeatureBatch( + tuple(k.to(device=device, dtype=dtype, non_blocking=non_blocking) for k in self.cross_k), + tuple(v.to(device=device, dtype=dtype, non_blocking=non_blocking) for v in self.cross_v), + self.causal_offsets.to(device=device, non_blocking=non_blocking), + self.fingerprints, + ) + + +class ReasonerFeatureProvider(Protocol): + """Asynchronous provider contract shared by inline/offline/remote backends.""" + + def submit(self, requests: Sequence[ReasonerFeatureRequest]) -> Future[ReasonerFeatureBatch]: ... + + +@torch.inference_mode() +def extract_reasoner_feature_batch( + causal_lm: torch.nn.Module, + requests: Sequence[ReasonerFeatureRequest], +) -> ReasonerFeatureBatch: + """Run the UND-only prefill and return its exact generator-facing K/V. + + Requests are evaluated independently to preserve causal-document isolation. + The initial implementation intentionally rejects architectures that apply a + second generator-specific UND K normalization; Nano's Qwen dense model uses + the same normalized/RoPE K for Reasoner self-attention and GEN cross-attention. + """ + if not requests: + raise ValueError("At least one Reasoner feature request is required") + multi_document_requests = [request.sample_key for request in requests if request.causal_offsets.numel() != 2] + if multi_document_requests: + raise NotImplementedError( + "Reasoner feature extraction currently requires one causal document per request; " + f"multi-document requests={multi_document_requests}" + ) + try: + model = causal_lm.model + layers = model.layers + embedding = model.embed_tokens + except AttributeError as error: + raise TypeError("Expected a *TextForCausalLM wrapper with model.embed_tokens and model.layers") from error + if not getattr(model, "include_und_pathway", True): + raise ValueError("Reasoner feature extraction requires include_und_pathway=True") + if causal_lm.training: + raise ValueError("Reasoner feature extraction requires causal_lm.eval()") + unsupported_layers = [ + layer_idx + for layer_idx, layer in enumerate(layers) + if getattr(layer.self_attn, "k_norm_und_for_gen", None) is not None + ] + if unsupported_layers: + raise NotImplementedError( + "UND-only extraction does not yet support generator-specific K normalization; " + f"affected layers={unsupported_layers}" + ) + + device = embedding.weight.device + per_request_keys: list[tuple[torch.Tensor, ...]] = [] + per_request_values: list[tuple[torch.Tensor, ...]] = [] + offsets = [0] + for request in requests: + token_ids = request.token_ids.to(device=device, dtype=torch.long).unsqueeze(0) + positions = request.position_ids.to(device=device) + position_ids = positions.unsqueeze(0) if positions.ndim == 1 else positions.unsqueeze(1) + cache = ReasonerKVCache.empty(num_layers=len(layers)) + model.reasoner_forward(input_ids=token_ids, position_ids=position_ids, cache=cache) + if any(key is None for key in cache.keys) or any(value is None for value in cache.values): + raise RuntimeError("Reasoner prefill did not populate every K/V cache layer") + keys = tuple(key.squeeze(0).detach() for key in cache.keys if key is not None) + values = tuple(value.squeeze(0).detach() for value in cache.values if value is not None) + if any(key.shape[0] != request.token_ids.numel() for key in keys): + raise RuntimeError(f"Reasoner K/V length mismatch for request {request.sample_key!r}") + per_request_keys.append(keys) + per_request_values.append(values) + offsets.append(offsets[-1] + request.token_ids.numel()) + + cross_k = tuple( + torch.cat([request_keys[layer_idx] for request_keys in per_request_keys], dim=0) + for layer_idx in range(len(layers)) + ) + cross_v = tuple( + torch.cat([request_values[layer_idx] for request_values in per_request_values], dim=0) + for layer_idx in range(len(layers)) + ) + return ReasonerFeatureBatch( + cross_k=cross_k, + cross_v=cross_v, + causal_offsets=torch.tensor(offsets, dtype=torch.int64, device=device), + fingerprints=tuple(request.fingerprint for request in requests), + ) + + +@dataclass(frozen=True) +class ReasonerAttentionMetadata: + """Precomputed packed-attention layout shared by every decoder layer.""" + + gen_len: int + kv_reorder_indices: torch.Tensor | None = None + cumulative_seqlen_q: torch.Tensor | None = None + cumulative_seqlen_kv: torch.Tensor | None = None + max_seqlen_q: int = 0 + max_seqlen_kv: int = 0 + + +def _pack_offsets(pack: SequencePack, key: str, *, expected_total: int) -> torch.Tensor: + offsets = pack.get(key) + if not isinstance(offsets, torch.Tensor): + raise ValueError(f"SequencePack must provide tensor metadata {key!r}") + offsets = drop_pad_segment(pack, offsets) + return _validate_offsets(offsets, name=f"SequencePack[{key!r}]", expected_total=expected_total) + + +def build_reasoner_attention_metadata( + feature_batch: ReasonerFeatureBatch, + pack: SequencePack, + device: torch.device | str, +) -> ReasonerAttentionMetadata: + """Validate a SequencePack and construct isolated GEN-to-(UND+GEN) ranges. + + Canonical UND and live GEN tensors are each sample-major streams. For a + multi-sample pack, ``kv_reorder_indices`` interleaves those two streams as + ``[und_0, gen_0, und_1, gen_1, ...]`` and varlen offsets prevent leakage + between samples. A single sample uses the equivalent dense path. + """ + target_device = torch.device(device) + try: + gen_len = int(pack["_num_full_tokens"]) + und_len = int(pack["_num_causal_tokens"]) + except KeyError as error: + raise ValueError(f"SequencePack is missing real-token metadata {error.args[0]!r}") from error + if gen_len <= 0: + raise ValueError(f"SequencePack must contain at least one real GEN token, got {gen_len}") + if und_len != feature_batch.sequence_length: + raise ValueError( + f"SequencePack UND token count {und_len} does not match Reasoner features {feature_batch.sequence_length}" + ) + + runtime_und_offsets = _pack_offsets(pack, "_causal_seq_offsets", expected_total=und_len) + runtime_gen_offsets = _pack_offsets(pack, "_full_only_seq_offsets", expected_total=gen_len) + feature_offsets_host = feature_batch.causal_offsets.to(device="cpu", dtype=torch.int64) + runtime_und_offsets_host = runtime_und_offsets.to(device="cpu", dtype=torch.int64) + if not torch.equal(feature_offsets_host, runtime_und_offsets_host): + raise ValueError( + "Reasoner feature causal offsets do not match SequencePack offsets: " + f"features={feature_offsets_host.tolist()} pack={runtime_und_offsets_host.tolist()}" + ) + if runtime_und_offsets.numel() != runtime_gen_offsets.numel(): + raise ValueError( + "Reasoner/GEN sample counts disagree: " + f"reasoner={runtime_und_offsets.numel() - 1} gen={runtime_gen_offsets.numel() - 1}" + ) + + und_offsets = feature_batch.causal_offsets.to(device=target_device, dtype=torch.int32) + gen_offsets = runtime_gen_offsets.to(device=target_device, dtype=torch.int32) + if und_offsets.numel() == 2: + return ReasonerAttentionMetadata(gen_len=gen_len) + + und_lens = und_offsets[1:] - und_offsets[:-1] + gen_lens = gen_offsets[1:] - gen_offsets[:-1] + if bool(torch.any(und_lens <= 0)) or bool(torch.any(gen_lens <= 0)): + raise ValueError("Every sample must contain at least one Reasoner token and one GEN token") + sample_ids = torch.arange(und_lens.numel(), dtype=torch.int64, device=target_device) + und_sample_ids = torch.repeat_interleave(sample_ids, und_lens.to(dtype=torch.int64)) + gen_sample_ids = torch.repeat_interleave(sample_ids, gen_lens.to(dtype=torch.int64)) + reorder = torch.argsort(torch.cat((und_sample_ids, gen_sample_ids)), stable=True) + return ReasonerAttentionMetadata( + gen_len=gen_len, + kv_reorder_indices=reorder, + cumulative_seqlen_q=gen_offsets, + cumulative_seqlen_kv=und_offsets + gen_offsets, + max_seqlen_q=int(gen_lens.max()), + max_seqlen_kv=int((und_lens + gen_lens).max()), + ) + + +@dataclass +class StaticReasonerKVMemoryValue(MemoryValue): + """One layer's external Reasoner K/V plus immutable attention metadata.""" + + cross_k: torch.Tensor + cross_v: torch.Tensor + gen_len: int + kv_reorder_indices: torch.Tensor | None = None + cumulative_seqlen_q: torch.Tensor | None = None + cumulative_seqlen_kv: torch.Tensor | None = None + max_seqlen_q: int = 0 + max_seqlen_kv: int = 0 + frame_idx: int = 1 + for_cuda_graphs: bool = False + + @property + def supports_context_parallel_attention(self) -> bool: + return False + + +@dataclass +class ReasonerKVCaptureMemoryValue(MemoryValue): + """Marker requesting a normal joint pass whose UND K/V is captured.""" + + frame_idx: int = 0 + for_cuda_graphs: bool = False + + @property + def supports_context_parallel_attention(self) -> bool: + return False + + +class StaticReasonerKVMemoryState(MemoryState): + """Read-only memory state that starts in gen-only mode from prefilled K/V.""" + + def __init__(self, feature_batch: ReasonerFeatureBatch) -> None: + self.feature_batch = feature_batch + self._active_batch = feature_batch + self._metadata: ReasonerAttentionMetadata | None = None + + def init(self, hidden_states: dict, device: torch.device) -> None: + self._active_batch = self.feature_batch.to(device, non_blocking=True) + self._metadata = build_reasoner_attention_metadata(self._active_batch, hidden_states, device) + + def read_for_layer(self, layer_idx: int) -> StaticReasonerKVMemoryValue: + if self._metadata is None: + raise RuntimeError("StaticReasonerKVMemoryState.init() must be called before read_for_layer()") + layer = self._active_batch.layer(layer_idx) + metadata = self._metadata + return StaticReasonerKVMemoryValue( + cross_k=layer.cross_k, + cross_v=layer.cross_v, + gen_len=metadata.gen_len, + kv_reorder_indices=metadata.kv_reorder_indices, + cumulative_seqlen_q=metadata.cumulative_seqlen_q, + cumulative_seqlen_kv=metadata.cumulative_seqlen_kv, + max_seqlen_q=metadata.max_seqlen_q, + max_seqlen_kv=metadata.max_seqlen_kv, + ) + + def write_for_layer(self, layer_idx: int, kv_to_store: KVToStore) -> None: + # PackedAttentionMoT currently emits an empty-UND write even in gen-only + # mode. External features are immutable, so deliberately ignore it. + del kv_to_store + if not 0 <= layer_idx < self._active_batch.num_layers: + raise IndexError(f"Reasoner layer index {layer_idx} is out of range") + + def is_gen_only(self) -> bool: + return True + + def requires_natten_metadata(self) -> bool: + return False + + +class CapturingReasonerKVMemoryState(MemoryState): + """Capture inline Reasoner K/V once, then serve it through the static path. + + The first forward is a normal joint UND+GEN pass. ``write_for_layer`` + detaches the exact RoPE-applied K/V emitted by each layer. Once every layer + is populated, the next forward reports ``is_gen_only()`` and reuses the + captured canonical batch. + """ + + def __init__(self, num_layers: int, *, fingerprints: Sequence[str] = ()) -> None: + if num_layers <= 0: + raise ValueError(f"num_layers must be positive, got {num_layers}") + self._layers: list[ReasonerLayerKV | None] = [None] * num_layers + self._fingerprints = tuple(fingerprints) + self._causal_offsets: torch.Tensor | None = None + self._static_state: StaticReasonerKVMemoryState | None = None + + def init(self, hidden_states: dict, device: torch.device) -> None: + if self.is_gen_only(): + if self._static_state is None: + self._static_state = StaticReasonerKVMemoryState(self.to_feature_batch()) + self._static_state.init(hidden_states, device) + return + + try: + und_len = int(hidden_states["_num_causal_tokens"]) + except KeyError as error: + raise ValueError("SequencePack is missing real-token metadata '_num_causal_tokens'") from error + offsets = _pack_offsets(hidden_states, "_causal_seq_offsets", expected_total=und_len).detach() + if self._causal_offsets is None: + if self._fingerprints and len(self._fingerprints) != offsets.numel() - 1: + raise ValueError( + f"fingerprints contain {len(self._fingerprints)} entries for {offsets.numel() - 1} samples" + ) + self._causal_offsets = offsets + elif not torch.equal( + self._causal_offsets.to(device="cpu", dtype=torch.int64), + offsets.to(device="cpu", dtype=torch.int64), + ): + raise ValueError("SequencePack causal offsets changed during partial Reasoner K/V capture") + + def read_for_layer(self, layer_idx: int) -> MemoryValue: + if not 0 <= layer_idx < len(self._layers): + raise IndexError(f"Reasoner layer index {layer_idx} is out of range") + if self.is_gen_only(): + if self._static_state is None: + raise RuntimeError("CapturingReasonerKVMemoryState.init() must run before cached reads") + return self._static_state.read_for_layer(layer_idx) + return ReasonerKVCaptureMemoryValue() + + def write_for_layer(self, layer_idx: int, kv_to_store: KVToStore) -> None: + if not 0 <= layer_idx < len(self._layers): + raise IndexError(f"Reasoner layer index {layer_idx} is out of range") + if self.is_gen_only(): + # See StaticReasonerKVMemoryState.write_for_layer. + return + if self._causal_offsets is None: + raise RuntimeError("CapturingReasonerKVMemoryState.init() must be called before writes") + _gen_k, _gen_v, und_k, und_v = kv_to_store + expected_total = int(self._causal_offsets[-1].to(device="cpu")) + self._layers[layer_idx] = ReasonerLayerKV( + self._flatten_captured(und_k, expected_total, name="und_k").clone(), + self._flatten_captured(und_v, expected_total, name="und_v").clone(), + ) + + def _flatten_captured(self, tensor: torch.Tensor, expected_total: int, *, name: str) -> torch.Tensor: + tensor = _detached_tensor(tensor, name=name) + if tensor.ndim == 3: + if tensor.shape[0] < expected_total: + raise ValueError(f"Captured {name} has {tensor.shape[0]} tokens, expected {expected_total}") + return tensor[:expected_total] + if tensor.ndim != 4: + raise ValueError(f"Captured {name} must have shape [S,H,D] or [B,S,H,D], got {tensor.shape}") + if tensor.shape[0] == 1: + if tensor.shape[1] < expected_total: + raise ValueError(f"Captured {name} has {tensor.shape[1]} tokens, expected {expected_total}") + return tensor[0, :expected_total] + + assert self._causal_offsets is not None + lengths = torch.diff(self._causal_offsets.to(device="cpu", dtype=torch.int64)).tolist() + if tensor.shape[0] != len(lengths): + raise ValueError(f"Captured {name} batch dimension {tensor.shape[0]} does not match {len(lengths)} samples") + if any(length > tensor.shape[1] for length in lengths): + raise ValueError(f"Captured {name} rows are too short for per-sample lengths {lengths}") + return torch.cat([tensor[sample_idx, :length] for sample_idx, length in enumerate(lengths)], dim=0) + + def to_feature_batch(self) -> ReasonerFeatureBatch: + """Return the completed canonical batch.""" + if self._causal_offsets is None or not self.is_gen_only(): + missing = [idx for idx, layer in enumerate(self._layers) if layer is None] + raise RuntimeError(f"Reasoner K/V capture is incomplete; missing layers {missing}") + layers = tuple(layer for layer in self._layers if layer is not None) + return ReasonerFeatureBatch( + tuple(layer.cross_k for layer in layers), + tuple(layer.cross_v for layer in layers), + self._causal_offsets, + self._fingerprints, + ) + + def is_gen_only(self) -> bool: + return all(layer is not None for layer in self._layers) + + def requires_natten_metadata(self) -> bool: + return False + + +def _validate_cached_attention_mode( + packed_query_states: SequencePack, + packed_key_states: SequencePack, + packed_value_states: SequencePack, + attention_mask: object, + natten_metadata: dict | None, +) -> None: + if natten_metadata is not None: + raise ValueError("Static Reasoner K/V supports only two-way attention; NATTEN metadata was provided") + if any(pack.get("is_sharded", False) for pack in (packed_query_states, packed_key_states, packed_value_states)): + raise ValueError("Static Reasoner K/V does not support context-parallel sharded SequencePacks") + if not hasattr(attention_mask, "is_three_way"): + raise TypeError(f"Unsupported attention metadata: {type(attention_mask)}") + if bool(getattr(attention_mask, "is_three_way")): + raise ValueError("Static Reasoner K/V supports only two-way attention, not three-way attention") + unsupported_fields = ( + "control_stream_token_ranges", + "flex_block_mask", + "multiview_maskless", + ) + enabled = [field for field in unsupported_fields if getattr(attention_mask, field, None) is not None] + if enabled: + raise ValueError(f"Static Reasoner K/V does not support specialized attention metadata: {enabled}") + + +def _attention_gen_with_reasoner_features( + packed_query_states: SequencePack, + packed_key_states: SequencePack, + packed_value_states: SequencePack, + memory_value: StaticReasonerKVMemoryValue, +) -> tuple[SequencePack, KVToStore | None]: + q_gen = get_gen_seq(packed_query_states) + k_gen = get_gen_seq(packed_key_states) + v_gen = get_gen_seq(packed_value_states) + if q_gen.ndim != 3 or k_gen.ndim != 3 or v_gen.ndim != 3: + raise ValueError( + "GEN Q/K/V must use [sequence, heads, head_dim] layout, " + f"got Q={q_gen.shape} K={k_gen.shape} V={v_gen.shape}" + ) + if k_gen.shape != v_gen.shape: + raise ValueError(f"GEN K/V shapes must match, got K={k_gen.shape} V={v_gen.shape}") + if memory_value.gen_len > min(q_gen.shape[0], k_gen.shape[0], v_gen.shape[0]): + raise ValueError( + f"GEN real-token count {memory_value.gen_len} exceeds Q/K/V stream lengths " + f"{q_gen.shape[0]}/{k_gen.shape[0]}/{v_gen.shape[0]}" + ) + cross_k = memory_value.cross_k + cross_v = memory_value.cross_v + if cross_k.ndim != 3 or cross_v.shape != cross_k.shape: + raise ValueError(f"Cached Reasoner K/V have invalid shapes K={cross_k.shape} V={cross_v.shape}") + if cross_k.shape[1:] != k_gen.shape[1:]: + raise ValueError( + f"Cached Reasoner and live GEN K/V head shapes disagree: {cross_k.shape[1:]} vs {k_gen.shape[1:]}" + ) + if cross_k.device != k_gen.device or cross_v.device != v_gen.device: + raise ValueError( + f"Cached Reasoner and live GEN K/V devices disagree: {cross_k.device}/{cross_v.device} " + f"vs {k_gen.device}/{v_gen.device}" + ) + if cross_k.dtype != k_gen.dtype or cross_v.dtype != v_gen.dtype: + raise TypeError( + f"Cached Reasoner and live GEN K/V dtypes disagree: {cross_k.dtype}/{cross_v.dtype} " + f"vs {k_gen.dtype}/{v_gen.dtype}" + ) + + gen_len = memory_value.gen_len + q_real = q_gen[:gen_len].unsqueeze(0) + k_real = k_gen[:gen_len] + v_real = v_gen[:gen_len] + k_full = torch.cat((cross_k, k_real), dim=0) + v_full = torch.cat((cross_v, v_real), dim=0) + + if memory_value.kv_reorder_indices is None: + attn_result = attention( + query=q_real, + key=k_full.unsqueeze(0), + value=v_full.unsqueeze(0), + is_causal=False, + return_lse=False, + ) + else: + if ( + memory_value.cumulative_seqlen_q is None + or memory_value.cumulative_seqlen_kv is None + or memory_value.max_seqlen_q <= 0 + or memory_value.max_seqlen_kv <= 0 + ): + raise ValueError("Multi-sample Reasoner K/V is missing varlen attention metadata") + k_full = k_full.index_select(0, memory_value.kv_reorder_indices) + v_full = v_full.index_select(0, memory_value.kv_reorder_indices) + attn_result = attention( + query=q_real, + key=k_full.unsqueeze(0), + value=v_full.unsqueeze(0), + is_causal=False, + return_lse=False, + cumulative_seqlen_Q=memory_value.cumulative_seqlen_q, + cumulative_seqlen_KV=memory_value.cumulative_seqlen_kv, + max_seqlen_Q=memory_value.max_seqlen_q, + max_seqlen_KV=memory_value.max_seqlen_kv, + ) + + if not isinstance(attn_result, torch.Tensor): + raise TypeError(f"Attention returned {type(attn_result).__name__}, expected torch.Tensor") + gen_out_real = attn_result.squeeze(0).flatten(-2, -1) + gen_out = q_gen.new_zeros((q_gen.shape[0], gen_out_real.shape[-1])) + gen_out[:gen_len] = gen_out_real + empty_und = gen_out.new_empty((0, gen_out.shape[-1])) + return from_und_gen_splits(empty_und, gen_out, packed_query_states), None + + +def dispatch_attention_with_reasoner_features( + packed_query_states: SequencePack, + packed_key_states: SequencePack, + packed_value_states: SequencePack, + attention_mask: object | SplitInfo, + natten_metadata: dict | None = None, + memory_value: MemoryValue | None = None, + packed_key_states_normalized: SequencePack | None = None, +) -> tuple[SequencePack, KVToStore | None]: + """Dispatch ordinary attention, inline capture, or cached GEN attention.""" + if isinstance(memory_value, (StaticReasonerKVMemoryValue, ReasonerKVCaptureMemoryValue)): + _validate_cached_attention_mode( + packed_query_states, + packed_key_states, + packed_value_states, + attention_mask, + natten_metadata, + ) + if isinstance(memory_value, StaticReasonerKVMemoryValue): + return _attention_gen_with_reasoner_features( + packed_query_states, + packed_key_states, + packed_value_states, + memory_value, + ) + + # The capture marker must reach PackedAttentionMoT (so it emits + # kv_to_store), but the ordinary dispatcher does not consume memory. + delegated_memory = None if isinstance(memory_value, ReasonerKVCaptureMemoryValue) else memory_value + return dispatch_attention( + packed_query_states, + packed_key_states, + packed_value_states, + attention_mask, + natten_metadata=natten_metadata, + memory_value=delegated_memory, + packed_key_states_normalized=packed_key_states_normalized, + ) + + +ReasonerDispatchSnapshot = list[tuple[torch.nn.Module, object]] + + +def install_reasoner_feature_attention_dispatch(net: torch.nn.Module) -> ReasonerDispatchSnapshot: + """Install the cached dispatcher and return an exact restoration snapshot.""" + try: + layers = net.language_model.model.layers + except AttributeError as error: + raise TypeError("Expected a model with net.language_model.model.layers") from error + + previous: ReasonerDispatchSnapshot = [] + try: + for layer in layers: + attn = layer.self_attn + current = attn.dispatch_attention_fn + if current is not dispatch_attention and current is not dispatch_attention_with_reasoner_features: + current_name = getattr(current, "__name__", type(current).__name__) + raise RuntimeError( + "Cannot install Reasoner feature attention over " + f"{current_name}; only the default non-CP dispatcher is supported" + ) + previous.append((attn, current)) + attn.dispatch_attention_fn = dispatch_attention_with_reasoner_features + except Exception: + restore_reasoner_feature_attention_dispatch(previous) + raise + return previous + + +def restore_reasoner_feature_attention_dispatch(previous: ReasonerDispatchSnapshot) -> None: + """Restore dispatchers returned by :func:`install_reasoner_feature_attention_dispatch`.""" + for attn, previous_fn in previous: + attn.dispatch_attention_fn = previous_fn + + +__all__ = [ + "CapturingReasonerKVMemoryState", + "ReasonerAttentionMetadata", + "ReasonerFeatureBatch", + "ReasonerFeatureProvider", + "ReasonerFeatureRequest", + "ReasonerKVCaptureMemoryValue", + "ReasonerLayerKV", + "StaticReasonerKVMemoryState", + "StaticReasonerKVMemoryValue", + "build_reasoner_attention_metadata", + "dispatch_attention_with_reasoner_features", + "extract_reasoner_feature_batch", + "install_reasoner_feature_attention_dispatch", + "restore_reasoner_feature_attention_dispatch", +] diff --git a/cosmos_framework/model/generator/reasoner_features_test.py b/cosmos_framework/model/generator/reasoner_features_test.py new file mode 100644 index 000000000..fcd4a4353 --- /dev/null +++ b/cosmos_framework/model/generator/reasoner_features_test.py @@ -0,0 +1,501 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +from types import SimpleNamespace + +import pytest +import torch + +import cosmos_framework.model.generator.reasoner_features as reasoner_features +from cosmos_framework.data.generator.sequence_packing.runtime import get_gen_seq +from cosmos_framework.model.attention.utils import is_blackwell_dc, is_hopper +from cosmos_framework.model.generator.mot.attention import build_packed_sequence +from cosmos_framework.model.generator.mot.unified_mot import ( + Qwen3VLMoTConfig, + Qwen3VLTextForCausalLM, + prune_und_pathway_, +) +from cosmos_framework.model.generator.reasoner_features import ( + CapturingReasonerKVMemoryState, + ReasonerFeatureBatch, + ReasonerFeatureRequest, + ReasonerLayerKV, + StaticReasonerKVMemoryState, + StaticReasonerKVMemoryValue, + extract_reasoner_feature_batch, +) + + +def _pack( + sequence: torch.Tensor, + und_offsets: torch.Tensor, + gen_offsets: torch.Tensor, + *, + is_sharded: bool = False, +) -> dict: + num_samples = gen_offsets.numel() - 1 + return { + "causal_seq": sequence.new_empty((0, *sequence.shape[1:])), + "full_only_seq": sequence, + "sample_offsets": torch.arange(num_samples + 1, dtype=torch.int32), + "max_sample_len": 1, + "max_causal_len": int(torch.diff(und_offsets).max()), + "max_full_len": int(torch.diff(gen_offsets).max()), + "_causal_indices": torch.empty(0, dtype=torch.int64), + "_full_indices": torch.arange(sequence.shape[0]), + "_causal_seq_offsets": und_offsets, + "_full_only_seq_offsets": gen_offsets, + "_num_causal_tokens": int(und_offsets[-1]), + "_num_full_tokens": int(gen_offsets[-1]), + "is_sharded": is_sharded, + } + + +def _two_way_mask() -> SimpleNamespace: + return SimpleNamespace( + is_three_way=False, + control_stream_token_ranges=None, + flex_block_mask=None, + multiview_maskless=None, + ) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_canonical_batch_validates_shapes_detaches_and_round_trips_cache_tensors() -> None: + cross_k = torch.arange(24.0, requires_grad=True).reshape(2, 3, 2, 2) + cross_v = (cross_k + 100).detach().requires_grad_() + offsets = torch.tensor([0, 2, 3], dtype=torch.int32) + + batch = ReasonerFeatureBatch.from_stacked(cross_k, cross_v, offsets, ("first", "second")) + + assert batch.num_layers == 2 + assert batch.num_samples == 2 + assert batch.sequence_length == 3 + assert all(not tensor.requires_grad for tensor in (*batch.cross_k, *batch.cross_v)) + + restored = ReasonerFeatureBatch.from_cache_tensors( + batch.to_cache_tensors(), + fingerprints=batch.fingerprints, + ) + restored_k, restored_v = restored.to_stacked() + assert torch.equal(restored_k, cross_k) + assert torch.equal(restored_v, cross_v) + assert torch.equal(restored.causal_offsets, offsets) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_extractor_rejects_multiple_causal_documents_per_request() -> None: + request = ReasonerFeatureRequest( + sample_key="per-view", + token_ids=torch.tensor([1, 2]), + position_ids=torch.zeros(3, 2), + causal_offsets=torch.tensor([0, 1, 2]), + fingerprint="fingerprint", + ) + + with pytest.raises(NotImplementedError, match="one causal document"): + extract_reasoner_feature_batch(object(), [request]) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +@pytest.mark.parametrize( + ("cross_k", "cross_v", "error"), + [ + (torch.zeros(3, 2), torch.zeros(3, 2), "shape"), + (torch.zeros(3, 2, 4), torch.zeros(4, 2, 4), "shapes must match"), + (torch.zeros(3, 2, 4), torch.zeros(3, 3, 4), "shapes must match"), + ], +) +def test_layer_kv_rejects_invalid_shapes( + cross_k: torch.Tensor, + cross_v: torch.Tensor, + error: str, +) -> None: + with pytest.raises(ValueError, match=error): + ReasonerLayerKV(cross_k, cross_v) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_static_state_rejects_sequence_pack_offset_mismatch() -> None: + features = ReasonerFeatureBatch( + (torch.zeros(3, 1, 2),), + (torch.zeros(3, 1, 2),), + torch.tensor([0, 2, 3], dtype=torch.int32), + ("first", "second"), + ) + state = StaticReasonerKVMemoryState(features) + hidden = torch.zeros(3, 1, 2) + pack = _pack( + hidden, + torch.tensor([0, 1, 3], dtype=torch.int32), + torch.tensor([0, 1, 3], dtype=torch.int32), + ) + + with pytest.raises(ValueError, match="causal offsets do not match"): + state.init(pack, torch.device("cpu")) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_static_state_rejects_reasoner_gen_sample_count_mismatch() -> None: + features = ReasonerFeatureBatch( + (torch.zeros(3, 1, 2),), + (torch.zeros(3, 1, 2),), + torch.tensor([0, 2, 3], dtype=torch.int32), + ("first", "second"), + ) + state = StaticReasonerKVMemoryState(features) + hidden = torch.zeros(3, 1, 2) + pack = _pack( + hidden, + torch.tensor([0, 2, 3], dtype=torch.int32), + torch.tensor([0, 3], dtype=torch.int32), + ) + + with pytest.raises(ValueError, match="sample counts disagree"): + state.init(pack, torch.device("cpu")) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_multi_sample_cached_attention_isolates_samples_and_zero_fills_padding( + monkeypatch: pytest.MonkeyPatch, +) -> None: + und_offsets = torch.tensor([0, 2, 3], dtype=torch.int32) + gen_offsets = torch.tensor([0, 1, 3], dtype=torch.int32) + cached_k = torch.tensor([10.0, 11.0, 20.0]).reshape(3, 1, 1).requires_grad_() + cached_v = cached_k.detach().clone().requires_grad_() + features = ReasonerFeatureBatch((cached_k,), (cached_v,), und_offsets, ("first", "second")) + state = StaticReasonerKVMemoryState(features) + + q_gen = torch.tensor([1.0, 2.0, 3.0, 999.0]).reshape(4, 1, 1).requires_grad_() + k_gen = torch.tensor([100.0, 200.0, 201.0, 999.0]).reshape(4, 1, 1).requires_grad_() + v_gen = k_gen.detach().clone().requires_grad_() + state.init(_pack(q_gen, und_offsets, gen_offsets), torch.device("cpu")) + memory_value = state.read_for_layer(0) + captured: dict[str, object] = {} + + def fake_attention( + *, query: torch.Tensor, key: torch.Tensor, value: torch.Tensor, **kwargs: object + ) -> torch.Tensor: + captured.update(query=query, key=key, value=value, **kwargs) + return query + 5 + + monkeypatch.setattr(reasoner_features, "attention", fake_attention) + output, kv_to_store = reasoner_features.dispatch_attention_with_reasoner_features( + _pack(q_gen, und_offsets, gen_offsets), + _pack(k_gen, und_offsets, gen_offsets), + _pack(v_gen, und_offsets, gen_offsets), + _two_way_mask(), + memory_value=memory_value, + ) + + assert kv_to_store is None + assert isinstance(memory_value, StaticReasonerKVMemoryValue) + assert captured["query"].flatten().tolist() == [1.0, 2.0, 3.0] + assert captured["key"].flatten().tolist() == [10.0, 11.0, 100.0, 20.0, 200.0, 201.0] + assert captured["value"].flatten().tolist() == [10.0, 11.0, 100.0, 20.0, 200.0, 201.0] + assert torch.equal(captured["cumulative_seqlen_Q"], gen_offsets) + assert torch.equal(captured["cumulative_seqlen_KV"], torch.tensor([0, 3, 6], dtype=torch.int32)) + assert captured["max_seqlen_Q"] == 2 + assert captured["max_seqlen_KV"] == 3 + assert output["full_only_seq"].flatten().tolist() == [6.0, 7.0, 8.0, 0.0] + + output["full_only_seq"][:3].sum().backward() + assert q_gen.grad is not None + assert k_gen.grad is None # fake attention consumes only Q; cached tensors stay detached either way + assert cached_k.grad is None + assert cached_v.grad is None + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_capture_state_detaches_inline_kv_and_switches_to_gen_only() -> None: + und_offsets = torch.tensor([0, 2, 3], dtype=torch.int32) + gen_offsets = torch.tensor([0, 1, 3], dtype=torch.int32) + state = CapturingReasonerKVMemoryState(1, fingerprints=("first", "second")) + state.init(_pack(torch.zeros(3, 1, 1), und_offsets, gen_offsets), torch.device("cpu")) + + assert not state.is_gen_only() + captured_k = torch.tensor([10.0, 11.0, 20.0]).reshape(1, 3, 1, 1).requires_grad_() + captured_v = (captured_k + 1).detach().requires_grad_() + generated = torch.zeros(1, 3, 1, 1) + state.write_for_layer(0, (generated, generated, captured_k, captured_v)) + + assert state.is_gen_only() + batch = state.to_feature_batch() + assert not batch.cross_k[0].requires_grad + assert not batch.cross_v[0].requires_grad + assert batch.cross_k[0].untyped_storage().data_ptr() != captured_k.untyped_storage().data_ptr() + assert batch.cross_v[0].untyped_storage().data_ptr() != captured_v.untyped_storage().data_ptr() + assert batch.cross_k[0].flatten().tolist() == [10.0, 11.0, 20.0] + + state.init(_pack(torch.zeros(3, 1, 1), und_offsets, gen_offsets), torch.device("cpu")) + assert isinstance(state.read_for_layer(0), StaticReasonerKVMemoryValue) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_cached_dispatch_fails_closed_for_context_parallel_pack() -> None: + offsets = torch.tensor([0, 2], dtype=torch.int32) + features = ReasonerFeatureBatch((torch.zeros(2, 1, 1),), (torch.zeros(2, 1, 1),), offsets, ("sample",)) + state = StaticReasonerKVMemoryState(features) + sequence = torch.zeros(2, 1, 1) + ordinary_pack = _pack(sequence, offsets, offsets) + state.init(ordinary_pack, torch.device("cpu")) + sharded_pack = _pack(sequence, offsets, offsets, is_sharded=True) + + with pytest.raises(ValueError, match="context-parallel"): + reasoner_features.dispatch_attention_with_reasoner_features( + sharded_pack, + sharded_pack, + sharded_pack, + _two_way_mask(), + memory_value=state.read_for_layer(0), + ) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_install_and_restore_reasoner_feature_dispatch() -> None: + attention_module = SimpleNamespace(dispatch_attention_fn=reasoner_features.dispatch_attention) + net = SimpleNamespace( + language_model=SimpleNamespace( + model=SimpleNamespace( + layers=[SimpleNamespace(self_attn=attention_module)], + ) + ) + ) + + previous = reasoner_features.install_reasoner_feature_attention_dispatch(net) + assert attention_module.dispatch_attention_fn is reasoner_features.dispatch_attention_with_reasoner_features + + reasoner_features.restore_reasoner_feature_attention_dispatch(previous) + assert attention_module.dispatch_attention_fn is reasoner_features.dispatch_attention + + +def _tiny_two_way_model_pack( + packed_sequence: torch.Tensor, +) -> tuple[dict, object, torch.Tensor]: + """Build a two-sample, two-way pack in original ``[UND_i, GEN_i]`` order.""" + und_lens = (3, 4) + gen_lens = (5, 3) + split_lens: list[int] = [] + sample_lens: list[int] = [] + und_indices: list[int] = [] + gen_indices: list[int] = [] + position_ids: list[torch.Tensor] = [] + offset = 0 + for und_len, gen_len in zip(und_lens, gen_lens): + split_lens.extend((und_len, gen_len)) + sample_lens.append(und_len + gen_len) + und_indices.extend(range(offset, offset + und_len)) + gen_indices.extend(range(offset + und_len, offset + und_len + gen_len)) + position_ids.append(torch.arange(und_len + gen_len, device=packed_sequence.device)) + offset += und_len + gen_len + + pack, attention_mask, natten_metadata = build_packed_sequence( + "two_way", + packed_sequence=packed_sequence, + attn_modes=["causal", "full", "causal", "full"], + split_lens=split_lens, + sample_lens=sample_lens, + packed_und_token_indexes=torch.tensor(und_indices, device=packed_sequence.device), + packed_gen_token_indexes=torch.tensor(gen_indices, device=packed_sequence.device), + num_heads=4, + head_dim=64, + num_layers=2, + is_image_batch=True, + ) + assert natten_metadata is None + return pack, attention_mask, torch.cat(position_ids) + + +@pytest.mark.level(1) +@pytest.mark.gpus(1) +@pytest.mark.skipif( + not torch.cuda.is_available() or (not is_hopper() and not is_blackwell_dc()), + reason="MoT attention parity requires a Hopper or Blackwell CUDA GPU", +) +def test_und_only_extractor_matches_joint_capture() -> None: + """The reasoner-only prefill emits the same per-layer K/V as the joint path.""" + device = torch.device("cuda") + dtype = torch.bfloat16 + torch.manual_seed(4321) + config = Qwen3VLMoTConfig( + { + "text_config": { + "vocab_size": 128, + "hidden_size": 256, + "intermediate_size": 512, + "num_hidden_layers": 2, + "num_attention_heads": 4, + "num_key_value_heads": 2, + "head_dim": 64, + "rms_norm_eps": 1e-6, + "rope_theta": 5000000.0, + "max_position_embeddings": 128, + "tie_word_embeddings": False, + } + } + ) + model = Qwen3VLTextForCausalLM(config).to(device=device, dtype=dtype).eval() + model.model.rotary_emb.init_weights(buffer_device=device) + token_ids = (torch.tensor([4, 5, 6], device=device), torch.tensor([10, 11, 12, 13], device=device)) + requests = tuple( + ReasonerFeatureRequest( + sample_key=f"sample-{sample_idx}", + token_ids=tokens, + position_ids=torch.arange(tokens.numel(), device=device), + causal_offsets=torch.tensor([0, tokens.numel()], device=device), + fingerprint=f"fingerprint-{sample_idx}", + ) + for sample_idx, tokens in enumerate(token_ids) + ) + + extracted = extract_reasoner_feature_batch(model, requests) + + base_embeddings = torch.randn(15, 256, device=device, dtype=dtype) + und_indexes = torch.tensor([0, 1, 2, 8, 9, 10, 11], device=device) + base_embeddings[und_indexes] = model.model.embed_tokens(torch.cat(token_ids)) + joint_pack, joint_mask, position_ids = _tiny_two_way_model_pack(base_embeddings) + capture = CapturingReasonerKVMemoryState(2, fingerprints=("fingerprint-0", "fingerprint-1")) + owner = SimpleNamespace(language_model=model) + previous_dispatchers = reasoner_features.install_reasoner_feature_attention_dispatch(owner) + try: + with torch.no_grad(): + model(joint_pack, attention_mask=joint_mask, position_ids=position_ids, memory=capture) + captured = capture.to_feature_batch() + finally: + reasoner_features.restore_reasoner_feature_attention_dispatch(previous_dispatchers) + + assert captured.causal_offsets.tolist() == extracted.causal_offsets.tolist() + for layer_idx in range(2): + torch.testing.assert_close(extracted.cross_k[layer_idx], captured.cross_k[layer_idx], rtol=2e-2, atol=2e-2) + torch.testing.assert_close(extracted.cross_v[layer_idx], captured.cross_v[layer_idx], rtol=2e-2, atol=2e-2) + + +@pytest.mark.level(1) +@pytest.mark.gpus(1) +@pytest.mark.skipif( + not torch.cuda.is_available() or (not is_hopper() and not is_blackwell_dc()), + reason="MoT attention parity requires a Hopper or Blackwell CUDA GPU", +) +def test_tiny_unified_mot_pruned_structure_matches_full_cached_forward_and_gradients() -> None: + """A structurally pruned consumer preserves cached GEN outputs and gradients.""" + device = torch.device("cuda") + dtype = torch.bfloat16 + torch.manual_seed(1234) + config = Qwen3VLMoTConfig( + { + "text_config": { + "vocab_size": 128, + "hidden_size": 256, + "intermediate_size": 512, + "num_hidden_layers": 2, + "num_attention_heads": 4, + "num_key_value_heads": 2, + "head_dim": 64, + "rms_norm_eps": 1e-6, + "rope_theta": 5000000.0, + "max_position_embeddings": 128, + "tie_word_embeddings": False, + } + } + ) + full_model = Qwen3VLTextForCausalLM(config) + generator_only_model = Qwen3VLTextForCausalLM(config) + prune_und_pathway_(generator_only_model) + + generator_only_state_names = set(generator_only_model.state_dict()) + assert generator_only_state_names + assert all("moe_gen" in name for name in generator_only_state_names) + full_generator_state = { + name: tensor for name, tensor in full_model.state_dict().items() if name in generator_only_state_names + } + assert full_generator_state.keys() == generator_only_state_names + generator_only_model.load_state_dict(full_generator_state, strict=True) + + full_model = full_model.to(device=device, dtype=dtype).train() + generator_only_model = generator_only_model.to(device=device, dtype=dtype).train() + # Casting the module also casts non-persistent buffers; RoPE deliberately + # keeps inv_freq in FP32, so restore it exactly as production materialization does. + full_model.model.rotary_emb.init_weights(buffer_device=device) + generator_only_model.model.rotary_emb.init_weights(buffer_device=device) + for name, parameter in full_model.named_parameters(): + parameter.requires_grad_("moe_gen" in name) + full_generator_parameters = { + name: parameter for name, parameter in full_model.named_parameters() if parameter.requires_grad + } + pruned_generator_parameters = dict(generator_only_model.named_parameters()) + assert full_generator_parameters.keys() == pruned_generator_parameters.keys() == generator_only_state_names + + base_embeddings = torch.randn(15, 256, device=device, dtype=dtype) + target = torch.randn(8, 256, device=device, dtype=torch.float32) + capture = CapturingReasonerKVMemoryState(2, fingerprints=("sample-0", "sample-1")) + full_owner = SimpleNamespace(language_model=full_model) + pruned_owner = SimpleNamespace(language_model=generator_only_model) + full_dispatchers = reasoner_features.install_reasoner_feature_attention_dispatch(full_owner) + pruned_dispatchers = reasoner_features.install_reasoner_feature_attention_dispatch(pruned_owner) + try: + capture_pack, capture_mask, capture_position_ids = _tiny_two_way_model_pack(base_embeddings.clone()) + with torch.no_grad(): + full_model( + capture_pack, + attention_mask=capture_mask, + position_ids=capture_position_ids, + memory=capture, + ) + assert capture.is_gen_only() + + full_pack, full_mask, full_position_ids = _tiny_two_way_model_pack(base_embeddings.clone()) + full_output, _ = full_model( + full_pack, + attention_mask=full_mask, + position_ids=full_position_ids, + memory=capture, + ) + full_gen = get_gen_seq(full_output)[: full_pack["_num_full_tokens"]] + full_loss = (full_gen.float() - target).square().mean() + full_loss.backward() + full_gradients = { + name: parameter.grad.detach().clone() + for name, parameter in full_generator_parameters.items() + if parameter.grad is not None + } + assert full_gradients.keys() == full_generator_parameters.keys() + + pruned_pack, pruned_mask, pruned_position_ids = _tiny_two_way_model_pack(base_embeddings.clone()) + pruned_output, _ = generator_only_model( + pruned_pack, + attention_mask=pruned_mask, + position_ids=pruned_position_ids, + memory=capture, + ) + pruned_gen = get_gen_seq(pruned_output)[: pruned_pack["_num_full_tokens"]] + pruned_loss = (pruned_gen.float() - target).square().mean() + pruned_loss.backward() + pruned_gradients = { + name: parameter.grad.detach().clone() + for name, parameter in pruned_generator_parameters.items() + if parameter.grad is not None + } + + torch.testing.assert_close(pruned_gen, full_gen, rtol=2e-2, atol=2e-2) + assert pruned_gradients.keys() == full_gradients.keys() + for name in full_gradients: + torch.testing.assert_close( + pruned_gradients[name], + full_gradients[name], + rtol=3e-2, + atol=3e-3, + msg=lambda message, name=name: f"Generator gradient mismatch for {name}: {message}", + ) + assert all(parameter.grad is None for name, parameter in full_model.named_parameters() if "moe_gen" not in name) + finally: + reasoner_features.restore_reasoner_feature_attention_dispatch(pruned_dispatchers) + reasoner_features.restore_reasoner_feature_attention_dispatch(full_dispatchers) diff --git a/cosmos_framework/scripts/extract_reasoner_features.py b/cosmos_framework/scripts/extract_reasoner_features.py new file mode 100644 index 000000000..188965796 --- /dev/null +++ b/cosmos_framework/scripts/extract_reasoner_features.py @@ -0,0 +1,669 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +"""Build an immutable offline Reasoner K/V cache for Nano generator SFT. + +Each torchrun rank owns a deterministic slice of the SFT metadata, loads only +the frozen Reasoner tower, and writes resumable rank-local safetensors shards. +Rank zero publishes the cache only after every rank has completed successfully. + +Example:: + + torchrun --nproc-per-node=8 -m cosmos_framework.scripts.extract_reasoner_features \\ + --sft-toml examples/toml/sft_config/vision_sft_nano.toml \\ + --checkpoint /shared/checkpoints/iter_000000100 \\ + --checkpoint-source regular \\ + --output /shared/caches/nano-sft-reasoner-kv \\ + --reasoner-fingerprint \\ + --tokenizer-fingerprint \\ + --framing-fingerprint \\ + -- model.config.compile.enabled=false + +The current MVP supports local/shared-POSIX DCP checkpoints and cache output, +Nano's standard single-caption two-way-attention recipe, BF16, and CP=1. +""" + +from __future__ import annotations + +import argparse +import copy +import hashlib +import json +import os +import tempfile +import time +import traceback +from collections.abc import Iterator, Mapping, Sequence +from contextlib import contextmanager +from dataclasses import asdict, dataclass +from pathlib import Path +from typing import Any + +import torch +from omegaconf import open_dict + +from cosmos_framework.checkpoint.reasoner_only import load_reasoner_only_dcp +from cosmos_framework.configs.toml_config.sft_config import load_experiment_from_toml +from cosmos_framework.data.generator.local_datasets.sft_dataset import ( + SFTDataset, + get_sft_dataset, + sft_metadata_sort_key, +) +from cosmos_framework.data.generator.local_datasets.sft_reasoner_documents import ( + iter_sft_reasoner_documents, +) +from cosmos_framework.data.generator.sequence_packing.modalities import add_special_tokens +from cosmos_framework.model.generator.reasoner_feature_cache import ( + IncrementalReasonerFeatureCacheWriter, + OfflineReasonerFeatureProvider, + ReasonerFeatureCacheEntry, + ReasonerFeatureCacheIdentity, + build_reasoner_feature_request_from_text_tokens, + finalize_incremental_reasoner_feature_cache, +) +from cosmos_framework.model.generator.reasoner_features import ( + ReasonerFeatureRequest, + extract_reasoner_feature_batch, +) +from cosmos_framework.utils.lazy_config import instantiate as lazy_instantiate + +_NULL_PROMPT_KEY = "__null__" +_RANK_STATS_FILE = "extraction.stats.json" +_RANK_FAILURE_FILE = "extraction.failure.json" + + +@dataclass(frozen=True) +class DistributedContext: + rank: int + world_size: int + local_rank: int + device: torch.device + run_id: str + + +@dataclass +class ExtractionStats: + rank: int + assigned_documents: int = 0 + cache_hits: int = 0 + written_records: int = 0 + extracted_tokens: int = 0 + reasoner_construction_seconds: float = 0.0 + checkpoint_load_seconds: float = 0.0 + extraction_seconds: float = 0.0 + peak_allocated_bytes: int = 0 + peak_reserved_bytes: int = 0 + + +def _target_matches_get_sft_dataset(target: object) -> bool: + if target is get_sft_dataset: + return True + if isinstance(target, str): + return target.rsplit(".", 1)[-1] == "get_sft_dataset" + return getattr(target, "__name__", None) == "get_sft_dataset" + + +def _find_sft_dataset_config(root: object) -> object: + """Find exactly one lazy ``get_sft_dataset`` node below a dataloader config.""" + + matches: list[object] = [] + visited: set[int] = set() + + def visit(value: object) -> None: + value_id = id(value) + if value_id in visited: + return + visited.add(value_id) + if isinstance(value, Mapping): + target = value.get("_target_") + if _target_matches_get_sft_dataset(target): + matches.append(value) + return + for child in value.values(): + visit(child) + elif isinstance(value, Sequence) and not isinstance(value, (str, bytes, bytearray)): + for child in value: + visit(child) + + visit(root) + if len(matches) != 1: + raise ValueError( + "Reasoner extraction requires exactly one get_sft_dataset node under dataloader_train; " + f"found {len(matches)}" + ) + return matches[0] + + +def _instantiate_sft_dataset(config: object) -> SFTDataset: + dataloader_config = getattr(config, "dataloader_train") + dataset_config = _find_sft_dataset_config(dataloader_config) + dataset = lazy_instantiate(dataset_config) + if not isinstance(dataset, SFTDataset): + raise TypeError(f"Expected get_sft_dataset to return SFTDataset, got {type(dataset).__name__}") + return dataset + + +def _validate_extraction_config(config: object) -> None: + model = getattr(getattr(config, "model"), "config") + if model.joint_attn_implementation != "two_way": + raise ValueError("Reasoner extraction currently requires joint_attn_implementation='two_way'") + if model.parallelism.context_parallel_shard_degree != 1: + raise ValueError("Reasoner extraction currently requires context_parallel_shard_degree=1") + if model.video_temporal_causal: + raise ValueError("Reasoner extraction does not support video_temporal_causal training") + if model.causal_training_strategy != "none": + raise ValueError("Reasoner extraction currently requires causal_training_strategy='none'") + if model.diffusion_expert_config.vision_temporal_position_mode != "latent_index": + raise ValueError("Reasoner extraction currently supports vision_temporal_position_mode='latent_index' only") + + model_instance = model.vlm_config.model_instance + if model_instance is None: + raise ValueError("Reasoner extraction requires model.config.vlm_config.model_instance") + target = model_instance.get("_target_") + target_name = target if isinstance(target, str) else getattr(target, "__name__", repr(target)) + if "Qwen3VLTextForCausalLM" not in target_name: + raise ValueError(f"Reasoner extraction MVP supports the Nano Qwen3VLTextForCausalLM target, got {target_name}") + + +def _prepare_reasoner_model_config(config: object) -> object: + """Return a lazy Nano LM config that cannot instantiate the Generator or ViT.""" + + model = getattr(getattr(config, "model"), "config") + model_instance = copy.deepcopy(model.vlm_config.model_instance) + nested_config = model_instance["config"] + if isinstance(nested_config, dict): + nested_config.update( + include_gen_pathway=False, + include_und_pathway=True, + include_visual=False, + ) + else: + with open_dict(nested_config): + nested_config.include_gen_pathway = False + nested_config.include_und_pathway = True + nested_config.include_visual = False + return model_instance + + +@contextmanager +def _temporary_default_dtype(dtype: torch.dtype) -> Iterator[None]: + previous = torch.get_default_dtype() + torch.set_default_dtype(dtype) + try: + yield + finally: + torch.set_default_dtype(previous) + + +def _build_reasoner(config: object, *, device: torch.device, dtype: torch.dtype) -> torch.nn.Module: + """Construct a materialized Reasoner replica directly in its compute dtype. + + Direct construction intentionally avoids meta ``to_empty`` here: HuggingFace + rotary buffers are non-persistent and therefore absent from DCP. Constructing + on the final device initializes those buffers correctly before parameters are + overwritten by the strict Reasoner-only checkpoint load. + """ + + model_instance = _prepare_reasoner_model_config(config) + with _temporary_default_dtype(dtype), torch.device(device): + reasoner = lazy_instantiate(model_instance) + generation_parameters = [name for name, _ in reasoner.named_parameters() if "moe_gen" in name] + if generation_parameters: + raise RuntimeError(f"Reasoner-only construction retained Generator parameters: {generation_parameters}") + reasoner.requires_grad_(False) + reasoner.eval() + return reasoner + + +def _special_tokens(dataset: SFTDataset) -> dict[str, int]: + tokenizer, special_tokens = add_special_tokens(dataset.vlm_tokenizer) + eos_token_id = tokenizer.eos_token_id + if eos_token_id is None: + raise ValueError("SFT tokenizer must define eos_token_id") + return {**special_tokens, "eos_token_id": int(eos_token_id)} + + +def _rank_dataset(dataset: SFTDataset, *, rank: int, world_size: int) -> SFTDataset: + """Shallow-copy a dataset and assign it a deterministic video-level slice.""" + + result = copy.copy(dataset) + ordered = sorted(dataset.metadata, key=sft_metadata_sort_key) + result.metadata = ordered[rank::world_size] + return result + + +def _request_owner(request: ReasonerFeatureRequest, world_size: int) -> int: + digest = hashlib.sha256(f"{request.sample_key}\0{request.fingerprint}".encode()).digest() + return int.from_bytes(digest[:8], byteorder="big") % world_size + + +def _iter_rank_requests( + dataset: SFTDataset, + *, + identity: ReasonerFeatureCacheIdentity, + rank: int, + world_size: int, + use_float_positions: bool, + max_documents_per_rank: int | None, +) -> Iterator[ReasonerFeatureRequest]: + """Yield framed requests for one deterministic metadata slice. + + The standard CFG mode (drop the complete caption) gets one shared null + record. Its content-based owner is deterministic, so no cross-rank duplicate + is published. Other requests stay attached to their diagnostic sample key. + """ + + tokens = _special_tokens(dataset) + produced = 0 + if dataset.cfg_dropout_rate > 0 and not dataset.cfg_dropout_keep_metadata: + null_text_ids, _ = dataset._tokenize_caption("") + null_request = build_reasoner_feature_request_from_text_tokens( + sample_key=_NULL_PROMPT_KEY, + text_ids=null_text_ids, + special_tokens=tokens, + use_float_positions=use_float_positions, + identity=identity, + ) + if _request_owner(null_request, world_size) == rank: + yield null_request + produced += 1 + if max_documents_per_rank is not None and produced >= max_documents_per_rank: + return + + rank_dataset = _rank_dataset(dataset, rank=rank, world_size=world_size) + for document in iter_sft_reasoner_documents(rank_dataset): + if document.caption == "" and dataset.cfg_dropout_rate > 0 and not dataset.cfg_dropout_keep_metadata: + continue + yield build_reasoner_feature_request_from_text_tokens( + sample_key=document.sample_key, + text_ids=document.text_token_ids, + special_tokens=tokens, + use_float_positions=use_float_positions, + identity=identity, + ) + produced += 1 + if max_documents_per_rank is not None and produced >= max_documents_per_rank: + return + + +def _initialize_distributed(coordination_run_id: str | None = None) -> DistributedContext: + """Resolve torchrun worker identity without creating a process group. + + Extraction workers are intentionally independent: each owns a complete + Reasoner replica, a disjoint metadata slice, and rank-local staging files. + Avoiding NCCL collectives means one worker can fail without stranding peers + in a mismatched or hours-long collective; torchrun remains the fail-fast + process supervisor. + """ + if not torch.cuda.is_available(): + raise RuntimeError("Reasoner feature extraction requires CUDA") + topology_names = ("WORLD_SIZE", "RANK", "LOCAL_RANK", "LOCAL_WORLD_SIZE") + present_topology = {name for name in topology_names if name in os.environ} + if not present_topology: + world_size, rank, local_rank, local_world_size = 1, 0, 0, 1 + elif present_topology != set(topology_names): + missing = sorted(set(topology_names) - present_topology) + raise ValueError( + "Incomplete torchrun topology environment; either set none for a direct single-worker run " + f"or set all of {topology_names}. Missing: {missing}" + ) + else: + world_size = int(os.environ["WORLD_SIZE"]) + rank = int(os.environ["RANK"]) + local_rank = int(os.environ["LOCAL_RANK"]) + local_world_size = int(os.environ["LOCAL_WORLD_SIZE"]) + if world_size <= 0 or not 0 <= rank < world_size: + raise ValueError(f"Invalid torchrun rank topology: rank={rank}, world_size={world_size}") + if local_world_size <= 0 or not 0 <= local_rank < local_world_size: + raise ValueError( + f"Invalid local torchrun topology: local_rank={local_rank}, local_world_size={local_world_size}" + ) + if local_rank < 0 or local_rank >= torch.cuda.device_count(): + raise ValueError(f"LOCAL_RANK={local_rank} is outside the {torch.cuda.device_count()} visible CUDA devices") + torch.cuda.set_device(local_rank) + restart_count = os.environ.get("TORCHELASTIC_RESTART_COUNT", "0") + if coordination_run_id is not None: + # A restart is a new coordination attempt. Durable shards/completion + # remain resumable, while stale failure/status files from the previous + # attempt cannot poison the replacement workers. + run_id = f"{coordination_run_id}:{restart_count}" + else: + if world_size > local_world_size: + raise ValueError("Multi-node extraction requires an explicit --coordination-run-id shared by every node") + elastic_run_id = os.environ.get("TORCHELASTIC_RUN_ID", "manual") + # All workers launched by one local torchrun agent share a parent PID; + # adding it prevents stale status from a later default-rdzv invocation. + # A direct single-worker invocation uses its own PID for the same reason. + launch_pid = os.getpid() if world_size == 1 else os.getppid() + run_id = f"{elastic_run_id}:{restart_count}:{launch_pid}" + return DistributedContext(rank, world_size, local_rank, torch.device("cuda", local_rank), run_id) + + +def _extract_rank( + *, + reasoner: torch.nn.Module, + requests: Iterator[ReasonerFeatureRequest], + writer: IncrementalReasonerFeatureCacheWriter, + context: DistributedContext, +) -> ExtractionStats: + stats = ExtractionStats(rank=context.rank) + torch.cuda.synchronize(context.device) + started = time.perf_counter() + for request in requests: + stats.assigned_documents += 1 + if writer.contains(request.sample_key, request.fingerprint): + stats.cache_hits += 1 + continue + features = extract_reasoner_feature_batch(reasoner, [request]) + stats.extracted_tokens += request.token_ids.numel() + if not writer.append(ReasonerFeatureCacheEntry(request.sample_key, features)): + raise RuntimeError(f"Writer unexpectedly rejected newly extracted record {request.sample_key!r}") + stats.written_records += 1 + del features + # Commit resumable data first. The caller writes the current-run stats and + # only then creates rank.complete as the last successful worker action. + writer.flush() + torch.cuda.synchronize(context.device) + stats.extraction_seconds = time.perf_counter() - started + stats.peak_allocated_bytes = torch.cuda.max_memory_allocated(context.device) + stats.peak_reserved_bytes = torch.cuda.max_memory_reserved(context.device) + return stats + + +def _atomic_write_json(path: Path, payload: Mapping[str, Any]) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + descriptor, temporary_name = tempfile.mkstemp(prefix=f".{path.name}.tmp-", dir=path.parent) + try: + with os.fdopen(descriptor, "w", encoding="utf-8") as handle: + json.dump(payload, handle, indent=2, sort_keys=True) + handle.write("\n") + handle.flush() + os.fsync(handle.fileno()) + os.replace(temporary_name, path) + except BaseException: + try: + os.close(descriptor) + except OSError: + pass + try: + os.unlink(temporary_name) + except FileNotFoundError: + pass + raise + + +def _rank_status_path(staging_root: Path, rank: int, name: str) -> Path: + return staging_root / f"rank-{rank:05d}" / name + + +def _write_rank_stats( + writer: IncrementalReasonerFeatureCacheWriter, + stats: ExtractionStats, + *, + context: DistributedContext, + publishable: bool, +) -> None: + _atomic_write_json( + writer.rank_dir / _RANK_STATS_FILE, + { + "run_id": context.run_id, + "world_size": context.world_size, + "publishable": publishable, + "stats": asdict(stats), + }, + ) + + +def _commit_rank_result( + writer: IncrementalReasonerFeatureCacheWriter, + stats: ExtractionStats, + *, + context: DistributedContext, + publishable: bool, +) -> None: + """Atomically record stats, then optionally make the rank publishable.""" + _write_rank_stats(writer, stats, context=context, publishable=publishable) + if publishable: + # rank.complete is deliberately the final fallible worker commit. Rank + # zero waits for both this marker and the current-run stats record. + writer.finalize() + + +def _read_current_rank_stats(path: Path, *, context: DistributedContext) -> dict[str, Any] | None: + try: + payload = json.loads(path.read_text(encoding="utf-8")) + except FileNotFoundError: + return None + if not isinstance(payload, dict) or payload.get("run_id") != context.run_id: + return None + if payload.get("world_size") != context.world_size or payload.get("publishable") is not True: + return None + stats = payload.get("stats") + if not isinstance(stats, dict): + raise ValueError(f"Invalid extraction stats payload: {path}") + return stats + + +def _wait_for_rank_completion( + staging_root: Path, + *, + context: DistributedContext, + timeout_seconds: float, +) -> list[dict[str, Any]]: + """Wait through shared POSIX markers, never through a GPU collective.""" + deadline = time.monotonic() + timeout_seconds + while True: + failures: list[tuple[int, str]] = [] + stats: list[dict[str, Any]] = [] + missing: list[int] = [] + for rank in range(context.world_size): + failure_path = _rank_status_path(staging_root, rank, _RANK_FAILURE_FILE) + try: + failure = json.loads(failure_path.read_text(encoding="utf-8")) + except FileNotFoundError: + failure = None + if isinstance(failure, dict) and failure.get("run_id") == context.run_id: + failures.append((rank, str(failure.get("traceback", "unknown worker failure")))) + continue + + completion_path = _rank_status_path(staging_root, rank, "rank.complete.json") + rank_stats = _read_current_rank_stats( + _rank_status_path(staging_root, rank, _RANK_STATS_FILE), + context=context, + ) + if not completion_path.is_file() or rank_stats is None: + missing.append(rank) + else: + stats.append(rank_stats) + + if failures: + details = "\n".join(f"--- rank {rank} ---\n{error}" for rank, error in failures) + raise RuntimeError(f"Reasoner extraction worker failed:\n{details}") + if not missing: + return sorted(stats, key=lambda item: int(item["rank"])) + if time.monotonic() >= deadline: + raise TimeoutError( + "Timed out waiting for Reasoner extraction ranks to finish; " + f"missing current-run completion/stats from ranks {missing}" + ) + time.sleep(min(1.0, max(0.0, deadline - time.monotonic()))) + + +def _write_worker_failure(args: argparse.Namespace, context: DistributedContext, error: str) -> None: + output = Path(args.output) + staging_root = output.parent / f".{output.name}.reasoner-kv-staging" + _atomic_write_json( + _rank_status_path(staging_root, context.rank, _RANK_FAILURE_FILE), + {"rank": context.rank, "run_id": context.run_id, "traceback": error}, + ) + + +def _run(args: argparse.Namespace, context: DistributedContext) -> Path: + config = load_experiment_from_toml(args.sft_toml, extra_overrides=args.overrides) + _validate_extraction_config(config) + + dtype = torch.bfloat16 + identity = ReasonerFeatureCacheIdentity( + reasoner=args.reasoner_fingerprint, + tokenizer=args.tokenizer_fingerprint, + framing=args.framing_fingerprint, + ) + output = Path(args.output) + if output.exists(): + provider = OfflineReasonerFeatureProvider( + output, + expected_identity=identity, + expected_dtype=dtype, + verify_checksums=args.verify_resume_checksums, + ) + if context.rank == 0: + provider.verify_all_shards() + manifest = output / "manifest.json" + if context.rank == 0: + print(f"Reasoner feature cache is already complete: {manifest}", flush=True) + return manifest + + dataset = _instantiate_sft_dataset(config) + writer = IncrementalReasonerFeatureCacheWriter.resume( + output, + identity=identity, + rank=context.rank, + world_size=context.world_size, + max_shard_bytes=args.max_shard_bytes, + verify_checksums=args.verify_resume_checksums, + ) + use_float_positions = bool(config.model.config.diffusion_expert_config.enable_fps_modulation) + requests = _iter_rank_requests( + dataset, + identity=identity, + rank=context.rank, + world_size=context.world_size, + use_float_positions=use_float_positions, + max_documents_per_rank=args.max_documents_per_rank, + ) + + torch.cuda.reset_peak_memory_stats(context.device) + torch.cuda.synchronize(context.device) + reasoner_construction_started = time.perf_counter() + reasoner = _build_reasoner(config, device=context.device, dtype=dtype) + torch.cuda.synchronize(context.device) + reasoner_construction_seconds = time.perf_counter() - reasoner_construction_started + torch.cuda.synchronize(context.device) + checkpoint_load_started = time.perf_counter() + load_reasoner_only_dcp(reasoner, args.checkpoint, source=args.checkpoint_source) + torch.cuda.synchronize(context.device) + checkpoint_load_seconds = time.perf_counter() - checkpoint_load_started + publishable = args.max_documents_per_rank is None + stats = _extract_rank( + reasoner=reasoner, + requests=requests, + writer=writer, + context=context, + ) + stats.reasoner_construction_seconds = reasoner_construction_seconds + stats.checkpoint_load_seconds = checkpoint_load_seconds + _commit_rank_result(writer, stats, context=context, publishable=publishable) + + if not publishable: + summary = { + "published": False, + "reason": ("INCOMPLETE / NOT FOR TRAINING: --max-documents-per-rank is a resumable smoke run"), + "staging_root": str(writer.staging_root), + "rank": asdict(stats), + } + print(json.dumps(summary, indent=2, sort_keys=True), flush=True) + return writer.staging_root + + if context.rank != 0: + return output / "manifest.json" + + all_stats = _wait_for_rank_completion( + writer.staging_root, + context=context, + timeout_seconds=args.coordination_timeout_seconds, + ) + + manifest = finalize_incremental_reasoner_feature_cache( + output, + identity=identity, + world_size=context.world_size, + verify_checksums=args.verify_resume_checksums, + ) + summary = { + "manifest": str(manifest), + "checkpoint": str(args.checkpoint), + "checkpoint_source": args.checkpoint_source, + "identity": asdict(identity), + "world_size": context.world_size, + "ranks": all_stats, + } + print(json.dumps(summary, indent=2, sort_keys=True), flush=True) + return manifest + + +def _parse_args(argv: Sequence[str] | None = None) -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) + parser.add_argument("--sft-toml", required=True) + parser.add_argument("--checkpoint", required=True, help="Local DCP model component or iteration directory") + parser.add_argument("--checkpoint-source", choices=("regular", "ema"), default="regular") + parser.add_argument("--output", required=True, help="New immutable cache root on shared POSIX storage") + parser.add_argument("--reasoner-fingerprint", required=True) + parser.add_argument("--tokenizer-fingerprint", required=True) + parser.add_argument("--framing-fingerprint", required=True) + parser.add_argument("--dtype", choices=("bfloat16",), default="bfloat16") + parser.add_argument("--max-shard-bytes", type=int, default=512 * 1024**2) + parser.add_argument( + "--max-documents-per-rank", + type=int, + default=None, + help="Resumable smoke-test limit; flushes staging shards but never publishes a cache", + ) + parser.add_argument( + "--coordination-timeout-seconds", + type=float, + default=24 * 60 * 60, + help="Rank-0 timeout while polling shared-POSIX completion markers", + ) + parser.add_argument( + "--coordination-run-id", + default=None, + help="Unique shared job ID; required for multi-node extraction", + ) + parser.add_argument( + "--verify-resume-checksums", + action=argparse.BooleanOptionalAction, + default=True, + ) + parser.add_argument( + "overrides", + nargs=argparse.REMAINDER, + help="Hydra overrides applied after TOML; prefix the list with --", + ) + args = parser.parse_args(argv) + args.overrides = [item for item in args.overrides if item != "--"] + if args.max_shard_bytes <= 0: + parser.error("--max-shard-bytes must be positive") + if args.max_documents_per_rank is not None and args.max_documents_per_rank <= 0: + parser.error("--max-documents-per-rank must be positive") + if args.coordination_timeout_seconds <= 0: + parser.error("--coordination-timeout-seconds must be positive") + return args + + +def main(argv: Sequence[str] | None = None) -> int: + args = _parse_args(argv) + context = _initialize_distributed(args.coordination_run_id) + try: + _run(args, context) + return 0 + except BaseException: + error = traceback.format_exc() + try: + _write_worker_failure(args, context, error) + except BaseException: + pass + raise + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/cosmos_framework/scripts/extract_reasoner_features_test.py b/cosmos_framework/scripts/extract_reasoner_features_test.py new file mode 100644 index 000000000..2cd941ea9 --- /dev/null +++ b/cosmos_framework/scripts/extract_reasoner_features_test.py @@ -0,0 +1,267 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +from __future__ import annotations + +import json +from pathlib import Path +from types import SimpleNamespace + +import pytest +import torch + +import cosmos_framework.scripts.extract_reasoner_features as extraction_cli +from cosmos_framework.data.generator.local_datasets.sft_dataset import get_sft_dataset +from cosmos_framework.model.generator.reasoner_feature_cache import ReasonerFeatureCacheIdentity +from cosmos_framework.model.generator.reasoner_features import ReasonerFeatureRequest +from cosmos_framework.scripts.extract_reasoner_features import ( + DistributedContext, + ExtractionStats, + _atomic_write_json, + _commit_rank_result, + _find_sft_dataset_config, + _initialize_distributed, + _iter_rank_requests, + _parse_args, + _prepare_reasoner_model_config, + _request_owner, + _validate_extraction_config, + _wait_for_rank_completion, +) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_find_sft_dataset_config_requires_exactly_one_lazy_node() -> None: + node = {"_target_": get_sft_dataset, "jsonl_paths": ["data.jsonl"]} + assert _find_sft_dataset_config({"outer": {"dataset": node}}) is node + + with pytest.raises(ValueError, match="found 0"): + _find_sft_dataset_config({"outer": {"dataset": {}}}) + with pytest.raises(ValueError, match="found 2"): + _find_sft_dataset_config({"first": node, "second": dict(node)}) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_request_owner_is_stable_and_content_sensitive() -> None: + request = ReasonerFeatureRequest( + sample_key="sample", + token_ids=torch.tensor([1, 2]), + position_ids=torch.tensor([0, 1]), + causal_offsets=torch.tensor([0, 2]), + fingerprint="fingerprint-a", + ) + assert _request_owner(request, 8) == _request_owner(request, 8) + other = ReasonerFeatureRequest( + sample_key=request.sample_key, + token_ids=request.token_ids, + position_ids=request.position_ids, + causal_offsets=request.causal_offsets, + fingerprint="fingerprint-b", + ) + assert any(_request_owner(request, size) != _request_owner(other, size) for size in range(2, 17)) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_rank_requests_frame_tokens_and_deduplicate_standard_cfg_null(monkeypatch: pytest.MonkeyPatch) -> None: + dataset = SimpleNamespace( + cfg_dropout_rate=0.1, + cfg_dropout_keep_metadata=False, + metadata=[{"uuid": "video"}], + _tokenize_caption=lambda caption: ([7] if caption else [8], caption), + ) + documents = [ + SimpleNamespace(sample_key="video_w0", caption="", text_token_ids=(8,)), + SimpleNamespace(sample_key="video_w0", caption="caption", text_token_ids=(7,)), + ] + monkeypatch.setattr( + extraction_cli, "_special_tokens", lambda _dataset: {"eos_token_id": 9, "start_of_generation": 10} + ) + monkeypatch.setattr(extraction_cli, "iter_sft_reasoner_documents", lambda _dataset: iter(documents)) + + requests = list( + _iter_rank_requests( + dataset, # type: ignore[arg-type] + identity=ReasonerFeatureCacheIdentity("reasoner", "tokenizer", "framing"), + rank=0, + world_size=1, + use_float_positions=True, + max_documents_per_rank=None, + ) + ) + + assert [request.sample_key for request in requests] == ["__null__", "video_w0"] + assert requests[0].token_ids.tolist() == [8, 9, 10] + assert requests[1].token_ids.tolist() == [7, 9, 10] + assert all(request.position_ids.dtype == torch.float32 for request in requests) + + +def _config() -> SimpleNamespace: + model_instance = { + "_target_": "cosmos_framework.model.generator.mot.unified_mot.Qwen3VLTextForCausalLM", + "config": { + "_target_": "cosmos_framework.configs.base.defaults.reasoner.create_vlm_config", + }, + } + model = SimpleNamespace( + joint_attn_implementation="two_way", + parallelism=SimpleNamespace(context_parallel_shard_degree=1), + video_temporal_causal=False, + causal_training_strategy="none", + diffusion_expert_config=SimpleNamespace(vision_temporal_position_mode="latent_index"), + vlm_config=SimpleNamespace(model_instance=model_instance), + ) + return SimpleNamespace(model=SimpleNamespace(config=model)) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_prepare_reasoner_model_config_disables_generator_and_visual() -> None: + prepared = _prepare_reasoner_model_config(_config()) + assert prepared["config"]["include_gen_pathway"] is False + assert prepared["config"]["include_und_pathway"] is True + assert prepared["config"]["include_visual"] is False + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_validate_extraction_config_fails_closed_on_unsupported_layouts() -> None: + config = _config() + _validate_extraction_config(config) + + config.model.config.parallelism.context_parallel_shard_degree = 2 + with pytest.raises(ValueError, match="context_parallel"): + _validate_extraction_config(config) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_parse_args_validates_bounded_writer_controls() -> None: + base = [ + "--sft-toml", + "recipe.toml", + "--checkpoint", + "checkpoint", + "--output", + "cache", + "--reasoner-fingerprint", + "r", + "--tokenizer-fingerprint", + "t", + "--framing-fingerprint", + "f", + ] + args = _parse_args([*base, "--max-documents-per-rank", "2", "--", "optimizer.lr=1e-5"]) + assert args.max_documents_per_rank == 2 + assert args.max_shard_bytes == 512 * 1024**2 + assert args.overrides == ["optimizer.lr=1e-5"] + + with pytest.raises(SystemExit): + _parse_args([*base, "--max-shard-bytes", "0"]) + with pytest.raises(SystemExit): + _parse_args([*base, "--coordination-timeout-seconds", "0"]) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_torchrun_context_uses_global_rank_without_initializing_collectives( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("RANK", "5") + monkeypatch.setenv("WORLD_SIZE", "8") + monkeypatch.setenv("LOCAL_RANK", "2") + monkeypatch.setenv("LOCAL_WORLD_SIZE", "4") + monkeypatch.setenv("TORCHELASTIC_RUN_ID", "job") + monkeypatch.setenv("TORCHELASTIC_RESTART_COUNT", "3") + selected: list[int] = [] + monkeypatch.setattr(torch.cuda, "is_available", lambda: True) + monkeypatch.setattr(torch.cuda, "device_count", lambda: 4) + monkeypatch.setattr(torch.cuda, "set_device", selected.append) + + context = _initialize_distributed("explicit-run") + + assert (context.rank, context.world_size, context.local_rank) == (5, 8, 2) + assert context.device == torch.device("cuda", 2) + assert context.run_id == "explicit-run:3" + assert selected == [2] + + monkeypatch.delenv("RANK") + with pytest.raises(ValueError, match="Incomplete torchrun topology"): + _initialize_distributed("explicit-run") + + monkeypatch.setenv("RANK", "5") + with pytest.raises(ValueError, match="Multi-node extraction requires"): + _initialize_distributed() + + +class _FakeWriter: + def __init__(self, rank_dir: Path) -> None: + self.rank_dir = rank_dir + self.finalized = False + + def finalize(self) -> Path: + stats = json.loads((self.rank_dir / "extraction.stats.json").read_text(encoding="utf-8")) + assert stats["publishable"] is True + marker = self.rank_dir / "rank.complete.json" + marker.write_text("{}", encoding="utf-8") + self.finalized = True + return marker + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_smoke_result_never_writes_completion_but_full_result_commits_stats_first(tmp_path: Path) -> None: + context = DistributedContext(0, 1, 0, torch.device("cpu"), "run:0") + smoke_writer = _FakeWriter(tmp_path / "smoke" / "rank-00000") + _commit_rank_result( # type: ignore[arg-type] + smoke_writer, + ExtractionStats(rank=0, written_records=2), + context=context, + publishable=False, + ) + assert not smoke_writer.finalized + smoke_payload = json.loads((smoke_writer.rank_dir / "extraction.stats.json").read_text(encoding="utf-8")) + assert smoke_payload["publishable"] is False + assert not (smoke_writer.rank_dir / "rank.complete.json").exists() + + full_writer = _FakeWriter(tmp_path / "full" / "rank-00000") + _commit_rank_result( # type: ignore[arg-type] + full_writer, + ExtractionStats(rank=0, written_records=3), + context=context, + publishable=True, + ) + assert full_writer.finalized + assert (full_writer.rank_dir / "rank.complete.json").is_file() + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_rank_zero_coordinates_through_current_run_files_and_surfaces_failures(tmp_path: Path) -> None: + context = DistributedContext(0, 2, 0, torch.device("cpu"), "run:0") + for rank in range(2): + rank_dir = tmp_path / f"rank-{rank:05d}" + _atomic_write_json( + rank_dir / "extraction.stats.json", + { + "run_id": context.run_id, + "world_size": 2, + "publishable": True, + "stats": {"rank": rank}, + }, + ) + _atomic_write_json(rank_dir / "rank.complete.json", {}) + + assert _wait_for_rank_completion(tmp_path, context=context, timeout_seconds=0.01) == [ + {"rank": 0}, + {"rank": 1}, + ] + + _atomic_write_json( + tmp_path / "rank-00001" / "extraction.failure.json", + {"run_id": context.run_id, "traceback": "rank one failed"}, + ) + with pytest.raises(RuntimeError, match="rank one failed"): + _wait_for_rank_completion(tmp_path, context=context, timeout_seconds=0.01) diff --git a/docs/nano_sft_decoupled_reasoner.md b/docs/nano_sft_decoupled_reasoner.md new file mode 100644 index 000000000..a93794b71 --- /dev/null +++ b/docs/nano_sft_decoupled_reasoner.md @@ -0,0 +1,670 @@ +# Decoupling the frozen Reasoner from Nano generator SFT + +Status: Phase 1 and the Phase 2 offline extraction MVP are implemented; an +eight-H20 short-run FSDP/EMA A/B is complete, while full-corpus and +compile-enabled production validation remain pending + +Target recipe: `vision_sft_nano` and later Nano multiview SFT variants + +## Current implementation status + +Implemented in this branch: + +- a provider-neutral per-layer Reasoner K/V contract, exact packed-sample + isolation, inline capture, and static generator-only replay; +- an UND-only extraction primitive, `extract_reasoner_feature_batch`, for a + supplied set of finalized `ReasonerFeatureRequest` objects; +- structural `prune_und_pathway_` before materialization/FSDP, including the + regular and EMA model construction paths; +- a generator-only VFM text-layout path that does not require `embed_tokens`; +- an immutable, layer-major safetensors cache writer, manifest, checksums, + per-record fingerprints, and `OfflineReasonerFeatureProvider`; +- a bounded-memory, resumable rank-local shard writer plus shared-POSIX + distributed finalizer for production cache extraction; +- a finite deterministic SFT document enumerator that reuses the training + caption/window framing and expands every reachable caption/CFG variant; +- a strict Reasoner-only DCP loader for either `net.language_model.*` or + `net_ema.language_model.*`, without constructing Generator, VAE, or EMA; +- a shared-POSIX `torchrun` extraction CLI with independent rank-local workers, + bounded file-based coordination, shared CFG-null deduplication, and atomic + publication; +- strict external-backend fingerprints plus cache/model layer, KV-head, and + head-dimension validation before FSDP materialization; +- offline-provider submission before noising and synchronized failure handling + before the FSDP forward; and +- CPU contract/storage tests plus a tiny H20 CUDA test comparing all generator + outputs and gradients between the full cached model and the structurally + pruned model. + +Not implemented yet: + +- remote and read-through clients/services (the backend names and `endpoint` + field are reserved, but selecting either backend currently fails explicitly); +- layerwise H2D staging (`layerwise_h2d=true` is rejected); +- automatic content-digest derivation for Reasoner/tokenizer/dataset artifacts; + the CLI currently requires the three pinned fingerprints explicitly; and +- a real full-corpus extraction plus a compile-enabled, representative-duration + Cosmos3-Nano memory and throughput benchmark. A short eager-mode eight-H20 A/B + on the official eight-video sample is complete. + +## Decision + +Use a single external-conditioning boundary at the per-layer Reasoner +cross-attention K/V tensors. The configuration contract names these backends: + +1. `joint` preserves the original one-pass dual-pathway forward and is the default. +2. `inline` captures Reasoner K/V locally, then runs a second cached GEN-only pass; + it is the numerical-reference path. +3. `offline` reads precomputed immutable Reasoner K/V shards and is implemented. +4. `remote` will obtain the same tensors asynchronously from dedicated Reasoner workers. +5. `read_through` will check the offline cache first and send misses to remote workers. + +Only `joint`, `inline`, and `offline` are executable today. `remote` and +`read_through` intentionally raise `NotImplementedError` during provider creation. + +For a fixed SFT corpus, `offline` should be the default. It removes the Reasoner from +every training rank, is deterministic, and turns Reasoner work into a one-time dataset +preparation cost. `remote` is useful for very large corpora, changing prompts, or +augmentations that make an exhaustive cache too large. The hybrid backend is the +long-term operational recommendation, but it should be built only after the offline +path establishes numerical parity. + +Do **not** cache QKV or every layer's hidden state. The generator computes Q from its +own trainable, noised tokens. It only consumes the frozen Reasoner's K and V at each +layer. For Nano's grouped-query attention, Q is also four times wider than K or V, so +caching QKV would triple the storage of the minimal K/V boundary. + +## Why this is the correct boundary + +Nano uses a dual-pathway MoT decoder. At every layer: + +- UND/Reasoner tokens use `q_proj`, `k_proj`, `v_proj`, `o_proj`, and `mlp`. +- GEN tokens use the corresponding `*_moe_gen` modules. +- GEN queries attend to both UND K/V and live GEN K/V. +- UND attention is causal and independent of GEN tokens, so all UND K/V can be + produced without running the generator. + +The current `vision_sft_nano` optimizer selects only `moe_gen`, `time_embedder`, +`vae2llm`, and `llm2vae`. The optimizer factory marks every parameter outside that +allowlist as `requires_grad=False`. The Reasoner is therefore frozen, but the current +joint forward still constructs, shards, all-gathers, and executes it. + +There is already a close precedent in `inference_text_kv_memory.py`: it stores +RoPE-applied UND K/V per layer, reports `is_gen_only()` after all layers are populated, +and then executes only the GEN pathway. The implemented training contract generalizes +that pattern without introducing a second attention abstraction. + +### Canonical tensor contract + +For each causal text segment and each decoder layer `l`, persist: + +```text +cross_k[l]: [S_und, num_kv_heads, head_dim] +cross_v[l]: [S_und, num_kv_heads, head_dim] +``` + +`cross_k` is the exact key seen by GEN cross-attention: after the Reasoner K projection, +text K normalization, any generator-facing UND K normalization, and RoPE. `cross_v` is +the V projection output (V does not receive RoPE). For Nano, `cross_k` is the same K +used by UND self-attention. The name deliberately does not promise that equality for +other tiers; a future model can materialize a generator-specific normalized K. + +The batch-level feature object also carries the per-sample/per-caption offsets needed +to isolate packed samples. The existing `PackedSequence` remains the source of truth +for GEN layout, attention masks, view IDs, and current noisy/conditioning tokens. +The persistent representation is canonical and unpadded, with all K/V heads present; +batch padding and any context-parallel sharding are runtime concerns and must never be +written into a cache shard. + +The target feature identity must cover: + +- the ordered, fully framed token IDs, including EOS and start-of-generation; +- exact mRoPE position IDs and causal-document boundaries; +- Reasoner checkpoint content hash and model/config hash; +- tokenizer and special-token hash; +- K/V dtype or quantization recipe; and +- cache schema and producer-code versions. + +Today, `compute_reasoner_feature_fingerprint` hashes the exact token IDs, position IDs, +causal offsets, cache schema, and the configured Reasoner/tokenizer/framing identity. +The manifest separately pins dtype and tensor geometry. Until an explicit producer +revision field is added, deployments should include it in one of the three configured +identity strings. A mismatch is an error, never a warning or an implicit fallback. + +## Size and memory model + +Cosmos3-Nano has 36 layers, 8 K/V heads, and head dimension 128. In BF16: + +```text +bytes per UND token + = layers * (K + V) * kv_heads * head_dim * bytes_per_element + = 36 * 2 * 8 * 128 * 2 + = 147,456 bytes + = 144 KiB +``` + +| Framed UND tokens | BF16 K/V per example | +|---:|---:| +| 256 | 36 MiB | +| 512 | 72 MiB | +| 1,024 | 144 MiB | +| 1,790 | 251.7 MiB | +| 2,048 | 288 MiB | + +The local full BridgeData manifest has 1,222 examples. With the recipe's real Qwen +tokenization, it has 1,082 framed tokens on average (p50 1,004, p95 1,586, maximum +1,797), which gives about **181.6 GiB** of BF16 K/V. This is a reasonable offline-cache +MVP. At one million 1,024-token examples, however, the cache is about **137 TiB**, so a +remote or read-through backend becomes attractive. + +The Nano language model contains approximately: + +| Component | Parameters | BF16 logical size | +|---|---:|---:| +| GEN layer pathway | 6.946B | 12.94 GiB | +| UND layer pathway | 6.946B | 12.94 GiB | +| UND embeddings + LM head + final norm | 1.245B | 2.32 GiB | +| Total removable Reasoner | 8.191B | 15.26 GiB | + +A real Nano meta-model audit measured 15,136,811,008 parameters in the full +language model and 6,946,075,648 after pruning: 8,190,735,360 parameters +(54.111%) removed. The 397 remaining parameter tensors are exactly the original +`moe_gen` FQNs, and the pruned model completed meta initialization, activation +checkpoint wrapping, and block/root FSDP2 wrapping without a missing attribute. + +The standard recipe stores FSDP master parameters in FP32 and enables a second FP32 +EMA network. With eight-way FSDP, removing the 8.191B Reasoner parameters from both +networks saves about **7.63 GiB of resident parameter shards per GPU**, before counting +smaller layer all-gathers, activations, and allocator fragmentation. The short +eight-H20 run measured an average peak-allocated reduction of 15.630 GiB/GPU and an +average peak-reserved reduction of 19.419 GiB/GPU; representative full-corpus and +compile-enabled measurements remain required. + +## Common provider API + +The provider-neutral request and result live in +`cosmos_framework/model/generator/reasoner_features.py`: + +```python +@dataclass(frozen=True) +class ReasonerFeatureRequest: + sample_key: str + token_ids: torch.Tensor + position_ids: torch.Tensor + causal_offsets: torch.Tensor + fingerprint: str + +@dataclass(frozen=True) +class ReasonerFeatureBatch: + cross_k: tuple[torch.Tensor, ...] + cross_v: tuple[torch.Tensor, ...] + causal_offsets: torch.Tensor + fingerprints: tuple[str, ...] = () + +class ReasonerFeatureProvider(Protocol): + def submit( + self, requests: Sequence[ReasonerFeatureRequest] + ) -> Future[ReasonerFeatureBatch]: ... +``` + +`submit` is a future-shaped contract for every provider. The current offline provider +performs its local read synchronously and returns an already-completed future. The +training path submits after final sequence packing and before noising, then resolves +the future before entering the FSDP forward. A future asynchronous provider can use +the same seam to overlap provider work with noising or other preparation. + +`extract_reasoner_feature_batch(causal_lm, requests)` is the current UND-only +reference extractor. The immutable cache APIs are in +`cosmos_framework/model/generator/reasoner_feature_cache.py`: + +- `ReasonerFeatureCacheIdentity(reasoner, tokenizer, framing)`; +- `build_reasoner_feature_requests(packed_sequence, sample_keys, identity)`, the + shared training/extraction framing boundary; +- `build_reasoner_feature_request_from_text_tokens(...)`, which uses + `PackedSequenceBuilder.pack_text_tokens` for the supported single-caption + offline producer path; +- `ReasonerFeatureCacheEntry(sample_key, features)`; +- `write_reasoner_feature_cache(cache_root, entries, identity=...)`; and +- `IncrementalReasonerFeatureCacheWriter(...).append/flush/finalize()` plus + `finalize_incremental_reasoner_feature_cache(...)` for distributed extraction; and +- `OfflineReasonerFeatureProvider(cache_root, expected_identity=..., + expected_dtype=..., strict_fingerprint=True)`. + +The incremental writer keeps only one target-sized shard payload in CPU memory. Each +rank writes uniquely named safetensors into a hidden sibling staging directory. A +shard is committed by publishing its checksum-bearing JSON sidecar only after the +safetensors file has been fsynced and atomically renamed. Reopening the same rank with +`IncrementalReasonerFeatureCacheWriter.resume(...)` validates those commits and skips +already-seen `(sample_key, fingerprint)` pairs. Extraction loops can call +`contains(sample_key, fingerprint)` before running the Reasoner and `append(...)` +also returns `False` for a repeated pair. `finalize()` writes a rank-complete marker +binding all sidecar checksums. + +After every rank is complete, one coordinator calls +`finalize_incremental_reasoner_feature_cache(cache_root, identity=..., world_size=...)`. +The finalizer requires all expected rank markers, validates cache identity, tensor +geometry, dtype, checksums, and cross-rank duplicate identities, then atomically +publishes the ordinary immutable `manifest.json` layout consumed by +`OfflineReasonerFeatureProvider`. The current implementation assumes a shared POSIX +filesystem. The low-level writer deliberately does not own dataset enumeration, +process launch, or a remote/object-store writer; the implemented extraction CLI +supplies enumeration and process coordination for a shared-POSIX deployment. Staging +data is retained after publication so finalization is idempotent and can be audited +or retried. + +The actual offline TOML configuration is a single flat table: + +```toml +[model.reasoner_conditioning] +backend = "offline" # joint | inline | offline | remote | read_through +cache_root = "/path/to/reasoner-kv-cache" +reasoner_fingerprint = "" +tokenizer_fingerprint = "" +framing_fingerprint = "" +strict_fingerprint = true +request_timeout_s = 300.0 +``` + +External backends deliberately require `strict_fingerprint = true`. A sample +key is only a diagnostic alias: the content fingerprint over exact token IDs, +positions, offsets, and cache identity is the authoritative lookup key. This +prevents a same-key, same-length caption variant from silently reusing the +wrong K/V record. + +With `strict_fingerprint=true` (the default), all three fingerprint fields are +required for an external backend. `request_timeout_s` bounds the +`Future.result()` wait; the current offline provider performs its read inside +`submit()`, so this deadline becomes operationally useful once submission is truly +asynchronous. The schema also currently accepts +`prefetch_batches`, `endpoint`, and `layerwise_h2d`; prefetch-depth scheduling is +not wired yet, `endpoint` is reserved for `remote`/`read_through`, and +`layerwise_h2d` must remain `false`. + +External modes initially require: + +- frozen UND weights and `predict_text_tokens=False`; +- Qwen3-VL-8B Nano dense layers; +- the base `OmniMoTModel` (`OmniMoTCausalModel` fails closed until its + AR/teacher-forcing memory dispatcher can compose with external K/V); +- `joint_attn_implementation="two_way"`; +- no context parallelism; and +- an initialized generator checkpoint (fresh initialization cannot copy GEN weights + from an UND tower that is absent). + +Multiview attention, context parallelism, teacher forcing, and other model tiers should +be enabled only after their exact cached-attention parity tests exist. + +## Generator-only structural model + +Skipping UND execution is not enough: the user's goal requires that training ranks do +not materialize UND parameters at all. + +The implemented external-backend path instantiates the existing model on the meta +device, calls `prune_und_pathway_` while the removed modules still occupy no storage, +and only then applies compile/FSDP and materializes the model. Every GEN module keeps +its existing fully-qualified name. The pruning operation removes: + +- `embed_tokens`, `lm_head`, and the UND final norm; +- the optional Reasoner-side `visual` encoder when configured; +- per-layer UND Q/K/V/O projections and Q/K norms; +- per-layer UND MLP and UND pre-attention/post-attention norms. + +It keeps `rotary_emb`, every `*_moe_gen` module, and VFM-level encoders/decoders. The +VFM input construction allocates the packed hidden-state buffer without calling +`embed_tokens`; the text rows are layout placeholders while GEN rows are filled by +`vae2llm`, the timestep embedder, and other live modality projectors. + +The transformer loop now has an explicit generator-only branch: + +1. `StaticReasonerKVMemoryState` is initialized from a `ReasonerFeatureBatch`. +2. `is_gen_only()` is true before layer 0. +3. Each decoder layer reads its external UND K/V and executes only GEN norms, + projections, attention, MLP, and residuals. +4. The final output applies only `norm_moe_gen`; it never touches an UND tensor or + parameter. + +Both regular and EMA networks go through the same `build_net` path and therefore have +the identical pruned structure. This prevents the frozen Reasoner from being +materialized, FSDP-sharded, all-gathered, or duplicated in EMA for external backends. +The tiny CUDA parity test additionally asserts that the pruned state dict exposes only +GEN pathway parameters and that all generator gradients match the full cached model. + +### Checkpoint behavior + +Implemented startup behavior skips pretrained Reasoner loading for an external +backend. If `load_weights_from_pretrained=true`, an external run must supply a resume +or warm-start checkpoint because the deleted UND tower cannot seed the GEN weights; +explicit random GEN initialization remains possible by disabling that option. + +For a model-only warm start whose DCP configuration skips the complete `net_ema.*` +subtree, the callback explicitly reports that fact and initializes the pruned EMA +model from the newly loaded regular generator. If EMA was loaded in full, it is +preserved. A partial EMA skip is rejected because it would mix restored and random +EMA leaves. Same-job resume always preserves its restored EMA trajectory. A +non-strict warm start with EMA and no complete EMA skip is also rejected because it +cannot prove whether EMA was loaded or left randomly initialized. + +The shipped full Nano DCP was audited against the pruned vision-SFT target: all +405 target tensors (397 language-model GEN tensors plus 8 VFM tensors) exist with +matching shapes, and strict DCP subset loading succeeds while ignoring source-only +UND tensors. This validates the model-only warm-start shape contract; a distributed +generator-only save/optimizer/resume integration test is still pending. + +Still to validate before production use: + +- full-checkpoint to generator-only warm start for both DCP and safetensors; +- generator-only save/resume while preserving immutable cache identity; +- reconstruction of a full inference model from the generator checkpoint and pinned + Reasoner; and +- mismatched generator-checkpoint/cache provenance rejection on resume. + +## Offline backend + +### Extraction + +The in-process `extract_reasoner_feature_batch` primitive is implemented and runs +each finalized causal document independently under `torch.inference_mode()`. It uses +the existing `ReasonerKVCache` to return normalized, RoPE-applied K/V. It currently +rejects architectures that require a second generator-specific UND K normalization. + +A resumable distributed CLI is implemented. Its current interface is: + +```shell +torchrun --nproc-per-node=8 -m cosmos_framework.scripts.extract_reasoner_features \ + --sft-toml examples/toml/sft_config/vision_sft_nano.toml \ + --checkpoint /shared/checkpoints/iter_000000100 \ + --checkpoint-source regular \ + --output /shared/caches/nano-sft-reasoner-kv \ + --reasoner-fingerprint \ + --tokenizer-fingerprint \ + --framing-fingerprint +``` + +Each rank constructs only the Reasoner (`include_gen_pathway=false`, +`include_visual=false`), loads the explicitly selected regular or EMA subtree from a +local DCP, and processes a deterministic video-level metadata slice. The finite +enumerator probes and decodes each assigned video through the same FFmpeg geometry as +training, then expands all retained windows, reachable positive-weight captions, and +CFG variants. It never enters the infinite/random `SFTDataset.__iter__` path. Positive +FPS noise and random fixed-length frame selection create unbounded or non-canonical +variants and therefore fail closed. + +The shipped Nano DCP keeps the Generator and optional visual encoder below the same +language-model root. The Reasoner-only loader requires every one of the 399 text +Reasoner target leaves, while safely omitting the 397 source-only Generator leaves +and 351 source-only visual leaves. Any other source extra remains an error. Every +worker performs a complete independent `no_dist=True` DCP read; the extraction CLI +does not initialize a process group or issue GPU collectives. + +The CLI frames each tokenizer output through the same +`PackedSequenceBuilder.pack_text_tokens` boundary as training, including EOS, +start-of-generation, and float mRoPE positions when FPS modulation is active. It +checks `writer.contains(...)` before Reasoner execution, reports construction/load/ +extraction time and peak allocated/reserved GPU memory per rank, and lets rank zero +publish only after every rank's current-run stats and completion marker exist. Rank +zero polls shared-POSIX status files with a bounded timeout, then invokes the strict +finalizer to validate every marker, sidecar, checksum, identity, tensor signature, +and cross-rank duplicate before publication. Worker exceptions are recorded and +allowed to escape so `torchrun` can terminate peers immediately; there is no +long-running collective in which a failed rank can strand the others. + +`--max-documents-per-rank` is a resumable smoke mode: it flushes committed shards but +never writes a rank-complete marker or publishes `manifest.json`, and its output is +therefore not valid for training. A later unlimited run resumes those shards. The +first version intentionally runs the single-document reference extractor; +token-budget dynamic batching is pending. One live extraction job must own a given +output root; concurrent jobs targeting the same root are not supported. Single-node +`torchrun` derives a per-launch coordination ID automatically; multi-node launches +must pass the same unique `--coordination-run-id` on every node. The CLI appends +TorchElastic's restart count so a restarted attempt ignores stale failure/status +records while still resuming durable shards. + +The unit of caching should be a causal document, not necessarily an entire video: + +- normal SFT has one document per sample; +- per-view caption SFT has independent documents that can be concatenated using the + recorded offsets; and +- unconditional CFG uses one shared empty-caption entry instead of duplicating it for + every sample. + +### On-disk format + +The implemented writer does not create one filesystem object per sample. It groups +records into immutable safetensors shards plus a JSON manifest. The low-level writer +defaults to a 4 GiB target; the extraction CLI uses 512 MiB per rank by default to +bound aggregate host memory on an eight-rank job. The current safetensors keys are: + +```text +cross_k.000: [total_tokens_in_shard, H_kv, D] +cross_v.000: [total_tokens_in_shard, H_kv, D] +... +cross_k.035 +cross_v.035 +record_offsets: [num_records + 1] +``` + +The cache root must not already exist. The writer builds a sibling temporary +directory, writes every shard and the manifest, and publishes the complete root with +one atomic rename. The manifest records cache identity, tensor geometry, dtype, +per-record `(sample_key, fingerprint, token range)`, and a SHA-256 per shard. The +offline provider validates the manifest eagerly, verifies each shard checksum before +its first read, and slices only requested records with `safe_open`. The extraction +CLI additionally performs a streaming existence/checksum/metadata/shape audit of all +shards on rank zero when it encounters an already-published output. + +Layer-major storage supports future double-buffered, layerwise H2D staging. The +current implementation stages the whole requested `ReasonerFeatureBatch`; even a +1,024-token example is 144 MiB, far smaller than the removed Reasoner weights. + +For multiview/per-view text, each record must additionally preserve caption offsets, +stable camera/view IDs, and the shared mRoPE temporal origin. Per-view captions are +independent causal documents, while the GEN temporal origin follows the longest +caption rather than the sum of all caption lengths. The existing inference cache does +not encode this geometry, so it must not be advertised as multiview-training support. + +Start with BF16 for parity. FP8/INT8 K/V with per-layer or per-head scales can halve +storage and bandwidth, but is an explicit approximation mode with its own quality and +gradient-parity gate. + +### Dataset nondeterminism + +Cache the exact text feature, not only a video UUID. For the current SFT dataset: + +- the random window index is already represented by `uuid_w`; +- all possible finite windows must be enumerated during extraction; +- CFG dropout selects either the normal feature or the single null-caption feature; +- random T2V/I2V/V2V conditioning does not change Reasoner text K/V; and +- any stochastic caption/FPS/resolution transform that changes token IDs must either + be frozen, represented in the cache key, or handled by `remote/read_through`. + +## Remote backend + +This section is a design target. No remote transport, client, service, or +read-through write path is implemented in the current branch. + +Run Reasoner workers as a separate service allocation, not as ranks in the generator's +FSDP/DDP process group. Mixing service ranks into the training world size would make +generator collectives hang and would waste Reasoner ranks in every all-reduce. + +Recommended data flow: + +```text +generator rank -- submit(token IDs, positions, fingerprint) --> request queue + | reasoner replica + +-- VAE encode + noise + pack (overlap) | + |<----------- per-layer K/V or cache URI -----------------+ + +-- validate fingerprint --> GEN-only forward/backward +``` + +Operational requirements: + +- replicate the 8B Reasoner one per service GPU; prefer replica/data parallelism over + tensor parallelism because Nano fits on one modern accelerator; +- dynamic-batch requests by total UND tokens, not request count; +- use immutable request IDs, bounded queues, backpressure, deadlines, and idempotent + retries; +- keep a memory and/or NVMe content-addressed cache in front of Reasoner execution; +- expose queue time, Reasoner compute time, serialization time, bytes sent, cache-hit + rate, and client wait time; and +- fail closed on version mismatches. An unavailable service may fall back to an exact + disk hit, but not silently to a different Reasoner or quantization. + +For a 1,024-token prompt the BF16 response is 144 MiB. At `R` examples/s, required +payload bandwidth is approximately `144 * R MiB/s`, before transport framing. This is +usually modest relative to long-video generator step time, but it must be measured at +the intended number of training ranks. Start with a simple streaming RPC into pinned +CPU buffers plus asynchronous H2D copies. Add CUDA IPC/NVLink for same-node workers or +CUDA-aware UCX/RDMA for cross-node workers only if profiling shows transport on the +critical path. + +When a node has exactly eight GPUs, reserving GPU 0 for the service means the generator +job must be launched as a seven-rank job (with FSDP degree adjusted accordingly). A +cleaner production topology is a separate Reasoner pool shared by several eight-GPU +generator nodes. + +## Code ownership and status + +Keep the feature contract in training/model code; a training job must not import Ray or +other heavyweight inference-only dependencies. + +| Area | Current status | +|---|---| +| `configs/base/defaults/model_config.py` | Implemented typed runtime configuration and compatibility validation. | +| `configs/toml_config/sft_config.py` | Implemented the flat `[model.reasoner_conditioning]` TOML schema. | +| `model/generator/mot/unified_mot.py` | Implemented optional UND construction, GEN-only execution, and `prune_und_pathway_`. | +| `model/generator/mot/cosmos3_vfm_network.py` | Implemented generator-only packed layout without text embedding. | +| `model/generator/omni_mot_model.py` | Implemented inline/offline lifecycle, pre-noise submission, synchronized resolution, structural pruning, and startup guards. Remote/read-through remain pending. | +| `model/generator/reasoner_features.py` | Implemented request/result types, UND-only extraction, inline/static memory, and exact two-way external-K/V attention. | +| `model/generator/reasoner_feature_cache.py` | Implemented immutable cache identity, fingerprints, bounded/resumable sharded writers, distributed finalization, manifest validation, and offline provider. | +| `model/generator/mot/multiview_attention.py` | Pending cache-aware per-view/maskless/Flex support. | +| `data/generator/local_datasets/sft_reasoner_documents.py` | Implemented finite deterministic standard-SFT document enumeration with training framing parity. | +| `data/generator/` | Centralized multi-batch prefetch/cache-handle plumbing remains pending. | +| `checkpoint/reasoner_only.py` | Implemented strict regular/EMA Reasoner-only local-DCP loading with independent reads and safe Generator/visual source-leaf omission; generator resume/composition validation remains pending. | +| `scripts/extract_reasoner_features.py` | Implemented resumable distributed shared-POSIX cache extraction and atomic publication without NCCL collectives. | +| `examples/toml/sft_config/` | Pending opt-in Nano cached-Reasoner recipe after full-scale parity passes. | + +An optional remote server can live under the inference/serving tree, but it should +implement the neutral wire schema without making the training client depend on that +server implementation. + +## Implementation sequence + +### Phase 0: golden reference and instrumentation — partial + +- Implemented a fixed-seed tiny H20 parity harness and provider/capture wait timers. +- Implemented overall step/VAE timing and per-rank peak allocated/reserved memory on + a short real-Nano eager run. +- Still pending comprehensive UND/GEN/H2D/backward/optimizer/EMA timing on a + representative compile-enabled run. + +### Phase 1: common boundary and inline replay — implemented + +- `ReasonerFeatureRequest`, `ReasonerFeatureBatch`, and `ReasonerFeatureProvider` + define the common contract. +- `CapturingReasonerKVMemoryState` detaches exact per-layer UND K/V, while + `StaticReasonerKVMemoryState` replays them with live GEN Q/K/V and autograd. +- Single- and multi-sample packing, fingerprint/offset validation, and unsupported + CP/attention modes fail closed. +- A tiny CUDA test validates GEN output and every generator gradient across the full + cached and structurally pruned executions. + +### Phase 2: generator-only model and offline cache — MVP implemented + +Completed: + +- meta-device UND pruning for regular and EMA models; +- generator-only VFM/model execution; +- immutable manifest plus sharded safetensors writer/reader; +- bounded-memory incremental shards, resume markers, and distributed finalization; +- deterministic SFT document enumeration and exact training text framing; +- strict Reasoner-only regular/EMA DCP loading; +- the eight-rank-capable extraction CLI with per-rank timings/memory statistics; +- an eight-H20 end-to-end validation on the official eight-video BridgeData sample + (nine records including the shared CFG-null prompt, 8,594 tokens, eight shards, + atomic publication, and successful eager revalidation); +- an eight-H20 FSDP training A/B with full activation checkpointing, FP32 EMA, + and gradient accumulation two: steps 5--10 used identical token traces and + measured 57.887 seconds/step for joint versus 48.015 seconds/step for offline, + while average peak allocated memory fell from 52.128 GiB/GPU to + 36.498 GiB/GPU; +- strict cache identity, per-record fingerprints, and shard checksums; and +- offline provider integration into the SFT step. + +Still pending: + +- a real full-corpus eight-GPU extraction run and optimizer-step numerical + cache-vs-joint audit; +- automatic fingerprint derivation and token-budget dynamic batching; +- DCP/safetensors full-checkpoint warm-start and generator-only resume tests; +- an opt-in `vision_sft_nano_reasoner_cache` example recipe/launcher; and +- a representative-duration, compile-enabled full-size Nano benchmark. + +In the small eight-GPU validation, every rank independently loaded the 8.19B BF16 +Reasoner in 10.45--10.91 seconds. Rank-local extraction took 1.15--2.77 seconds and +peak allocated memory was 15.55--15.76 GiB (maximum reserved 15.77 GiB). This proves +the multi-worker coordination/publication path, but is not a substitute for the +pending full-corpus throughput and storage benchmark. + +The short training A/B used the full Nano model, FSDP, full activation +checkpointing, FP32 EMA, and eight H20 GPUs, but deliberately disabled compilation +and repeatedly sampled the eight-video cache. Across the six paired steady steps, +offline reduced mean step time by 17.05%, increased aggregate token throughput by +20.56%, and reduced average peak allocated memory by 15.630 GiB/GPU. This validates +the expected direction and magnitude, but the hot 1.18-GiB cache is not a proxy for +full-corpus shared-filesystem behavior. + +### Phase 3: remote/read-through provider — pending + +1. Reuse the implemented request/result and fingerprint schema. +2. Add remote client/service transport and asynchronous prefetch. +3. Benchmark one service GPU against increasing generator-rank counts and scale + replicas based on measured queueing, not a fixed assumed ratio. +4. Add persistent read-through writes only after cache-key and atomicity tests pass. + +### Phase 4: broaden support — pending + +Add separately gated parity suites for multiview masks and per-view captions, context +parallelism, causal/teacher-forcing training, MoE tiers, quantized K/V, and layerwise +H2D staging. + +## Validation gates + +Validated in the current implementation: + +- canonical K/V shape/dtype/device/offset validation and multi-sample isolation; +- inline capture followed by static GEN-only replay; +- tiny dense Qwen output and all-generator-gradient parity between the full cached + model and `prune_und_pathway_` model on one H20; +- absence of UND parameters from the structurally pruned state dict; and +- immutable-cache round trips, batching across shards, content-addressed prompt + reuse, cache misses, stale fingerprints, malformed manifests, and corrupt shard + checksums. + +Remaining correctness gates: + +- one optimizer step produces equivalent selected weights; +- normal prompt, CFG-null prompt, maximum-length prompt, and per-view captions; +- full checkpoint to generator-only warm start, generator-only resume, and full-model + inference composition; and +- end-to-end stale cache/config provenance is rejected during distributed startup + and resume, before any FSDP forward. + +Distributed gates: + +- eight-GPU DP/FSDP with activation checkpointing and EMA is validated in eager + mode; compile-enabled validation remains pending; +- no Reasoner parameter appears in either regular or EMA model state on training ranks; +- no collective divergence when a cache/service request fails; and +- deterministic sample-to-feature association across resume and dataloader workers. + +Performance report: + +- peak allocated and reserved GPU memory for every rank; +- wall time split into data, VAE, provider wait, H2D, GEN forward, backward, optimizer, + and EMA; +- cache bytes/read bandwidth or service queue/compute/transport percentiles; and +- throughput against the unmodified joint baseline. + +The first production promotion should require zero cache misses, no silent fallbacks, +successful checkpoint resume, and a measured reduction in both peak memory and step +time on the real Nano SFT workload. From 858ed9c9b055c043865fb589830147cfb50ee4ad Mon Sep 17 00:00:00 2001 From: Codex Date: Tue, 22 Sep 2026 01:25:22 +0800 Subject: [PATCH 2/4] feat(training): add remote Reasoner feature serving --- .../configs/base/defaults/model_config.py | 3 + .../configs/toml_config/sft_config.py | 7 + .../configs/toml_config/sft_config_test.py | 28 + cosmos_framework/model/_base.py | 27 + .../model/generator/omni_mot_model.py | 253 ++++-- .../omni_mot_reasoner_conditioning_test.py | 293 ++++++- .../model/generator/reasoner_feature_cache.py | 55 +- .../model/generator/reasoner_features.py | 70 +- .../model/generator/reasoner_remote.py | 813 ++++++++++++++++++ .../reasoner_remote_lifecycle_test.py | 442 ++++++++++ .../model/generator/reasoner_remote_server.py | 366 ++++++++ .../generator/reasoner_remote_server_test.py | 384 +++++++++ .../model/generator/reasoner_remote_test.py | 442 ++++++++++ .../model/generator/reasoner_runtime.py | 321 +++++++ .../model/generator/reasoner_runtime_test.py | 292 +++++++ .../model/model_lifecycle_test.py | 123 +++ cosmos_framework/protos/__init__.py | 4 + .../protos/reasoner_features/__init__.py | 4 + .../protos/reasoner_features/v1/__init__.py | 4 + .../v1/reasoner_features.proto | 149 ++++ .../v1/reasoner_features_pb2.py | 67 ++ .../v1/reasoner_features_pb2.pyi | 316 +++++++ .../v1/reasoner_features_pb2_grpc.py | 155 ++++ cosmos_framework/scripts/_train.py | 36 +- .../scripts/extract_reasoner_features.py | 110 +-- .../scripts/extract_reasoner_features_test.py | 26 + .../scripts/serve_reasoner_features.py | 186 ++++ .../scripts/serve_reasoner_features_test.py | 49 ++ cosmos_framework/scripts/train.py | 39 +- cosmos_framework/trainer/__init__.py | 29 +- docs/nano_sft_decoupled_reasoner.md | 269 ++++-- docs/nano_sft_remote_reasoner_test_report.md | 218 +++++ pyproject.toml | 4 + uv.lock | 136 +-- 34 files changed, 5388 insertions(+), 332 deletions(-) create mode 100644 cosmos_framework/model/generator/reasoner_remote.py create mode 100644 cosmos_framework/model/generator/reasoner_remote_lifecycle_test.py create mode 100644 cosmos_framework/model/generator/reasoner_remote_server.py create mode 100644 cosmos_framework/model/generator/reasoner_remote_server_test.py create mode 100644 cosmos_framework/model/generator/reasoner_remote_test.py create mode 100644 cosmos_framework/model/generator/reasoner_runtime.py create mode 100644 cosmos_framework/model/generator/reasoner_runtime_test.py create mode 100644 cosmos_framework/model/model_lifecycle_test.py create mode 100644 cosmos_framework/protos/__init__.py create mode 100644 cosmos_framework/protos/reasoner_features/__init__.py create mode 100644 cosmos_framework/protos/reasoner_features/v1/__init__.py create mode 100644 cosmos_framework/protos/reasoner_features/v1/reasoner_features.proto create mode 100644 cosmos_framework/protos/reasoner_features/v1/reasoner_features_pb2.py create mode 100644 cosmos_framework/protos/reasoner_features/v1/reasoner_features_pb2.pyi create mode 100644 cosmos_framework/protos/reasoner_features/v1/reasoner_features_pb2_grpc.py create mode 100644 cosmos_framework/scripts/serve_reasoner_features.py create mode 100644 cosmos_framework/scripts/serve_reasoner_features_test.py create mode 100644 docs/nano_sft_remote_reasoner_test_report.md diff --git a/cosmos_framework/configs/base/defaults/model_config.py b/cosmos_framework/configs/base/defaults/model_config.py index 32899b742..14510a477 100644 --- a/cosmos_framework/configs/base/defaults/model_config.py +++ b/cosmos_framework/configs/base/defaults/model_config.py @@ -178,6 +178,9 @@ class ReasonerConditioningConfig: prefetch_batches: int = attrs.field(default=2, validator=attrs.validators.ge(0)) layerwise_h2d: bool = False request_timeout_s: float = attrs.field(default=300.0, validator=attrs.validators.gt(0.0)) + connect_timeout_s: float = attrs.field(default=30.0, validator=attrs.validators.gt(0.0)) + request_max_retries: int = attrs.field(default=2, validator=attrs.validators.ge(0)) + retry_backoff_s: float = attrs.field(default=0.25, validator=attrs.validators.ge(0.0)) def __attrs_post_init__(self) -> None: if self.backend in {"offline", "read_through"} and not self.cache_root: diff --git a/cosmos_framework/configs/toml_config/sft_config.py b/cosmos_framework/configs/toml_config/sft_config.py index 6cdf8abdd..7c988d796 100644 --- a/cosmos_framework/configs/toml_config/sft_config.py +++ b/cosmos_framework/configs/toml_config/sft_config.py @@ -273,6 +273,13 @@ class ReasonerConditioningConfig(BaseModel): description="Stage K/V one layer at a time instead of keeping the entire batch resident on GPU.", ) request_timeout_s: float = Field(default=300.0, gt=0.0, description="Remote request deadline in seconds.") + connect_timeout_s: float = Field(default=30.0, gt=0.0, description="Remote startup handshake deadline in seconds.") + request_max_retries: int = Field( + default=2, + ge=0, + description="Retries for transient remote UNAVAILABLE/RESOURCE_EXHAUSTED failures.", + ) + retry_backoff_s: float = Field(default=0.25, ge=0.0, description="Initial remote retry backoff in seconds.") @model_validator(mode="after") def validate_backend_inputs(self) -> "ReasonerConditioningConfig": diff --git a/cosmos_framework/configs/toml_config/sft_config_test.py b/cosmos_framework/configs/toml_config/sft_config_test.py index 155043e72..5b1ff4ab5 100644 --- a/cosmos_framework/configs/toml_config/sft_config_test.py +++ b/cosmos_framework/configs/toml_config/sft_config_test.py @@ -146,6 +146,34 @@ def test_reasoner_conditioning_requires_backend_inputs(self, settings: dict[str, } ) + def test_remote_reasoner_transport_controls_are_validated(self) -> None: + raw = { + "job": {"task": "vfm", "experiment": "vision_sft_nano"}, + "model": { + "reasoner_conditioning": { + "backend": "remote", + "endpoint": "dns:///reasoner:50051", + "reasoner_fingerprint": "reasoner", + "tokenizer_fingerprint": "tokenizer", + "framing_fingerprint": "framing", + "connect_timeout_s": 7.5, + "request_timeout_s": 90.0, + "request_max_retries": 4, + "retry_backoff_s": 0.5, + } + }, + } + + conditioning = SFTExperimentConfig.model_validate(raw).model.reasoner_conditioning + assert conditioning.connect_timeout_s == 7.5 + assert conditioning.request_timeout_s == 90.0 + assert conditioning.request_max_retries == 4 + assert conditioning.retry_backoff_s == 0.5 + + raw["model"]["reasoner_conditioning"]["request_max_retries"] = -1 + with pytest.raises(ValidationError): + SFTExperimentConfig.model_validate(raw) + # --------------------------------------------------------------------------- # # 2. build_hydra_overrides must NOT emit [custom] as per-leaf overrides # diff --git a/cosmos_framework/model/_base.py b/cosmos_framework/model/_base.py index ed9e74dc1..edaad7db0 100644 --- a/cosmos_framework/model/_base.py +++ b/cosmos_framework/model/_base.py @@ -1,12 +1,30 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: OpenMDW-1.1 +import logging from typing import Any import torch from cosmos_framework.utils.lazy_config import LazyDict, instantiate +_LOG = logging.getLogger(__name__) + + +def close_model(model: "ImaginaireModel", *, primary_error: BaseException | None = None) -> None: + """Close a model without replacing an already-active primary failure.""" + + try: + model.close() + except BaseException as close_error: + if primary_error is None: + raise + note = f"Model cleanup also failed with {type(close_error).__name__}: {close_error}" + add_note = getattr(primary_error, "add_note", None) + if add_note is not None: + add_note(note) + _LOG.exception(note) + class ImaginaireModel(torch.nn.Module): """The base model class of Imaginaire. It is inherited from torch.nn.Module. @@ -128,3 +146,12 @@ def on_after_backward(self, iteration: int = 0) -> None: iteration (int): Current iteration number. """ pass + + def close(self) -> None: + """Release non-module resources owned by the model. + + Training entry points call this hook on successful and failed exits. + Implementations must be idempotent because ownership guards may nest. + """ + + pass diff --git a/cosmos_framework/model/generator/omni_mot_model.py b/cosmos_framework/model/generator/omni_mot_model.py index f4b77d092..42f9bbe32 100644 --- a/cosmos_framework/model/generator/omni_mot_model.py +++ b/cosmos_framework/model/generator/omni_mot_model.py @@ -38,7 +38,7 @@ from cosmos_framework.data.generator.sequence_packing.modality import add_special_tokens from cosmos_framework.data.generator.sequence_packing.packers import is_item_generated, uses_single_timestep from cosmos_framework.data.generator.utils import IMAGE_RES_SIZE_INFO, VIDEO_RES_SIZE_INFO -from cosmos_framework.model._base import ImaginaireModel +from cosmos_framework.model._base import ImaginaireModel, close_model from cosmos_framework.model.generator.algorithm.loss.flow_matching import ( ACTION_SLOT_SAMPLE_COUNT_KEY, ACTION_SLOT_SAMPLE_LOSS_KEY, @@ -76,6 +76,7 @@ CapturingReasonerKVMemoryState, ReasonerFeatureBatch, ReasonerFeatureProvider, + ReasonerFeatureSignature, StaticReasonerKVMemoryState, install_reasoner_feature_attention_dispatch, ) @@ -291,6 +292,67 @@ def _reasoner_cache_identity(config: OmniMoTModelConfig) -> ReasonerFeatureCache return ReasonerFeatureCacheIdentity(**values) +def _create_external_reasoner_provider_fail_closed( + create_provider: Callable[[], ReasonerFeatureProvider | None], +) -> ReasonerFeatureProvider | None: + """Create one rank-local provider, then make startup an all-rank gate. + + A remote provider performs its GetInfo/identity handshake in the + constructor. Without this gate, one rank can fail that handshake while + peers proceed into FSDP collectives and hang. Only serializable error text + crosses the process group; the originating rank retains the real exception + as the synchronized failure's cause. + """ + + provider: ReasonerFeatureProvider | None = None + local_error: Exception | None = None + try: + provider = create_provider() + except Exception as error: + if not dist.is_initialized(): + raise + local_error = error + + if not dist.is_initialized(): + return provider + + local_message = None if local_error is None else f"{type(local_error).__name__}: {local_error}" + all_messages: list[str | None] = [None] * dist.get_world_size() + try: + dist.all_gather_object(all_messages, local_message) + except BaseException as collective_error: + if provider is not None: + try: + provider.close() + except BaseException as close_error: + note = f"Reasoner provider cleanup also failed with {type(close_error).__name__}: {close_error}" + add_note = getattr(collective_error, "add_note", None) + if add_note is not None: + add_note(note) + log.exception(note, rank0_only=False) + raise + + failures = [f"rank {rank}: {message}" for rank, message in enumerate(all_messages) if message is not None] + if not failures: + return provider + + synchronized_error = RuntimeError( + "Reasoner feature provider failed during distributed startup: " + "; ".join(failures) + ) + if provider is not None: + try: + provider.close() + except BaseException as close_error: + note = f"Reasoner provider cleanup also failed with {type(close_error).__name__}: {close_error}" + add_note = getattr(synchronized_error, "add_note", None) + if add_note is not None: + add_note(note) + log.exception(note, rank0_only=False) + if local_error is not None: + raise synchronized_error from local_error + raise synchronized_error + + def _validate_reasoner_conditioning(config: OmniMoTModelConfig) -> str: """Validate the deliberately narrow first implementation of external K/V.""" backend = _reasoner_conditioning_backend(config) @@ -385,30 +447,47 @@ def __init__(self, config: OmniMoTModelConfig): "use the base OmniMoTModel Nano SFT path." ) self.reasoner_cache_identity = _reasoner_cache_identity(config) - self.reasoner_feature_provider = self._create_reasoner_feature_provider() - if self.reasoner_cache_identity is None and isinstance( - self.reasoner_feature_provider, OfflineReasonerFeatureProvider - ): - self.reasoner_cache_identity = self.reasoner_feature_provider.manifest.identity - log.info(f"OmniMoTModel: config {self.config}") + self.reasoner_feature_provider: ReasonerFeatureProvider | None = None + try: + create_provider = self._create_reasoner_feature_provider + if self.reasoner_conditioning_backend in _EXTERNAL_REASONER_BACKENDS: + self.reasoner_feature_provider = _create_external_reasoner_provider_fail_closed(create_provider) + else: + self.reasoner_feature_provider = create_provider() + if self.reasoner_cache_identity is None and isinstance( + self.reasoner_feature_provider, OfflineReasonerFeatureProvider + ): + self.reasoner_cache_identity = self.reasoner_feature_provider.manifest.identity + log.info(f"OmniMoTModel: config {self.config}") - # 0. Set up precision - self.set_precision() + # 0. Set up precision + self.set_precision() - # 1. Set data keys and data information - self.set_up_data_key() + # 1. Set data keys and data information + self.set_up_data_key() - # 2. Text, vision, audio, action tokenizers - self.set_up_tokenizers() + # 2. Text, vision, audio, action tokenizers + self.set_up_tokenizers() - # 3. FSDP setup. Note: call this before building the model. - self.set_up_parallelism() + # 3. FSDP setup. Note: call this before building the model. + self.set_up_parallelism() - # 4. Build the denoiser network - self.set_up_model() + # 4. Build the denoiser network + self.set_up_model() - # 5. Set up training time scheduler and inference time sampler - self.set_up_scheduler_and_sampler() + # 5. Set up training time scheduler and inference time sampler + self.set_up_scheduler_and_sampler() + except BaseException as error: + close_model(self, primary_error=error) + raise + + def close(self) -> None: + """Release the external Reasoner provider, if this model owns one.""" + + provider = self.reasoner_feature_provider + self.reasoner_feature_provider = None + if provider is not None: + provider.close() def _create_reasoner_feature_provider(self) -> ReasonerFeatureProvider | None: if self.reasoner_conditioning_backend in {"joint", "inline"}: @@ -422,31 +501,86 @@ def _create_reasoner_feature_provider(self) -> ReasonerFeatureProvider | None: expected_dtype=PRECISION_TO_TORCH_DTYPE[self.config.precision], strict_fingerprint=bool(_conditioning_value(self.config, "strict_fingerprint", True)), ) + if self.reasoner_conditioning_backend == "remote": + # Keep grpc/protobuf optional for joint/inline/offline training. + from cosmos_framework.model.generator.reasoner_remote import RemoteReasonerFeatureProvider + + endpoint = _conditioning_value(self.config, "endpoint") + assert endpoint + if self.reasoner_cache_identity is None: + raise ValueError("Remote Reasoner conditioning requires a complete pinned identity") + return RemoteReasonerFeatureProvider( + endpoint, + expected_identity=self.reasoner_cache_identity, + expected_dtype=PRECISION_TO_TORCH_DTYPE[self.config.precision], + connect_timeout_s=float(_conditioning_value(self.config, "connect_timeout_s", 30.0)), + request_timeout_s=float(_conditioning_value(self.config, "request_timeout_s", 300.0)), + request_max_retries=int(_conditioning_value(self.config, "request_max_retries", 2)), + retry_backoff_s=float(_conditioning_value(self.config, "retry_backoff_s", 0.25)), + ) raise NotImplementedError( f"reasoner_conditioning.backend={self.reasoner_conditioning_backend!r} is reserved by the common " - "feature contract, but its service client has not been implemented yet; use 'offline' or 'inline'." + "feature contract, but read-through composition has not been implemented yet; " + "use 'offline', 'remote', or 'inline'." ) def _validate_reasoner_feature_signature(self, net: torch.nn.Module) -> None: - """Reject a cache built for a different decoder before FSDP/materialization.""" - if not isinstance(self.reasoner_feature_provider, OfflineReasonerFeatureProvider): + """Reject provider features for a different decoder before FSDP/materialization. + + Provider endpoints can sit behind a load balancer, so ranks may receive + different replicas even after every GetInfo call succeeds. Synchronize + the complete architecture signature result before any rank is allowed + to enter ``parallelize_vfm_network`` collectives. + """ + + if self.reasoner_feature_provider is None: return - manifest = self.reasoner_feature_provider.manifest - expected = ( - int(net.num_hidden_layers), - int(net.num_kv_heads), - int(net.head_dim), - ) - actual = ( - manifest.num_layers, - manifest.num_kv_heads, - manifest.head_dim, - ) - if actual != expected: - raise ValueError( - "Reasoner cache/model architecture mismatch: " - f"cache(layers,kv_heads,head_dim)={actual}, model={expected}" + + local_error: Exception | None = None + try: + signature = self.reasoner_feature_provider.signature + if not isinstance(signature, ReasonerFeatureSignature): + raise TypeError( + "Reasoner feature provider signature must be a ReasonerFeatureSignature, " + f"got {type(signature).__name__}" + ) + expected = ReasonerFeatureSignature( + num_layers=int(net.num_hidden_layers), + num_kv_heads=int(net.num_kv_heads), + head_dim=int(net.head_dim), + dtype=PRECISION_TO_TORCH_DTYPE[self.config.precision], ) + if signature != expected: + raise ValueError( + f"Reasoner feature provider/model architecture mismatch: provider={signature!r}, model={expected!r}" + ) + except Exception as error: + if not dist.is_initialized(): + raise + local_error = error + + if not dist.is_initialized(): + return + + local_message = None if local_error is None else f"{type(local_error).__name__}: {local_error}" + all_messages: list[str | None] = [None] * dist.get_world_size() + try: + dist.all_gather_object(all_messages, local_message) + except BaseException as collective_error: + close_model(self, primary_error=collective_error) + raise + + failures = [f"rank {rank}: {message}" for rank, message in enumerate(all_messages) if message is not None] + if not failures: + return + + synchronized_error = RuntimeError( + "Reasoner feature provider signature validation failed before FSDP: " + "; ".join(failures) + ) + close_model(self, primary_error=synchronized_error) + if local_error is not None: + raise synchronized_error from local_error + raise synchronized_error def set_precision(self) -> None: self.precision = PRECISION_TO_TORCH_DTYPE[self.config.precision] @@ -1292,7 +1426,17 @@ def memory_init_training( "initial_temporal_offset": 0, } if self.reasoner_conditioning_backend in _EXTERNAL_REASONER_BACKENDS: - memory_info[REASONER_SAMPLE_KEYS_KEY] = _reasoner_sample_keys(data_batch, gen_data_clean.batch_size) + try: + memory_info[REASONER_SAMPLE_KEYS_KEY] = _reasoner_sample_keys(data_batch, gen_data_clean.batch_size) + except Exception as error: + # Sample-key validity is data dependent, so one DP rank can fail + # while its peers have valid batches. Defer the local failure to + # _resolve_reasoner_feature_batch(), where every rank participates + # in the same pre-FSDP all-gather error gate. BaseException is + # deliberately not captured so process-control signals propagate. + future: Future[ReasonerFeatureBatch] = Future() + future.set_exception(error) + memory_info[REASONER_FEATURE_FUTURE_KEY] = future return gen_data_clean, memory_info def build_memory_state( @@ -1414,19 +1558,28 @@ def pre_noise_memory_hook( if self.reasoner_conditioning_backend in _EXTERNAL_REASONER_BACKENDS: if REASONER_FEATURE_BATCH_KEY in memory_info or REASONER_FEATURE_FUTURE_KEY in memory_info: return memory_info - if self.reasoner_feature_provider is None or self.reasoner_cache_identity is None: - raise RuntimeError( - f"reasoner_conditioning.backend={self.reasoner_conditioning_backend!r} has no feature provider" + try: + if self.reasoner_feature_provider is None or self.reasoner_cache_identity is None: + raise RuntimeError( + f"reasoner_conditioning.backend={self.reasoner_conditioning_backend!r} has no feature provider" + ) + sample_keys = memory_info.get(REASONER_SAMPLE_KEYS_KEY) + if not isinstance(sample_keys, (list, tuple)): + raise RuntimeError(f"memory_info[{REASONER_SAMPLE_KEYS_KEY!r}] is missing") + requests = build_reasoner_feature_requests( + packed_sequence, + tuple(str(key) for key in sample_keys), + self.reasoner_cache_identity, ) - sample_keys = memory_info.get(REASONER_SAMPLE_KEYS_KEY) - if not isinstance(sample_keys, (list, tuple)): - raise RuntimeError(f"memory_info[{REASONER_SAMPLE_KEYS_KEY!r}] is missing") - requests = build_reasoner_feature_requests( - packed_sequence, - tuple(str(key) for key in sample_keys), - self.reasoner_cache_identity, - ) - memory_info[REASONER_FEATURE_FUTURE_KEY] = self.reasoner_feature_provider.submit(requests) + future = self.reasoner_feature_provider.submit(requests) + except Exception as error: + # External request setup and remote submit perform rank-local + # validation synchronously. Preserve those failures in the + # provider contract so every rank reaches the synchronized + # resolution gate instead of leaving peers in FSDP collectives. + future = Future() + future.set_exception(error) + memory_info[REASONER_FEATURE_FUTURE_KEY] = future return memory_info def _prepare_training_data( diff --git a/cosmos_framework/model/generator/omni_mot_reasoner_conditioning_test.py b/cosmos_framework/model/generator/omni_mot_reasoner_conditioning_test.py index 7c10e0f2b..c9feb7a3c 100644 --- a/cosmos_framework/model/generator/omni_mot_reasoner_conditioning_test.py +++ b/cosmos_framework/model/generator/omni_mot_reasoner_conditioning_test.py @@ -8,6 +8,7 @@ import pytest import torch +import cosmos_framework.model.generator.omni_mot_model as omni_mot_model_module from cosmos_framework.callbacks.load_pretrained import ( _warm_start_partially_skips_ema, _warm_start_skips_complete_ema, @@ -16,18 +17,20 @@ from cosmos_framework.model.generator.omni_mot_model import ( REASONER_FEATURE_BATCH_KEY, REASONER_FEATURE_FUTURE_KEY, + REASONER_SAMPLE_KEYS_KEY, OmniMoTModel, + _create_external_reasoner_provider_fail_closed, _reasoner_cache_identity, _validate_reasoner_conditioning, ) from cosmos_framework.model.generator.reasoner_feature_cache import ( - OfflineReasonerFeatureProvider, ReasonerFeatureCacheIdentity, build_reasoner_feature_requests, ) from cosmos_framework.model.generator.reasoner_features import ( CapturingReasonerKVMemoryState, ReasonerFeatureBatch, + ReasonerFeatureSignature, ReasonerLayerKV, StaticReasonerKVMemoryState, ) @@ -109,6 +112,79 @@ def test_non_external_backend_ignores_partial_cache_identity() -> None: assert _reasoner_cache_identity(_config("inline", reasoner_fingerprint="bookkeeping-only")) is None +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_external_provider_startup_preserves_non_distributed_failure(monkeypatch: pytest.MonkeyPatch) -> None: + startup_error = ValueError("identity mismatch") + create_provider = MagicMock(side_effect=startup_error) + monkeypatch.setattr(torch.distributed, "is_initialized", lambda: False) + + with pytest.raises(ValueError) as raised: + _create_external_reasoner_provider_fail_closed(create_provider) + + assert raised.value is startup_error + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_external_provider_startup_synchronizes_local_failure(monkeypatch: pytest.MonkeyPatch) -> None: + startup_error = ValueError("identity mismatch") + create_provider = MagicMock(side_effect=startup_error) + monkeypatch.setattr(torch.distributed, "is_initialized", lambda: True) + monkeypatch.setattr(torch.distributed, "get_world_size", lambda: 2) + + def gather(messages: list[str | None], local_message: str | None) -> None: + assert local_message == "ValueError: identity mismatch" + messages[:] = [local_message, None] + + monkeypatch.setattr(torch.distributed, "all_gather_object", gather) + + with pytest.raises(RuntimeError, match="rank 0: ValueError: identity mismatch") as raised: + _create_external_reasoner_provider_fail_closed(create_provider) + + assert raised.value.__cause__ is startup_error + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_external_provider_startup_closes_local_provider_on_remote_failure( + monkeypatch: pytest.MonkeyPatch, +) -> None: + provider = MagicMock() + monkeypatch.setattr(torch.distributed, "is_initialized", lambda: True) + monkeypatch.setattr(torch.distributed, "get_world_size", lambda: 2) + + def gather(messages: list[str | None], local_message: str | None) -> None: + assert local_message is None + messages[:] = [None, "RuntimeError: GetInfo unavailable"] + + monkeypatch.setattr(torch.distributed, "all_gather_object", gather) + + with pytest.raises(RuntimeError, match="rank 1: RuntimeError: GetInfo unavailable"): + _create_external_reasoner_provider_fail_closed(lambda: provider) + + provider.close.assert_called_once_with() + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_external_provider_startup_returns_provider_when_all_ranks_succeed( + monkeypatch: pytest.MonkeyPatch, +) -> None: + provider = MagicMock() + monkeypatch.setattr(torch.distributed, "is_initialized", lambda: True) + monkeypatch.setattr(torch.distributed, "get_world_size", lambda: 2) + + def gather(messages: list[str | None], local_message: str | None) -> None: + assert local_message is None + messages[:] = [None, None] + + monkeypatch.setattr(torch.distributed, "all_gather_object", gather) + + assert _create_external_reasoner_provider_fail_closed(lambda: provider) is provider + provider.close.assert_not_called() + + @pytest.mark.level(0) @pytest.mark.gpus(0) def test_causal_model_rejects_uncomposed_reasoner_feature_dispatch() -> None: @@ -140,17 +216,68 @@ class UnsupportedModel(OmniMoTModel): @pytest.mark.gpus(0) def test_offline_cache_signature_is_checked_before_model_materialization() -> None: model = object.__new__(OmniMoTModel) - provider = MagicMock(spec=OfflineReasonerFeatureProvider) - provider.manifest = SimpleNamespace(num_layers=3, num_kv_heads=2, head_dim=8) + provider = SimpleNamespace(signature=ReasonerFeatureSignature(3, 2, 8, torch.bfloat16)) model.reasoner_feature_provider = provider + model.config = SimpleNamespace(precision="bfloat16") net = SimpleNamespace(num_hidden_layers=3, num_kv_heads=2, head_dim=8) model._validate_reasoner_feature_signature(net) - provider.manifest.num_layers = 2 + provider.signature = ReasonerFeatureSignature(2, 2, 8, torch.bfloat16) with pytest.raises(ValueError, match="architecture mismatch"): model._validate_reasoner_feature_signature(net) +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_signature_validation_synchronizes_local_mismatch_before_fsdp(monkeypatch: pytest.MonkeyPatch) -> None: + model = object.__new__(OmniMoTModel) + provider = MagicMock() + provider.signature = ReasonerFeatureSignature(2, 2, 8, torch.bfloat16) + model.reasoner_feature_provider = provider + model.config = SimpleNamespace(precision="bfloat16") + net = SimpleNamespace(num_hidden_layers=3, num_kv_heads=2, head_dim=8) + monkeypatch.setattr(torch.distributed, "is_initialized", lambda: True) + monkeypatch.setattr(torch.distributed, "get_world_size", lambda: 2) + + def gather(messages: list[str | None], local_message: str | None) -> None: + assert local_message is not None and "architecture mismatch" in local_message + messages[:] = [local_message, None] + + monkeypatch.setattr(torch.distributed, "all_gather_object", gather) + + with pytest.raises(RuntimeError, match="rank 0: ValueError:.*architecture mismatch") as raised: + model._validate_reasoner_feature_signature(net) + + assert isinstance(raised.value.__cause__, ValueError) + provider.close.assert_called_once_with() + assert model.reasoner_feature_provider is None + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_signature_validation_closes_peer_provider_on_remote_mismatch(monkeypatch: pytest.MonkeyPatch) -> None: + model = object.__new__(OmniMoTModel) + provider = MagicMock() + provider.signature = ReasonerFeatureSignature(3, 2, 8, torch.bfloat16) + model.reasoner_feature_provider = provider + model.config = SimpleNamespace(precision="bfloat16") + net = SimpleNamespace(num_hidden_layers=3, num_kv_heads=2, head_dim=8) + monkeypatch.setattr(torch.distributed, "is_initialized", lambda: True) + monkeypatch.setattr(torch.distributed, "get_world_size", lambda: 2) + + def gather(messages: list[str | None], local_message: str | None) -> None: + assert local_message is None + messages[:] = [None, "ValueError: architecture mismatch"] + + monkeypatch.setattr(torch.distributed, "all_gather_object", gather) + + with pytest.raises(RuntimeError, match="rank 1: ValueError: architecture mismatch"): + model._validate_reasoner_feature_signature(net) + + provider.close.assert_called_once_with() + assert model.reasoner_feature_provider is None + + class _MemoryBuilder: build_memory_state = OmniMoTModel.build_memory_state @@ -367,6 +494,92 @@ def __init__(self) -> None: self.config = SimpleNamespace(reasoner_conditioning={"request_timeout_s": 1.0}) +class _FeatureSubmitter(_FeatureResolver): + pre_noise_memory_hook = OmniMoTModel.pre_noise_memory_hook + + def __init__(self, provider: object) -> None: + super().__init__() + self.reasoner_conditioning_backend = "remote" + self.reasoner_feature_provider = provider + self.reasoner_cache_identity = ReasonerFeatureCacheIdentity("reasoner", "tokenizer", "framing") + + +class _MemoryInitializer(_FeatureResolver): + memory_init_training = OmniMoTModel.memory_init_training + pre_noise_memory_hook = OmniMoTModel.pre_noise_memory_hook + + def __init__(self, provider: object) -> None: + super().__init__() + self.reasoner_conditioning_backend = "remote" + self.reasoner_feature_provider = provider + self.reasoner_cache_identity = ReasonerFeatureCacheIdentity("reasoner", "tokenizer", "framing") + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +@pytest.mark.parametrize( + ("data_batch", "batch_size", "error_match"), + ( + ({}, 1, "requires a per-sample '__key__' field"), + ({"__key__": ["only-one"]}, 2, "Expected 2 non-empty Reasoner sample keys"), + ), +) +def test_sample_key_failure_is_deferred_to_distributed_resolution_gate( + data_batch: dict[str, object], + batch_size: int, + error_match: str, + monkeypatch: pytest.MonkeyPatch, +) -> None: + provider = MagicMock() + model = _MemoryInitializer(provider) + gen_data_clean = SimpleNamespace(batch_size=batch_size) + + returned, memory_info = model.memory_init_training(gen_data_clean, data_batch, []) + + assert returned is gen_data_clean + assert REASONER_SAMPLE_KEYS_KEY not in memory_info + failed_future = memory_info[REASONER_FEATURE_FUTURE_KEY] + assert isinstance(failed_future, Future) + assert failed_future.done() + assert isinstance(failed_future.exception(), ValueError) + + # The pre-noise hook must preserve the failed Future rather than replacing + # it with a later provider submission. + assert model.pre_noise_memory_hook(object(), object(), memory_info) is memory_info + assert memory_info[REASONER_FEATURE_FUTURE_KEY] is failed_future + provider.submit.assert_not_called() + + monkeypatch.setattr(torch.distributed, "is_initialized", lambda: True) + monkeypatch.setattr(torch.distributed, "get_world_size", lambda: 2) + + def gather(messages: list[str | None], local_message: str | None) -> None: + assert local_message is not None and error_match in local_message + messages[:] = [local_message, None] + + monkeypatch.setattr(torch.distributed, "all_gather_object", gather) + with pytest.raises(RuntimeError, match="failed before FSDP forward") as raised: + model._resolve_reasoner_feature_batch(memory_info) + + assert error_match in str(raised.value) + assert raised.value.__cause__ is failed_future.exception() + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_sample_key_extraction_preserves_base_exception(monkeypatch: pytest.MonkeyPatch) -> None: + provider = MagicMock() + model = _MemoryInitializer(provider) + gen_data_clean = SimpleNamespace(batch_size=1) + monkeypatch.setattr( + omni_mot_model_module, + "_reasoner_sample_keys", + MagicMock(side_effect=KeyboardInterrupt("stop")), + ) + + with pytest.raises(KeyboardInterrupt, match="stop"): + model.memory_init_training(gen_data_clean, {"__key__": ["sample"]}, []) + + @pytest.mark.level(0) @pytest.mark.gpus(0) def test_resolve_reasoner_features_consumes_future_and_surfaces_failure() -> None: @@ -388,3 +601,75 @@ def test_resolve_reasoner_features_consumes_future_and_surfaces_failure() -> Non failed.set_exception(KeyError("missing")) with pytest.raises(RuntimeError, match="failed before FSDP forward"): _FeatureResolver()._resolve_reasoner_feature_batch({REASONER_FEATURE_FUTURE_KEY: failed}) + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +@pytest.mark.parametrize("failure_stage", ("provider", "identity", "sample_keys", "build", "submit")) +def test_synchronous_reasoner_submission_failure_reaches_distributed_resolution_gate( + failure_stage: str, + monkeypatch: pytest.MonkeyPatch, +) -> None: + provider = MagicMock() + build_requests = MagicMock(return_value=(object(),)) + if failure_stage == "build": + build_requests.side_effect = ValueError("build failed") + else: + provider.submit.side_effect = ValueError("submit failed") if failure_stage == "submit" else None + monkeypatch.setattr(omni_mot_model_module, "build_reasoner_feature_requests", build_requests) + + submitter = _FeatureSubmitter(provider) + memory_info = {REASONER_SAMPLE_KEYS_KEY: ("sample",)} + if failure_stage == "provider": + submitter.reasoner_feature_provider = None + elif failure_stage == "identity": + submitter.reasoner_cache_identity = None + elif failure_stage == "sample_keys": + memory_info.clear() + assert submitter.pre_noise_memory_hook(object(), object(), memory_info) is memory_info + future = memory_info[REASONER_FEATURE_FUTURE_KEY] + assert isinstance(future, Future) + assert future.done() + deferred_error = future.exception() + assert isinstance(deferred_error, Exception) + local_message = f"{type(deferred_error).__name__}: {deferred_error}" + + monkeypatch.setattr(torch.distributed, "is_initialized", lambda: True) + monkeypatch.setattr(torch.distributed, "get_world_size", lambda: 2) + + def gather(messages: list[str | None], gathered_local_message: str | None) -> None: + assert gathered_local_message == local_message + messages[:] = [gathered_local_message, None] + + monkeypatch.setattr(torch.distributed, "all_gather_object", gather) + with pytest.raises(RuntimeError, match="failed before FSDP forward") as raised: + submitter._resolve_reasoner_feature_batch(memory_info) + + assert f"rank 0: {local_message}" in str(raised.value) + assert raised.value.__cause__ is deferred_error + if failure_stage in {"provider", "identity", "sample_keys", "build"}: + provider.submit.assert_not_called() + else: + provider.submit.assert_called_once() + + +@pytest.mark.level(0) +@pytest.mark.gpus(0) +@pytest.mark.parametrize("failure_stage", ("build", "submit")) +def test_reasoner_submission_preserves_base_exception_control_flow( + failure_stage: str, + monkeypatch: pytest.MonkeyPatch, +) -> None: + provider = MagicMock() + build_requests = MagicMock(return_value=(object(),)) + if failure_stage == "build": + build_requests.side_effect = KeyboardInterrupt("stop") + else: + provider.submit.side_effect = KeyboardInterrupt("stop") + monkeypatch.setattr(omni_mot_model_module, "build_reasoner_feature_requests", build_requests) + + memory_info = {REASONER_SAMPLE_KEYS_KEY: ("sample",)} + with pytest.raises(KeyboardInterrupt, match="stop"): + _FeatureSubmitter(provider).pre_noise_memory_hook(object(), object(), memory_info) + + assert REASONER_FEATURE_FUTURE_KEY not in memory_info diff --git a/cosmos_framework/model/generator/reasoner_feature_cache.py b/cosmos_framework/model/generator/reasoner_feature_cache.py index 29b346933..aae48a3bc 100644 --- a/cosmos_framework/model/generator/reasoner_feature_cache.py +++ b/cosmos_framework/model/generator/reasoner_feature_cache.py @@ -40,7 +40,10 @@ from cosmos_framework.data.generator.sequence_packing.sequence import PackedSequenceBuilder from cosmos_framework.model.generator.reasoner_features import ( ReasonerFeatureBatch, + ReasonerFeatureIdentity, ReasonerFeatureRequest, + ReasonerFeatureSignature, + compute_reasoner_feature_fingerprint, ) REASONER_FEATURE_CACHE_SCHEMA_VERSION = 1 @@ -55,18 +58,7 @@ _CacheSignature = tuple[int, int, int, str] -@dataclass(frozen=True) -class ReasonerFeatureCacheIdentity: - """Immutable global identity of every record in one cache.""" - - reasoner: str - tokenizer: str - framing: str - - def __post_init__(self) -> None: - for field_name, value in asdict(self).items(): - if not isinstance(value, str) or not value: - raise ValueError(f"Cache identity field {field_name!r} must be a non-empty string") +ReasonerFeatureCacheIdentity = ReasonerFeatureIdentity @dataclass(frozen=True) @@ -258,33 +250,6 @@ def _signature_from_dict(value: Any, *, field: str) -> _CacheSignature: ) -def compute_reasoner_feature_fingerprint( - token_ids: torch.Tensor, - position_ids: torch.Tensor, - causal_offsets: torch.Tensor, - *, - identity: ReasonerFeatureCacheIdentity, -) -> str: - """Hash the exact framed Reasoner input plus its immutable global identity.""" - - digest = hashlib.sha256() - digest.update(f"{_FORMAT_NAME}:{REASONER_FEATURE_CACHE_SCHEMA_VERSION}\n".encode()) - digest.update(json.dumps(asdict(identity), sort_keys=True, separators=(",", ":")).encode()) - for name, tensor in ( - ("token_ids", token_ids), - ("position_ids", position_ids), - ("causal_offsets", causal_offsets), - ): - if not isinstance(tensor, torch.Tensor) or tensor.device.type == "meta": - raise TypeError(f"{name} must be a materialized torch.Tensor") - value = tensor.detach().to(device="cpu").contiguous() - digest.update(name.encode()) - digest.update(str(value.dtype).encode()) - digest.update(json.dumps(list(value.shape), separators=(",", ":")).encode()) - digest.update(value.view(torch.uint8).numpy().tobytes()) - return digest.hexdigest() - - def build_reasoner_feature_requests( packed_sequence: PackedSequence, sample_keys: Sequence[str], @@ -1508,6 +1473,15 @@ def __init__( def cache_fingerprint(self) -> str: return self.manifest.cache_fingerprint + @property + def signature(self) -> ReasonerFeatureSignature: + return ReasonerFeatureSignature( + num_layers=self.manifest.num_layers, + num_kv_heads=self.manifest.num_kv_heads, + head_dim=self.manifest.head_dim, + dtype=getattr(torch, self.manifest.dtype), + ) + def submit(self, requests: Sequence[ReasonerFeatureRequest]) -> Future[ReasonerFeatureBatch]: future: Future[ReasonerFeatureBatch] = Future() try: @@ -1516,6 +1490,9 @@ def submit(self, requests: Sequence[ReasonerFeatureRequest]) -> Future[ReasonerF future.set_exception(error) return future + def close(self) -> None: + """Match the provider lifecycle contract; immutable local caches own no live resources.""" + def _verify_shard(self, shard: _Shard) -> Path: path = self.cache_root / shard.path if not path.is_file(): diff --git a/cosmos_framework/model/generator/reasoner_features.py b/cosmos_framework/model/generator/reasoner_features.py index a41cbaf3c..c513e8416 100644 --- a/cosmos_framework/model/generator/reasoner_features.py +++ b/cosmos_framework/model/generator/reasoner_features.py @@ -16,8 +16,10 @@ from __future__ import annotations +import hashlib +import json from concurrent.futures import Future -from dataclasses import dataclass +from dataclasses import asdict, dataclass from typing import Mapping, Protocol, Sequence import torch @@ -36,6 +38,49 @@ _CACHE_CROSS_K = "cross_k" _CACHE_CROSS_V = "cross_v" _CACHE_CAUSAL_OFFSETS = "causal_offsets" +_REASONER_FEATURE_INPUT_FORMAT = "cosmos3-reasoner-kv" +REASONER_FEATURE_INPUT_SCHEMA_VERSION = 1 + + +@dataclass(frozen=True) +class ReasonerFeatureIdentity: + """Pinned model, tokenizer, and prompt-framing identity for Reasoner inputs.""" + + reasoner: str + tokenizer: str + framing: str + + def __post_init__(self) -> None: + for field_name, value in asdict(self).items(): + if not isinstance(value, str) or not value: + raise ValueError(f"Reasoner identity field {field_name!r} must be a non-empty string") + + +def compute_reasoner_feature_fingerprint( + token_ids: torch.Tensor, + position_ids: torch.Tensor, + causal_offsets: torch.Tensor, + *, + identity: ReasonerFeatureIdentity, +) -> str: + """Hash the exact framed Reasoner input plus its immutable global identity.""" + + digest = hashlib.sha256() + digest.update(f"{_REASONER_FEATURE_INPUT_FORMAT}:{REASONER_FEATURE_INPUT_SCHEMA_VERSION}\n".encode()) + digest.update(json.dumps(asdict(identity), sort_keys=True, separators=(",", ":")).encode()) + for name, tensor in ( + ("token_ids", token_ids), + ("position_ids", position_ids), + ("causal_offsets", causal_offsets), + ): + if not isinstance(tensor, torch.Tensor) or tensor.device.type == "meta": + raise TypeError(f"{name} must be a materialized torch.Tensor") + value = tensor.detach().to(device="cpu").contiguous() + digest.update(name.encode()) + digest.update(str(value.dtype).encode()) + digest.update(json.dumps(list(value.shape), separators=(",", ":")).encode()) + digest.update(value.view(torch.uint8).numpy().tobytes()) + return digest.hexdigest() def _detached_tensor(value: torch.Tensor, *, name: str) -> torch.Tensor: @@ -320,11 +365,34 @@ def to( ) +@dataclass(frozen=True) +class ReasonerFeatureSignature: + """Provider-visible architecture and dtype of canonical Reasoner K/V.""" + + num_layers: int + num_kv_heads: int + head_dim: int + dtype: torch.dtype + + def __post_init__(self) -> None: + for name in ("num_layers", "num_kv_heads", "head_dim"): + value = getattr(self, name) + if not isinstance(value, int) or isinstance(value, bool) or value <= 0: + raise ValueError(f"{name} must be a positive integer, got {value!r}") + if self.dtype not in (torch.bfloat16, torch.float16, torch.float32): + raise TypeError(f"Unsupported Reasoner feature dtype: {self.dtype}") + + class ReasonerFeatureProvider(Protocol): """Asynchronous provider contract shared by inline/offline/remote backends.""" + @property + def signature(self) -> ReasonerFeatureSignature: ... + def submit(self, requests: Sequence[ReasonerFeatureRequest]) -> Future[ReasonerFeatureBatch]: ... + def close(self) -> None: ... + @torch.inference_mode() def extract_reasoner_feature_batch( diff --git a/cosmos_framework/model/generator/reasoner_remote.py b/cosmos_framework/model/generator/reasoner_remote.py new file mode 100644 index 000000000..09062a0b4 --- /dev/null +++ b/cosmos_framework/model/generator/reasoner_remote.py @@ -0,0 +1,813 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +"""gRPC wire codec and asynchronous client for remote Reasoner K/V features.""" + +from __future__ import annotations + +import hashlib +import json +import math +import sys +import threading +import time +from collections.abc import Iterable, Iterator, Sequence +from concurrent.futures import Future, ThreadPoolExecutor +from dataclasses import asdict, dataclass +from typing import Any, Protocol + +import torch + +try: + from cosmos_framework.protos.reasoner_features.v1 import reasoner_features_pb2 as reasoner_pb2 +except ModuleNotFoundError as error: # pragma: no cover - exercised only in minimal installations + raise ModuleNotFoundError( + "Remote Reasoner conditioning requires protobuf; install cosmos-framework[reasoner-remote]" + ) from error + +from cosmos_framework.model.generator.reasoner_features import ( + ReasonerFeatureBatch, + ReasonerFeatureIdentity, + ReasonerFeatureRequest, + ReasonerFeatureSignature, +) + +REASONER_FEATURE_PROTOCOL_VERSION = 1 +DEFAULT_MAX_CHUNK_BYTES = 2 * 1024**2 +_MAX_ACCEPTED_CHUNK_BYTES = 4 * 1024**2 + +_TORCH_TO_PROTO_DTYPE = { + torch.int32: reasoner_pb2.DTYPE_INT32, + torch.int64: reasoner_pb2.DTYPE_INT64, + torch.float16: reasoner_pb2.DTYPE_FLOAT16, + torch.bfloat16: reasoner_pb2.DTYPE_BFLOAT16, + torch.float32: reasoner_pb2.DTYPE_FLOAT32, + torch.float64: reasoner_pb2.DTYPE_FLOAT64, +} +_PROTO_TO_TORCH_DTYPE = {value: key for key, value in _TORCH_TO_PROTO_DTYPE.items()} + + +def _require_little_endian() -> None: + if sys.byteorder != "little": + raise RuntimeError("Reasoner feature protocol currently supports little-endian hosts only") + + +def _dtype_size(dtype: torch.dtype) -> int: + return torch.empty((), dtype=dtype).element_size() + + +def _shape_numel(shape: Sequence[int]) -> int: + if any(not isinstance(size, int) or isinstance(size, bool) or size < 0 for size in shape): + raise ValueError(f"Tensor shape contains an invalid dimension: {tuple(shape)!r}") + return math.prod(shape) + + +def encode_tensor_payload(tensor: torch.Tensor) -> Any: + """Encode a small tensor while preserving its exact dtype, shape, and bytes.""" + + _require_little_endian() + if not isinstance(tensor, torch.Tensor) or tensor.device.type == "meta": + raise TypeError("Tensor payload must be a materialized torch.Tensor") + value = tensor.detach().to(device="cpu").contiguous() + try: + dtype = _TORCH_TO_PROTO_DTYPE[value.dtype] + except KeyError as error: + raise TypeError(f"Unsupported Reasoner protocol dtype: {value.dtype}") from error + return reasoner_pb2.TensorPayload( + dtype=dtype, + shape=list(value.shape), + data=value.view(torch.uint8).numpy().tobytes(), + ) + + +def decode_tensor_payload(payload: Any, *, name: str) -> torch.Tensor: + """Decode a small tensor payload into owned CPU storage.""" + + _require_little_endian() + try: + dtype = _PROTO_TO_TORCH_DTYPE[int(payload.dtype)] + except KeyError as error: + raise ValueError(f"{name} uses unsupported protocol dtype {payload.dtype!r}") from error + shape = tuple(int(size) for size in payload.shape) + expected_bytes = _shape_numel(shape) * _dtype_size(dtype) + if len(payload.data) != expected_bytes: + raise ValueError(f"{name} has {len(payload.data)} bytes, expected {expected_bytes} for {shape} {dtype}") + storage = bytearray(payload.data) + return torch.frombuffer(storage, dtype=dtype).reshape(shape) + + +def encode_identity(identity: ReasonerFeatureIdentity) -> Any: + return reasoner_pb2.CacheIdentity(**asdict(identity)) + + +def decode_identity(value: Any) -> ReasonerFeatureIdentity: + return ReasonerFeatureIdentity( + reasoner=str(value.reasoner), + tokenizer=str(value.tokenizer), + framing=str(value.framing), + ) + + +def encode_signature(signature: ReasonerFeatureSignature) -> Any: + try: + dtype = _TORCH_TO_PROTO_DTYPE[signature.dtype] + except KeyError as error: + raise TypeError(f"Unsupported Reasoner feature signature dtype: {signature.dtype}") from error + return reasoner_pb2.FeatureSignature( + dtype=dtype, + num_layers=signature.num_layers, + num_kv_heads=signature.num_kv_heads, + head_dim=signature.head_dim, + ) + + +def decode_signature(value: Any) -> ReasonerFeatureSignature: + try: + dtype = _PROTO_TO_TORCH_DTYPE[int(value.dtype)] + except KeyError as error: + raise ValueError(f"Unsupported Reasoner feature signature dtype {value.dtype!r}") from error + return ReasonerFeatureSignature( + num_layers=int(value.num_layers), + num_kv_heads=int(value.num_kv_heads), + head_dim=int(value.head_dim), + dtype=dtype, + ) + + +def encode_reasoner_request(request: ReasonerFeatureRequest) -> Any: + return reasoner_pb2.ReasonerFeatureRequest( + sample_key=request.sample_key, + token_ids=encode_tensor_payload(request.token_ids), + position_ids=encode_tensor_payload(request.position_ids), + causal_offsets=encode_tensor_payload(request.causal_offsets), + fingerprint=request.fingerprint, + ) + + +def decode_reasoner_request(value: Any) -> ReasonerFeatureRequest: + return ReasonerFeatureRequest( + sample_key=str(value.sample_key), + token_ids=decode_tensor_payload(value.token_ids, name="token_ids"), + position_ids=decode_tensor_payload(value.position_ids, name="position_ids"), + causal_offsets=decode_tensor_payload(value.causal_offsets, name="causal_offsets"), + fingerprint=str(value.fingerprint), + ) + + +def compute_remote_request_id( + requests: Sequence[ReasonerFeatureRequest], + identity: ReasonerFeatureIdentity, +) -> str: + """Return a stable id for idempotent retries of one ordered request batch.""" + + payload = { + "protocol_version": REASONER_FEATURE_PROTOCOL_VERSION, + "identity": asdict(identity), + "requests": [{"sample_key": request.sample_key, "fingerprint": request.fingerprint} for request in requests], + } + return hashlib.sha256(json.dumps(payload, sort_keys=True, separators=(",", ":")).encode()).hexdigest() + + +def encode_generate_request( + requests: Sequence[ReasonerFeatureRequest], + identity: ReasonerFeatureIdentity, +) -> Any: + requests = tuple(requests) + if not requests: + raise ValueError("At least one Reasoner feature request is required") + return reasoner_pb2.GenerateRequest( + protocol_version=REASONER_FEATURE_PROTOCOL_VERSION, + request_id=compute_remote_request_id(requests, identity), + identity=encode_identity(identity), + requests=[encode_reasoner_request(request) for request in requests], + ) + + +def decode_generate_request(value: Any) -> tuple[str, ReasonerFeatureIdentity, tuple[ReasonerFeatureRequest, ...]]: + if int(value.protocol_version) != REASONER_FEATURE_PROTOCOL_VERSION: + raise ValueError( + f"Unsupported Reasoner protocol version {value.protocol_version}; " + f"expected {REASONER_FEATURE_PROTOCOL_VERSION}" + ) + identity = decode_identity(value.identity) + requests = tuple(decode_reasoner_request(request) for request in value.requests) + if not requests: + raise ValueError("At least one Reasoner feature request is required") + expected_id = compute_remote_request_id(requests, identity) + if value.request_id != expected_id: + raise ValueError(f"Reasoner request_id mismatch: request={value.request_id!r}, expected={expected_id!r}") + return str(value.request_id), identity, requests + + +@dataclass(frozen=True) +class RemoteReasonerServiceInfo: + instance_id: str + identity: ReasonerFeatureIdentity + signature: ReasonerFeatureSignature + capabilities: dict[str, int] + max_requests: int + max_batch_tokens: int + max_queued_tokens: int + max_chunk_bytes: int + + +def decode_service_info(value: Any) -> RemoteReasonerServiceInfo: + if int(value.protocol_version) != REASONER_FEATURE_PROTOCOL_VERSION: + raise ValueError( + f"Remote Reasoner protocol mismatch: service={value.protocol_version}, " + f"client={REASONER_FEATURE_PROTOCOL_VERSION}" + ) + if not value.ready: + raise RuntimeError("Remote Reasoner service is not ready") + if not value.service_instance_id: + raise ValueError("Remote Reasoner service returned an empty instance id") + capabilities = {str(item.name): int(item.version) for item in value.capabilities} + if any(not name or version <= 0 for name, version in capabilities.items()): + raise ValueError(f"Remote Reasoner service returned invalid capabilities: {capabilities!r}") + max_requests = int(value.max_requests) + max_batch_tokens = int(value.max_batch_tokens) + max_queued_tokens = int(value.max_queued_tokens) + max_chunk_bytes = int(value.max_chunk_bytes) + if max_requests <= 0: + raise ValueError("Remote Reasoner service returned an invalid request-count limit") + if max_batch_tokens <= 0 or max_queued_tokens < max_batch_tokens: + raise ValueError("Remote Reasoner service returned invalid token admission limits") + if max_chunk_bytes <= 0 or max_chunk_bytes > _MAX_ACCEPTED_CHUNK_BYTES: + raise ValueError( + f"Remote Reasoner max_chunk_bytes={max_chunk_bytes} is outside (0, {_MAX_ACCEPTED_CHUNK_BYTES}]" + ) + return RemoteReasonerServiceInfo( + instance_id=str(value.service_instance_id), + identity=decode_identity(value.identity), + signature=decode_signature(value.signature), + capabilities=capabilities, + max_requests=max_requests, + max_batch_tokens=max_batch_tokens, + max_queued_tokens=max_queued_tokens, + max_chunk_bytes=max_chunk_bytes, + ) + + +@dataclass(frozen=True) +class RemoteReasonerServerTiming: + queue_ns: int + compute_ns: int + device_to_host_ns: int + serialization_ns: int + + +class _TensorAssembler: + def __init__( + self, + *, + dtype: torch.dtype, + shape: tuple[int, ...], + total_bytes: int, + ) -> None: + expected = _shape_numel(shape) * _dtype_size(dtype) + if total_bytes != expected: + raise ValueError(f"Stream tensor declares {total_bytes} bytes, expected {expected} for {shape} {dtype}") + self.dtype = dtype + self.shape = shape + self.total_bytes = total_bytes + self.storage = bytearray(total_bytes) + self.next_offset = 0 + + @property + def complete(self) -> bool: + return self.next_offset == self.total_bytes + + def append(self, *, offset: int, data: bytes) -> None: + if not data: + raise ValueError("Remote Reasoner stream contains an empty tensor chunk") + if offset != self.next_offset: + raise ValueError(f"Remote Reasoner tensor chunk offset={offset}, expected {self.next_offset}") + end = offset + len(data) + if end > self.total_bytes: + raise ValueError(f"Remote Reasoner tensor chunk ends at {end}, beyond {self.total_bytes}") + self.storage[offset:end] = data + self.next_offset = end + + def tensor(self) -> torch.Tensor: + if not self.complete: + raise ValueError(f"Remote Reasoner tensor ended at {self.next_offset}, expected {self.total_bytes}") + return torch.frombuffer(self.storage, dtype=self.dtype).reshape(self.shape) + + +def decode_feature_stream( + responses: Iterable[Any], + *, + request: Any, + expected_identity: ReasonerFeatureIdentity, + expected_signature: ReasonerFeatureSignature, + max_chunk_bytes: int, +) -> tuple[ReasonerFeatureBatch, RemoteReasonerServerTiming]: + """Validate and assemble one streamed response entirely on CPU.""" + + header: Any | None = None + trailer: Any | None = None + offsets: torch.Tensor | None = None + assemblers: dict[tuple[int, int], _TensorAssembler] = {} + expected_keys = [ + (layer_idx, kind) + for layer_idx in range(expected_signature.num_layers) + for kind in (reasoner_pb2.TENSOR_KIND_CROSS_K, reasoner_pb2.TENSOR_KIND_CROSS_V) + ] + tensor_cursor = 0 + next_chunk_index = 0 + received_tensor_bytes = 0 + response_digest = hashlib.sha256() + + for response in responses: + body = response.WhichOneof("body") + if body == "header": + if header is not None or next_chunk_index or trailer is not None: + raise ValueError("Remote Reasoner stream header must be the first and only header") + header = response.header + if int(header.protocol_version) != REASONER_FEATURE_PROTOCOL_VERSION: + raise ValueError(f"Remote Reasoner response protocol mismatch: {header.protocol_version}") + if header.request_id != request.request_id: + raise ValueError( + f"Remote Reasoner response request id mismatch: {header.request_id!r} != {request.request_id!r}" + ) + if decode_identity(header.identity) != expected_identity: + raise ValueError("Remote Reasoner response identity changed after handshake") + if decode_signature(header.signature) != expected_signature: + raise ValueError("Remote Reasoner response signature changed after handshake") + if not header.service_instance_id: + raise ValueError("Remote Reasoner response contains an empty service instance id") + expected_fingerprints = tuple(item.fingerprint for item in request.requests) + if tuple(header.fingerprints) != expected_fingerprints: + raise ValueError("Remote Reasoner response fingerprints do not preserve request order") + offsets = decode_tensor_payload(header.causal_offsets, name="response causal_offsets") + if offsets.dtype != torch.int64 or offsets.ndim != 1: + raise ValueError("Remote Reasoner response causal_offsets must be a one-dimensional int64 tensor") + expected_offsets = [0] + for item in request.requests: + expected_offsets.append(expected_offsets[-1] + int(item.token_ids.shape[0])) + if offsets.tolist() != expected_offsets: + raise ValueError(f"Remote Reasoner response offsets={offsets.tolist()}, expected={expected_offsets}") + expected_total_bytes = ( + expected_offsets[-1] + * expected_signature.num_layers + * 2 + * expected_signature.num_kv_heads + * expected_signature.head_dim + * _dtype_size(expected_signature.dtype) + ) + if int(header.total_tensor_bytes) != expected_total_bytes: + raise ValueError( + f"Remote Reasoner header bytes={header.total_tensor_bytes}, expected={expected_total_bytes}" + ) + continue + + if body == "tensor_chunk": + if header is None or trailer is not None: + raise ValueError("Remote Reasoner tensor chunk appeared outside the header/trailer envelope") + chunk = response.tensor_chunk + if int(chunk.chunk_index) != next_chunk_index: + raise ValueError(f"Remote Reasoner chunk index={chunk.chunk_index}, expected={next_chunk_index}") + if not chunk.data or len(chunk.data) > max_chunk_bytes: + raise ValueError(f"Remote Reasoner chunk size={len(chunk.data)} is outside (0, {max_chunk_bytes}]") + if hashlib.sha256(chunk.data).digest() != chunk.sha256: + raise ValueError(f"Remote Reasoner chunk {chunk.chunk_index} checksum mismatch") + key = (int(chunk.layer_index), int(chunk.kind)) + if tensor_cursor >= len(expected_keys) or key != expected_keys[tensor_cursor]: + expected_key = None if tensor_cursor >= len(expected_keys) else expected_keys[tensor_cursor] + raise ValueError(f"Remote Reasoner tensor order={key}, expected={expected_key}") + try: + dtype = _PROTO_TO_TORCH_DTYPE[int(chunk.dtype)] + except KeyError as error: + raise ValueError(f"Remote Reasoner chunk uses unsupported dtype {chunk.dtype!r}") from error + if dtype != expected_signature.dtype: + raise ValueError(f"Remote Reasoner chunk dtype={dtype}, expected={expected_signature.dtype}") + shape = tuple(int(size) for size in chunk.shape) + expected_shape = ( + int(offsets[-1]), + expected_signature.num_kv_heads, + expected_signature.head_dim, + ) + if shape != expected_shape: + raise ValueError(f"Remote Reasoner chunk shape={shape}, expected={expected_shape}") + assembler = assemblers.get(key) + if assembler is None: + assembler = _TensorAssembler(dtype=dtype, shape=shape, total_bytes=int(chunk.total_bytes)) + assemblers[key] = assembler + elif ( + assembler.dtype != dtype or assembler.shape != shape or assembler.total_bytes != int(chunk.total_bytes) + ): + raise ValueError(f"Remote Reasoner metadata changed between chunks for tensor {key}") + assembler.append(offset=int(chunk.byte_offset), data=chunk.data) + response_digest.update(chunk.data) + received_tensor_bytes += len(chunk.data) + next_chunk_index += 1 + if assembler.complete: + tensor_cursor += 1 + continue + + if body == "trailer": + if header is None or trailer is not None: + raise ValueError("Remote Reasoner stream contains a misplaced or duplicate trailer") + trailer = response.trailer + continue + + raise ValueError("Remote Reasoner stream contains an empty response envelope") + + if header is None or trailer is None or offsets is None: + raise ValueError("Remote Reasoner stream ended without a complete header/trailer envelope") + if tensor_cursor != len(expected_keys) or any(not item.complete for item in assemblers.values()): + raise ValueError("Remote Reasoner stream ended with incomplete K/V tensors") + if int(trailer.total_chunks) != next_chunk_index: + raise ValueError(f"Remote Reasoner trailer chunks={trailer.total_chunks}, received={next_chunk_index}") + if int(trailer.total_tensor_bytes) != received_tensor_bytes: + raise ValueError( + f"Remote Reasoner trailer bytes={trailer.total_tensor_bytes}, received={received_tensor_bytes}" + ) + if int(header.total_tensor_bytes) != received_tensor_bytes: + raise ValueError(f"Remote Reasoner header bytes={header.total_tensor_bytes}, received={received_tensor_bytes}") + if trailer.response_sha256 != response_digest.digest(): + raise ValueError("Remote Reasoner full-response checksum mismatch") + + cross_k = tuple( + assemblers[(layer_idx, reasoner_pb2.TENSOR_KIND_CROSS_K)].tensor() + for layer_idx in range(expected_signature.num_layers) + ) + cross_v = tuple( + assemblers[(layer_idx, reasoner_pb2.TENSOR_KIND_CROSS_V)].tensor() + for layer_idx in range(expected_signature.num_layers) + ) + features = ReasonerFeatureBatch( + cross_k=cross_k, + cross_v=cross_v, + causal_offsets=offsets, + fingerprints=tuple(header.fingerprints), + ) + timing = RemoteReasonerServerTiming( + queue_ns=int(trailer.timing.queue_ns), + compute_ns=int(trailer.timing.compute_ns), + device_to_host_ns=int(trailer.timing.device_to_host_ns), + serialization_ns=int(trailer.timing.serialization_ns), + ) + return features, timing + + +def iter_feature_stream( + features: ReasonerFeatureBatch, + *, + request_id: str, + identity: ReasonerFeatureIdentity, + signature: ReasonerFeatureSignature, + service_instance_id: str, + max_chunk_bytes: int, + queue_ns: int, + compute_ns: int, + device_to_host_ns: int, +) -> Iterator[Any]: + """Serialize a CPU feature batch into the canonical response stream.""" + + if max_chunk_bytes <= 0 or max_chunk_bytes > _MAX_ACCEPTED_CHUNK_BYTES: + raise ValueError(f"max_chunk_bytes must be in (0, {_MAX_ACCEPTED_CHUNK_BYTES}]") + layer = features.layer(0) + actual_signature = ReasonerFeatureSignature( + num_layers=features.num_layers, + num_kv_heads=layer.num_kv_heads, + head_dim=layer.head_dim, + dtype=layer.cross_k.dtype, + ) + if actual_signature != signature: + raise ValueError(f"Feature stream signature={actual_signature!r}, expected={signature!r}") + serialization_ns = 0 + serialization_started = time.perf_counter_ns() + tensors = [tensor for pair in zip(features.cross_k, features.cross_v) for tensor in pair] + cpu_tensors = [tensor.detach().to(device="cpu").contiguous() for tensor in tensors] + total_tensor_bytes = sum(tensor.numel() * tensor.element_size() for tensor in cpu_tensors) + header = reasoner_pb2.GenerateResponse( + header=reasoner_pb2.GenerateHeader( + protocol_version=REASONER_FEATURE_PROTOCOL_VERSION, + request_id=request_id, + identity=encode_identity(identity), + signature=encode_signature(signature), + fingerprints=list(features.fingerprints), + causal_offsets=encode_tensor_payload(features.causal_offsets.to(dtype=torch.int64)), + total_tensor_bytes=total_tensor_bytes, + service_instance_id=service_instance_id, + ) + ) + serialization_ns += time.perf_counter_ns() - serialization_started + yield header + + response_digest = hashlib.sha256() + chunk_index = 0 + for layer_idx in range(signature.num_layers): + for kind, tensor in ( + (reasoner_pb2.TENSOR_KIND_CROSS_K, cpu_tensors[2 * layer_idx]), + (reasoner_pb2.TENSOR_KIND_CROSS_V, cpu_tensors[2 * layer_idx + 1]), + ): + raw = tensor.view(torch.uint8).reshape(-1).numpy() + total_bytes = raw.size + for offset in range(0, total_bytes, max_chunk_bytes): + serialization_started = time.perf_counter_ns() + data = raw[offset : offset + max_chunk_bytes].tobytes() + response_digest.update(data) + response = reasoner_pb2.GenerateResponse( + tensor_chunk=reasoner_pb2.TensorChunk( + chunk_index=chunk_index, + kind=kind, + layer_index=layer_idx, + dtype=_TORCH_TO_PROTO_DTYPE[tensor.dtype], + shape=list(tensor.shape), + byte_offset=offset, + total_bytes=total_bytes, + data=data, + sha256=hashlib.sha256(data).digest(), + ) + ) + chunk_index += 1 + serialization_ns += time.perf_counter_ns() - serialization_started + yield response + yield reasoner_pb2.GenerateResponse( + trailer=reasoner_pb2.GenerateTrailer( + total_chunks=chunk_index, + total_tensor_bytes=total_tensor_bytes, + response_sha256=response_digest.digest(), + timing=reasoner_pb2.ServerTiming( + queue_ns=queue_ns, + compute_ns=compute_ns, + device_to_host_ns=device_to_host_ns, + serialization_ns=serialization_ns, + ), + cache_hits=0, + cache_misses=len(features.fingerprints), + ) + ) + + +class RemoteReasonerRPCError(RuntimeError): + def __init__(self, code: str, details: str, *, retryable: bool) -> None: + super().__init__(f"Remote Reasoner RPC {code}: {details}") + self.code = code + self.details = details + self.retryable = retryable + + +class ReasonerRemoteTransport(Protocol): + def get_info(self, *, timeout_s: float) -> Any: ... + + def generate(self, request: Any, *, timeout_s: float) -> Iterable[Any]: ... + + def close(self) -> None: ... + + +class GrpcReasonerRemoteTransport: + """Thin lazy-imported gRPC transport; tensor validation stays transport-neutral.""" + + def __init__(self, endpoint: str) -> None: + if not endpoint or ":" not in endpoint: + raise ValueError(f"Remote Reasoner endpoint must be host:port, got {endpoint!r}") + try: + import grpc + + from cosmos_framework.protos.reasoner_features.v1 import reasoner_features_pb2_grpc + except ModuleNotFoundError as error: # pragma: no cover - exercised only in minimal installations + raise ModuleNotFoundError( + "Remote Reasoner conditioning requires grpcio; install cosmos-framework[reasoner-remote]" + ) from error + self._grpc = grpc + self._channel = grpc.insecure_channel( + endpoint, + options=( + ("grpc.max_send_message_length", 16 * 1024**2), + ("grpc.max_receive_message_length", 8 * 1024**2), + ), + ) + self._stub = reasoner_features_pb2_grpc.ReasonerFeatureServiceStub(self._channel) + + def _convert_error(self, error: Any) -> RemoteReasonerRPCError: + code = error.code() + return RemoteReasonerRPCError( + code.name, + error.details() or str(error), + retryable=code in {self._grpc.StatusCode.UNAVAILABLE, self._grpc.StatusCode.RESOURCE_EXHAUSTED}, + ) + + def get_info(self, *, timeout_s: float) -> Any: + try: + return self._stub.GetInfo( + reasoner_pb2.GetInfoRequest(protocol_version=REASONER_FEATURE_PROTOCOL_VERSION), + timeout=timeout_s, + wait_for_ready=True, + ) + except self._grpc.RpcError as error: + raise self._convert_error(error) from error + + def generate(self, request: Any, *, timeout_s: float) -> Iterator[Any]: + call: Any | None = None + completed = False + try: + call = self._stub.Generate(request, timeout=timeout_s, wait_for_ready=True) + yield from call + completed = True + except self._grpc.RpcError as error: + raise self._convert_error(error) from error + finally: + # A decoder may fail closed before consuming the full response. In + # that case explicitly cancel the server-streaming RPC rather than + # leaving its server handler and channel resources alive. Do not + # catch BaseException here: GeneratorExit, KeyboardInterrupt, and + # SystemExit must retain their normal control-flow semantics. + if call is not None and not completed: + call.cancel() + + def close(self) -> None: + self._channel.close() + + +def _request_snapshot(request: ReasonerFeatureRequest) -> ReasonerFeatureRequest: + return ReasonerFeatureRequest( + sample_key=request.sample_key, + token_ids=request.token_ids.detach().to(device="cpu").contiguous().clone(), + position_ids=request.position_ids.detach().to(device="cpu").contiguous().clone(), + causal_offsets=request.causal_offsets.detach().to(device="cpu").contiguous().clone(), + fingerprint=request.fingerprint, + ) + + +class RemoteReasonerFeatureProvider: + """Asynchronously fetch canonical K/V from a fail-closed Reasoner service.""" + + def __init__( + self, + endpoint: str, + *, + expected_identity: ReasonerFeatureIdentity, + expected_dtype: torch.dtype, + connect_timeout_s: float = 30.0, + request_timeout_s: float = 300.0, + request_max_retries: int = 2, + retry_backoff_s: float = 0.25, + transport: ReasonerRemoteTransport | None = None, + ) -> None: + # Initialize teardown state before creating either resource. Besides + # normal shutdown, this lets every constructor failure use the exact + # same idempotent cleanup path. + self._state_lock = threading.Lock() + self._close_complete = threading.Event() + self._close_requested = threading.Event() + self._closed = False + self._transport: ReasonerRemoteTransport | None = transport + self._executor: ThreadPoolExecutor | None = None + + self.endpoint = endpoint + self.expected_identity = expected_identity + self.last_server_timing: RemoteReasonerServerTiming | None = None + + try: + connect_timeout_s = float(connect_timeout_s) + self.request_timeout_s = float(request_timeout_s) + self.request_max_retries = int(request_max_retries) + self.retry_backoff_s = float(retry_backoff_s) + if connect_timeout_s <= 0 or self.request_timeout_s <= 0: + raise ValueError("Remote Reasoner connect/request timeouts must be positive") + if self.request_max_retries < 0 or self.retry_backoff_s < 0: + raise ValueError("Remote Reasoner retry count/backoff must be non-negative") + + self._executor = ThreadPoolExecutor(max_workers=1, thread_name_prefix="reasoner-remote") + if self._transport is None: + self._transport = GrpcReasonerRemoteTransport(endpoint) + info = decode_service_info(self._transport.get_info(timeout_s=connect_timeout_s)) + if info.identity != expected_identity: + raise ValueError( + f"Remote Reasoner identity mismatch: service={info.identity!r}, expected={expected_identity!r}" + ) + if info.signature.dtype != expected_dtype: + raise ValueError( + f"Remote Reasoner dtype mismatch: service={info.signature.dtype}, expected={expected_dtype}" + ) + required_capabilities = ( + "server_streaming", + "strict_fingerprint", + "ordered_requests", + "single_document", + ) + missing_capabilities = [ + capability for capability in required_capabilities if info.capabilities.get(capability, 0) < 1 + ] + if missing_capabilities: + raise ValueError( + f"Remote Reasoner service does not advertise required capability v1: {missing_capabilities!r}" + ) + self.info = info + self.signature = info.signature + except BaseException: + self.close(wait=False) + raise + + def submit(self, requests: Sequence[ReasonerFeatureRequest]) -> Future[ReasonerFeatureBatch]: + snapshots = tuple(_request_snapshot(request) for request in requests) + if not snapshots: + raise ValueError("At least one Reasoner feature request is required") + multi_document = [request.sample_key for request in snapshots if request.causal_offsets.numel() != 2] + if multi_document: + raise NotImplementedError( + f"Remote Reasoner service v1 supports one causal document per request; got {multi_document!r}" + ) + if len(snapshots) > self.info.max_requests: + raise ValueError( + f"Reasoner request has {len(snapshots)} samples, service limit is {self.info.max_requests}" + ) + total_tokens = sum(request.token_ids.numel() for request in snapshots) + if total_tokens > self.info.max_batch_tokens: + raise ValueError( + f"Reasoner request has {total_tokens} tokens, service limit is {self.info.max_batch_tokens}" + ) + with self._state_lock: + if self._closed: + raise RuntimeError("RemoteReasonerFeatureProvider is closed") + submitted_at = time.monotonic() + if self._executor is None: # pragma: no cover - guarded by successful construction + raise RuntimeError("RemoteReasonerFeatureProvider executor is unavailable") + return self._executor.submit(self._fetch, snapshots, submitted_at) + + def _fetch( + self, + requests: tuple[ReasonerFeatureRequest, ...], + submitted_at: float, + ) -> ReasonerFeatureBatch: + request = encode_generate_request(requests, self.expected_identity) + deadline = submitted_at + self.request_timeout_s + last_error: RemoteReasonerRPCError | None = None + for attempt in range(self.request_max_retries + 1): + remaining = deadline - time.monotonic() + if remaining <= 0: + if last_error is not None: + raise TimeoutError("Remote Reasoner request exhausted its absolute deadline") from last_error + raise TimeoutError("Remote Reasoner request expired before transport execution") + try: + responses = self._transport.generate(request, timeout_s=remaining) + try: + features, timing = decode_feature_stream( + responses, + request=request, + expected_identity=self.expected_identity, + expected_signature=self.signature, + max_chunk_bytes=self.info.max_chunk_bytes, + ) + finally: + close_responses = getattr(responses, "close", None) + if close_responses is not None: + close_responses() + self.last_server_timing = timing + return features + except RemoteReasonerRPCError as error: + last_error = error + if self._close_requested.is_set() or not error.retryable or attempt >= self.request_max_retries: + raise + sleep_seconds = self.retry_backoff_s * (2**attempt) + remaining = deadline - time.monotonic() + if sleep_seconds >= remaining: + raise TimeoutError("Remote Reasoner retry backoff exceeds the absolute deadline") from error + if self._close_requested.wait(timeout=sleep_seconds): + raise + raise AssertionError("unreachable") + + def close(self, *, wait: bool = True) -> None: + with self._state_lock: + if self._closed: + close_complete = self._close_complete + owns_close = False + else: + self._closed = True + self._close_requested.set() + close_complete = self._close_complete + transport = self._transport + executor = self._executor + owns_close = True + + if not owns_close: + if wait: + close_complete.wait() + return + + # Closing the channel first cancels an in-flight streaming RPC, allowing + # its worker to leave promptly before executor shutdown waits for it. + try: + if transport is not None: + transport.close() + finally: + try: + if executor is not None: + executor.shutdown(wait=wait, cancel_futures=True) + finally: + close_complete.set() + + def __enter__(self) -> RemoteReasonerFeatureProvider: + return self + + def __exit__(self, *_args: object) -> None: + self.close() + + def __del__(self) -> None: # pragma: no cover - best-effort process-exit cleanup + try: + self.close(wait=False) + except BaseException: + pass diff --git a/cosmos_framework/model/generator/reasoner_remote_lifecycle_test.py b/cosmos_framework/model/generator/reasoner_remote_lifecycle_test.py new file mode 100644 index 000000000..4c229abde --- /dev/null +++ b/cosmos_framework/model/generator/reasoner_remote_lifecycle_test.py @@ -0,0 +1,442 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +from __future__ import annotations + +import threading +import time +from collections.abc import Sequence +from typing import Any + +import pytest +import torch + +import cosmos_framework.model.generator.reasoner_remote as reasoner_remote_module +from cosmos_framework.model.generator.reasoner_features import ( + ReasonerFeatureIdentity, + ReasonerFeatureRequest, + ReasonerFeatureSignature, +) +from cosmos_framework.model.generator.reasoner_remote import ( + REASONER_FEATURE_PROTOCOL_VERSION, + GrpcReasonerRemoteTransport, + RemoteReasonerFeatureProvider, + RemoteReasonerRPCError, + encode_identity, + encode_signature, +) +from cosmos_framework.protos.reasoner_features.v1 import reasoner_features_pb2 as reasoner_pb2 + +pytestmark = [pytest.mark.level(0), pytest.mark.gpus(0)] + +_REQUIRED_CAPABILITIES = ( + "server_streaming", + "strict_fingerprint", + "ordered_requests", + "single_document", +) + + +def _identity() -> ReasonerFeatureIdentity: + return ReasonerFeatureIdentity( + reasoner="reasoner-digest", + tokenizer="tokenizer-digest", + framing="framing-digest", + ) + + +def _signature() -> ReasonerFeatureSignature: + return ReasonerFeatureSignature( + num_layers=2, + num_kv_heads=2, + head_dim=4, + dtype=torch.bfloat16, + ) + + +def _service_info( + *, + ready: bool = True, + capabilities: Sequence[str] = _REQUIRED_CAPABILITIES, + max_requests: int = 4, +) -> Any: + return reasoner_pb2.GetInfoResponse( + protocol_version=REASONER_FEATURE_PROTOCOL_VERSION, + service_instance_id="lifecycle-test-replica", + identity=encode_identity(_identity()), + signature=encode_signature(_signature()), + capabilities=tuple(reasoner_pb2.Capability(name=name, version=1) for name in capabilities), + max_requests=max_requests, + max_batch_tokens=64, + max_queued_tokens=128, + max_chunk_bytes=32, + ready=ready, + ) + + +def _request( + sample_key: str = "sample", + *, + causal_offsets: Sequence[int] = (0, 3), +) -> ReasonerFeatureRequest: + token_count = causal_offsets[-1] + return ReasonerFeatureRequest( + sample_key=sample_key, + token_ids=torch.arange(token_count, dtype=torch.int64), + position_ids=torch.arange(token_count, dtype=torch.int64), + causal_offsets=torch.tensor(causal_offsets, dtype=torch.int64), + fingerprint=f"{sample_key}-fingerprint", + ) + + +class _HandshakeTransport: + def __init__( + self, + *, + info: Any | None = None, + get_info_error: BaseException | None = None, + events: list[str] | None = None, + ) -> None: + self.info = info or _service_info() + self.get_info_error = get_info_error + self.events = events + self.close_count = 0 + self.generate_count = 0 + + def get_info(self, *, timeout_s: float) -> Any: + del timeout_s + if self.events is not None: + self.events.append("get_info") + if self.get_info_error is not None: + raise self.get_info_error + return self.info + + def generate(self, request: Any, *, timeout_s: float) -> list[Any]: + del request, timeout_s + self.generate_count += 1 + raise AssertionError("generate should not be called in this lifecycle test") + + def close(self) -> None: + self.close_count += 1 + if self.events is not None: + self.events.append("transport.close") + + +class _RecordingExecutor: + instances: list[_RecordingExecutor] = [] + + def __init__(self, *args: object, events: list[str], **kwargs: object) -> None: + del args, kwargs + self.events = events + self.shutdown_calls: list[tuple[bool, bool]] = [] + self.events.append("executor.init") + self.instances.append(self) + + def shutdown(self, *, wait: bool, cancel_futures: bool) -> None: + self.shutdown_calls.append((wait, cancel_futures)) + self.events.append("executor.shutdown") + + +@pytest.mark.parametrize( + ("transport", "error_type", "error_match"), + ( + (_HandshakeTransport(get_info_error=RuntimeError("get-info failed")), RuntimeError, "get-info failed"), + (_HandshakeTransport(get_info_error=KeyboardInterrupt("stop")), KeyboardInterrupt, "stop"), + (_HandshakeTransport(info=_service_info(ready=False)), RuntimeError, "not ready"), + (_HandshakeTransport(info=_service_info(max_requests=0)), ValueError, "request-count limit"), + ( + _HandshakeTransport(info=_service_info(capabilities=("server_streaming",))), + ValueError, + "required capability", + ), + ), +) +def test_constructor_failure_closes_transport_then_executor( + monkeypatch: pytest.MonkeyPatch, + transport: _HandshakeTransport, + error_type: type[BaseException], + error_match: str, +) -> None: + events: list[str] = [] + transport.events = events + _RecordingExecutor.instances.clear() + monkeypatch.setattr( + reasoner_remote_module, + "ThreadPoolExecutor", + lambda *args, **kwargs: _RecordingExecutor(*args, events=events, **kwargs), + ) + + with pytest.raises(error_type, match=error_match): + RemoteReasonerFeatureProvider( + "unused.test:1234", + expected_identity=_identity(), + expected_dtype=torch.bfloat16, + transport=transport, + ) + + assert transport.close_count == 1 + assert len(_RecordingExecutor.instances) == 1 + assert _RecordingExecutor.instances[0].shutdown_calls == [(False, True)] + assert events[-2:] == ["transport.close", "executor.shutdown"] + + +@pytest.mark.parametrize("missing_capability", _REQUIRED_CAPABILITIES) +def test_handshake_requires_each_v1_capability(missing_capability: str) -> None: + capabilities = tuple(name for name in _REQUIRED_CAPABILITIES if name != missing_capability) + transport = _HandshakeTransport(info=_service_info(capabilities=capabilities)) + + with pytest.raises(ValueError, match=missing_capability): + RemoteReasonerFeatureProvider( + "unused.test:1234", + expected_identity=_identity(), + expected_dtype=torch.bfloat16, + transport=transport, + ) + + assert transport.close_count == 1 + + +class _BlockingTransport(_HandshakeTransport): + def __init__(self) -> None: + super().__init__() + self.generate_entered = threading.Event() + self.channel_closed = threading.Event() + + def generate(self, request: Any, *, timeout_s: float) -> list[Any]: + del request, timeout_s + self.generate_count += 1 + self.generate_entered.set() + if not self.channel_closed.wait(timeout=2): + raise AssertionError("provider waited for its executor before closing the transport") + raise RemoteReasonerRPCError("UNAVAILABLE", "test channel closed", retryable=True) + + def close(self) -> None: + super().close() + self.channel_closed.set() + + +def test_close_cancels_inflight_transport_before_waiting_for_executor() -> None: + transport = _BlockingTransport() + provider = RemoteReasonerFeatureProvider( + "unused.test:1234", + expected_identity=_identity(), + expected_dtype=torch.bfloat16, + request_timeout_s=30, + request_max_retries=2, + retry_backoff_s=10, + transport=transport, + ) + future = provider.submit((_request(),)) + assert transport.generate_entered.wait(timeout=1) + + started = time.monotonic() + provider.close() + elapsed = time.monotonic() - started + + assert elapsed < 1 + with pytest.raises(RemoteReasonerRPCError, match="UNAVAILABLE"): + future.result(timeout=1) + provider.close() + assert transport.close_count == 1 + assert transport.generate_count == 1 + + +class _SlowCloseTransport(_HandshakeTransport): + def __init__(self) -> None: + super().__init__() + self.close_entered = threading.Event() + self.allow_close = threading.Event() + + def close(self) -> None: + self.close_count += 1 + self.close_entered.set() + if not self.allow_close.wait(timeout=2): + raise AssertionError("test did not release transport.close") + + +def test_concurrent_close_is_idempotent_and_waits_for_the_owner() -> None: + transport = _SlowCloseTransport() + provider = RemoteReasonerFeatureProvider( + "unused.test:1234", + expected_identity=_identity(), + expected_dtype=torch.bfloat16, + transport=transport, + ) + errors: list[BaseException] = [] + + def close_provider() -> None: + try: + provider.close() + except BaseException as error: + errors.append(error) + + first = threading.Thread(target=close_provider) + second = threading.Thread(target=close_provider) + first.start() + assert transport.close_entered.wait(timeout=1) + second.start() + time.sleep(0.05) + assert second.is_alive() + + transport.allow_close.set() + first.join(timeout=1) + second.join(timeout=1) + + assert not first.is_alive() + assert not second.is_alive() + assert not errors + assert transport.close_count == 1 + + +class _FakeStatus: + def __init__(self, name: str) -> None: + self.name = name + + +class _FakeRpcError(Exception): + def __init__(self, status: _FakeStatus, details: str) -> None: + super().__init__(details) + self.status = status + self.message = details + + def code(self) -> _FakeStatus: + return self.status + + def details(self) -> str: + return self.message + + +class _FakeGrpc: + RpcError = _FakeRpcError + + class StatusCode: + UNAVAILABLE = _FakeStatus("UNAVAILABLE") + RESOURCE_EXHAUSTED = _FakeStatus("RESOURCE_EXHAUSTED") + + +class _FakeStreamingCall: + def __init__(self, failure: BaseException | None = None) -> None: + self.failure = failure + self.index = 0 + self.cancel_count = 0 + + def __iter__(self) -> _FakeStreamingCall: + return self + + def __next__(self) -> str: + if self.index == 0: + self.index += 1 + return "first-response" + if self.failure is not None: + raise self.failure + raise StopIteration + + def cancel(self) -> None: + self.cancel_count += 1 + + +class _FakeStub: + def __init__(self, call: _FakeStreamingCall) -> None: + self.call = call + + def Generate(self, request: Any, *, timeout: float, wait_for_ready: bool) -> _FakeStreamingCall: # noqa: N802 + del request, timeout, wait_for_ready + return self.call + + +class _FakeUnaryStub: + def __init__(self, failure: BaseException) -> None: + self.failure = failure + + def GetInfo(self, request: Any, *, timeout: float, wait_for_ready: bool) -> Any: # noqa: N802 + del request, timeout, wait_for_ready + raise self.failure + + +def _grpc_transport(call: _FakeStreamingCall) -> GrpcReasonerRemoteTransport: + transport = object.__new__(GrpcReasonerRemoteTransport) + transport._grpc = _FakeGrpc() # type: ignore[attr-defined] + transport._stub = _FakeStub(call) # type: ignore[attr-defined] + return transport + + +def _grpc_unary_transport(failure: BaseException) -> GrpcReasonerRemoteTransport: + transport = object.__new__(GrpcReasonerRemoteTransport) + transport._grpc = _FakeGrpc() # type: ignore[attr-defined] + transport._stub = _FakeUnaryStub(failure) # type: ignore[attr-defined] + return transport + + +def test_grpc_stream_cancelled_when_consumer_closes_generator_early() -> None: + call = _FakeStreamingCall() + responses = _grpc_transport(call).generate(object(), timeout_s=1) + + assert next(responses) == "first-response" + responses.close() + + assert call.cancel_count == 1 + + +@pytest.mark.parametrize("control_flow_error", (KeyboardInterrupt(), SystemExit(3))) +def test_grpc_stream_does_not_convert_control_flow_exceptions(control_flow_error: BaseException) -> None: + call = _FakeStreamingCall(control_flow_error) + responses = _grpc_transport(call).generate(object(), timeout_s=1) + + assert next(responses) == "first-response" + with pytest.raises(type(control_flow_error)): + next(responses) + + assert call.cancel_count == 1 + + +def test_grpc_stream_only_converts_grpc_rpc_errors() -> None: + call = _FakeStreamingCall(_FakeRpcError(_FakeGrpc.StatusCode.UNAVAILABLE, "temporarily unavailable")) + responses = _grpc_transport(call).generate(object(), timeout_s=1) + + assert next(responses) == "first-response" + with pytest.raises(RemoteReasonerRPCError, match="temporarily unavailable") as error: + next(responses) + + assert error.value.retryable + assert call.cancel_count == 1 + + +@pytest.mark.parametrize("control_flow_error", (KeyboardInterrupt(), SystemExit(3))) +def test_grpc_get_info_does_not_convert_control_flow_exceptions(control_flow_error: BaseException) -> None: + with pytest.raises(type(control_flow_error)): + _grpc_unary_transport(control_flow_error).get_info(timeout_s=1) + + +def test_submit_rejects_multi_document_request_before_transport() -> None: + transport = _HandshakeTransport() + provider = RemoteReasonerFeatureProvider( + "unused.test:1234", + expected_identity=_identity(), + expected_dtype=torch.bfloat16, + transport=transport, + ) + try: + with pytest.raises(NotImplementedError, match="one causal document"): + provider.submit((_request(causal_offsets=(0, 1, 3)),)) + finally: + provider.close() + + assert transport.generate_count == 0 + + +def test_submit_rejects_request_count_above_advertised_limit_before_transport() -> None: + transport = _HandshakeTransport() + provider = RemoteReasonerFeatureProvider( + "unused.test:1234", + expected_identity=_identity(), + expected_dtype=torch.bfloat16, + transport=transport, + ) + try: + requests = tuple(_request(f"sample-{index}") for index in range(5)) + with pytest.raises(ValueError, match=r"5 samples, service limit is 4"): + provider.submit(requests) + finally: + provider.close() + + assert transport.generate_count == 0 diff --git a/cosmos_framework/model/generator/reasoner_remote_server.py b/cosmos_framework/model/generator/reasoner_remote_server.py new file mode 100644 index 000000000..5784ad2ab --- /dev/null +++ b/cosmos_framework/model/generator/reasoner_remote_server.py @@ -0,0 +1,366 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +"""Bounded one-replica gRPC service for frozen Reasoner K/V extraction.""" + +from __future__ import annotations + +import threading +import time +import uuid +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass +from typing import Any + +import torch + +try: + import grpc + + from cosmos_framework.protos.reasoner_features.v1 import reasoner_features_pb2 as reasoner_pb2 + from cosmos_framework.protos.reasoner_features.v1 import reasoner_features_pb2_grpc +except ModuleNotFoundError as error: # pragma: no cover - exercised only in minimal installations + raise ModuleNotFoundError( + "The Reasoner feature service requires grpcio and protobuf; install cosmos-framework[reasoner-remote]" + ) from error + +from cosmos_framework.model.generator.reasoner_features import ReasonerFeatureBatch +from cosmos_framework.model.generator.reasoner_remote import ( + DEFAULT_MAX_CHUNK_BYTES, + REASONER_FEATURE_PROTOCOL_VERSION, + decode_generate_request, + encode_identity, + encode_signature, + iter_feature_stream, +) +from cosmos_framework.model.generator.reasoner_runtime import ReasonerFeatureRuntime, ReasonerRuntimeInvariantError + +_EXECUTION_LOCK_POLL_SECONDS = 0.05 + + +@dataclass(frozen=True) +class ReasonerServiceMetrics: + accepted_batches: int + rejected_batches: int + failed_batches: int + completed_batches: int + accepted_tokens: int + queued_tokens: int + healthy: bool + + +class _TokenAdmission: + def __init__(self, max_queued_tokens: int) -> None: + if max_queued_tokens <= 0: + raise ValueError("max_queued_tokens must be positive") + self.max_queued_tokens = int(max_queued_tokens) + self.queued_tokens = 0 + self._lock = threading.Lock() + + def acquire(self, tokens: int) -> bool: + with self._lock: + if tokens <= 0 or self.queued_tokens + tokens > self.max_queued_tokens: + return False + self.queued_tokens += tokens + return True + + def release(self, tokens: int) -> None: + with self._lock: + self.queued_tokens -= tokens + if self.queued_tokens < 0: + raise RuntimeError("Reasoner service token admission accounting underflow") + + def snapshot(self) -> int: + with self._lock: + return self.queued_tokens + + +def _cpu_feature_batch(features: ReasonerFeatureBatch) -> ReasonerFeatureBatch: + return ReasonerFeatureBatch( + cross_k=tuple(tensor.detach().to(device="cpu").contiguous() for tensor in features.cross_k), + cross_v=tuple(tensor.detach().to(device="cpu").contiguous() for tensor in features.cross_v), + causal_offsets=features.causal_offsets.detach().to(device="cpu", dtype=torch.int64).contiguous(), + fingerprints=features.fingerprints, + ) + + +class ReasonerFeatureService(reasoner_features_pb2_grpc.ReasonerFeatureServiceServicer): + """Serve one runtime with bounded admission and exactly one GPU executor.""" + + def __init__( + self, + runtime: ReasonerFeatureRuntime, + *, + max_queued_tokens: int, + max_chunk_bytes: int = DEFAULT_MAX_CHUNK_BYTES, + service_instance_id: str | None = None, + ) -> None: + if max_queued_tokens < runtime.max_total_tokens: + raise ValueError("max_queued_tokens must be at least the runtime max_total_tokens") + if max_chunk_bytes <= 0 or max_chunk_bytes > 4 * 1024**2: + raise ValueError("max_chunk_bytes must be in (0, 4 MiB]") + self.runtime = runtime + self.max_chunk_bytes = int(max_chunk_bytes) + self.service_instance_id = service_instance_id or str(uuid.uuid4()) + self._admission = _TokenAdmission(max_queued_tokens) + self._execution_lock = threading.Lock() + self._state_lock = threading.Lock() + self._healthy = True + self._accepted_batches = 0 + self._rejected_batches = 0 + self._failed_batches = 0 + self._completed_batches = 0 + self._accepted_tokens = 0 + + def _record_accept(self, tokens: int) -> None: + with self._state_lock: + self._accepted_batches += 1 + self._accepted_tokens += tokens + + def _record_reject(self) -> None: + with self._state_lock: + self._rejected_batches += 1 + + def _record_failure(self, *, unhealthy: bool = False) -> None: + with self._state_lock: + self._failed_batches += 1 + if unhealthy: + self._healthy = False + + def _record_complete(self) -> None: + with self._state_lock: + self._completed_batches += 1 + + def _is_healthy(self) -> bool: + with self._state_lock: + return self._healthy + + def metrics(self) -> ReasonerServiceMetrics: + with self._state_lock: + return ReasonerServiceMetrics( + accepted_batches=self._accepted_batches, + rejected_batches=self._rejected_batches, + failed_batches=self._failed_batches, + completed_batches=self._completed_batches, + accepted_tokens=self._accepted_tokens, + queued_tokens=self._admission.snapshot(), + healthy=self._healthy, + ) + + def GetInfo(self, request: Any, context: Any) -> Any: # noqa: N802 - gRPC API name + if int(request.protocol_version) != REASONER_FEATURE_PROTOCOL_VERSION: + context.abort( + grpc.StatusCode.FAILED_PRECONDITION, + f"protocol_version={request.protocol_version}, expected={REASONER_FEATURE_PROTOCOL_VERSION}", + ) + return reasoner_pb2.GetInfoResponse( + protocol_version=REASONER_FEATURE_PROTOCOL_VERSION, + service_instance_id=self.service_instance_id, + identity=encode_identity(self.runtime.identity), + signature=encode_signature(self.runtime.signature), + capabilities=( + reasoner_pb2.Capability(name="server_streaming", version=1), + reasoner_pb2.Capability(name="strict_fingerprint", version=1), + reasoner_pb2.Capability(name="ordered_requests", version=1), + reasoner_pb2.Capability(name="single_document", version=1), + ), + max_batch_tokens=self.runtime.max_total_tokens, + max_queued_tokens=self._admission.max_queued_tokens, + max_chunk_bytes=self.max_chunk_bytes, + ready=self._is_healthy(), + max_requests=self.runtime.max_requests, + ) + + def Generate(self, request: Any, context: Any) -> Any: # noqa: N802 - gRPC API name + """Validate, enqueue, execute, stage to CPU, then stream an all-or-nothing result.""" + + if not self._is_healthy(): + self._record_reject() + context.abort(grpc.StatusCode.UNAVAILABLE, "Reasoner replica is unhealthy and requires restart") + if int(request.protocol_version) != REASONER_FEATURE_PROTOCOL_VERSION: + self._record_reject() + context.abort( + grpc.StatusCode.FAILED_PRECONDITION, + f"protocol_version={request.protocol_version}, expected={REASONER_FEATURE_PROTOCOL_VERSION}", + ) + try: + request_id, identity, requests = decode_generate_request(request) + except (TypeError, ValueError) as error: + self._record_reject() + context.abort(grpc.StatusCode.INVALID_ARGUMENT, str(error)) + if identity != self.runtime.identity: + self._record_reject() + context.abort( + grpc.StatusCode.FAILED_PRECONDITION, + f"Reasoner identity mismatch: request={identity!r}, service={self.runtime.identity!r}", + ) + multi_document = [item.sample_key for item in requests if item.causal_offsets.numel() != 2] + if multi_document: + self._record_reject() + context.abort( + grpc.StatusCode.UNIMPLEMENTED, + f"Reasoner service v1 supports one causal document per request; got {multi_document!r}", + ) + token_count = sum(item.token_ids.numel() for item in requests) + empty_requests = [item.sample_key for item in requests if item.token_ids.numel() == 0] + if empty_requests: + self._record_reject() + context.abort( + grpc.StatusCode.INVALID_ARGUMENT, + f"Reasoner requests must contain at least one token; empty requests={empty_requests!r}", + ) + if token_count > self.runtime.max_total_tokens or len(requests) > self.runtime.max_requests: + self._record_reject() + context.abort( + grpc.StatusCode.INVALID_ARGUMENT, + f"Request batch has {len(requests)} requests/{token_count} tokens; limits are " + f"{self.runtime.max_requests}/{self.runtime.max_total_tokens}", + ) + if not self._admission.acquire(token_count): + self._record_reject() + context.abort( + grpc.StatusCode.RESOURCE_EXHAUSTED, + f"Reasoner queue token budget {self._admission.max_queued_tokens} is full", + ) + self._record_accept(token_count) + + cancelled = threading.Event() + if context is not None: + context.add_callback(cancelled.set) + + def request_is_active() -> bool: + # Direct iterator tests use no gRPC context. Production RPCs always + # supply one and therefore exercise cancellation/deadline checks. + return context is None or context.is_active() + + try: + enqueued_ns = time.perf_counter_ns() + execution_lock_acquired = False + try: + while not execution_lock_acquired: + if cancelled.is_set() or not request_is_active(): + self._record_failure() + context.abort(grpc.StatusCode.CANCELLED, "Reasoner request was cancelled while queued") + if not self._is_healthy(): + self._record_failure() + context.abort( + grpc.StatusCode.UNAVAILABLE, + "Reasoner replica became unhealthy while the request was queued", + ) + execution_lock_acquired = self._execution_lock.acquire(timeout=_EXECUTION_LOCK_POLL_SECONDS) + + # Cancellation and health can change after the last poll but + # before the lock acquisition. Re-check both at the exact GPU + # execution boundary so expired calls never become ghost work, + # and requests queued behind a fatal CUDA failure do not touch + # the unhealthy replica. + if cancelled.is_set() or not request_is_active(): + self._record_failure() + context.abort(grpc.StatusCode.CANCELLED, "Reasoner request was cancelled before execution") + if not self._is_healthy(): + self._record_failure() + context.abort( + grpc.StatusCode.UNAVAILABLE, + "Reasoner replica became unhealthy while the request was queued", + ) + + try: + compute_started_ns = time.perf_counter_ns() + queue_ns = compute_started_ns - enqueued_ns + features = self.runtime.execute(requests) + if self.runtime.device.type == "cuda": + torch.cuda.synchronize(self.runtime.device) + compute_finished_ns = time.perf_counter_ns() + d2h_started_ns = compute_finished_ns + cpu_features = _cpu_feature_batch(features) + if self.runtime.device.type == "cuda": + torch.cuda.synchronize(self.runtime.device) + d2h_finished_ns = time.perf_counter_ns() + del features + except torch.cuda.OutOfMemoryError as error: + # Mark the replica unhealthy before releasing the execution + # lock. Otherwise a queued request can acquire the lock and + # enter the same broken CUDA context in the intervening race. + self._record_failure(unhealthy=True) + context.abort( + grpc.StatusCode.RESOURCE_EXHAUSTED, + f"Reasoner replica encountered CUDA OOM and was marked unhealthy: {error}", + ) + except (TypeError, ValueError) as error: + self._record_failure() + context.abort(grpc.StatusCode.INVALID_ARGUMENT, str(error)) + except NotImplementedError as error: + self._record_failure() + context.abort(grpc.StatusCode.UNIMPLEMENTED, str(error)) + except ReasonerRuntimeInvariantError as error: + # An output contract violation is deterministic for this + # loaded replica. Mark it unhealthy while still holding + # the execution lock so queued calls cannot enter the same + # broken runtime in the intervening race. + self._record_failure(unhealthy=True) + context.abort( + grpc.StatusCode.INTERNAL, + f"Reasoner runtime violated its output contract and was marked unhealthy: {error}", + ) + except RuntimeError as error: + is_cuda_failure = "CUDA" in str(error) or "cuda" in str(error) + self._record_failure(unhealthy=is_cuda_failure) + context.abort(grpc.StatusCode.INTERNAL, str(error)) + except Exception as error: + self._record_failure() + context.abort(grpc.StatusCode.INTERNAL, f"{type(error).__name__}: {error}") + finally: + if execution_lock_acquired: + self._execution_lock.release() + + try: + yield from iter_feature_stream( + cpu_features, + request_id=request_id, + identity=self.runtime.identity, + signature=self.runtime.signature, + service_instance_id=self.service_instance_id, + max_chunk_bytes=self.max_chunk_bytes, + queue_ns=queue_ns, + compute_ns=compute_finished_ns - compute_started_ns, + device_to_host_ns=d2h_finished_ns - d2h_started_ns, + ) + except GeneratorExit: + self._record_failure() + raise + except Exception as error: + self._record_failure() + context.abort( + grpc.StatusCode.INTERNAL, + f"Reasoner response serialization failed: {type(error).__name__}: {error}", + ) + self._record_complete() + finally: + # Keep the token reservation through streaming so slow/cancelled + # clients cannot accumulate unbounded CPU-resident K/V responses. + self._admission.release(token_count) + + +def create_reasoner_grpc_server( + service: ReasonerFeatureService, + *, + address: str, + max_rpc_workers: int = 16, +) -> tuple[Any, int]: + """Create a configured server and bind it without starting it.""" + + if max_rpc_workers <= 0: + raise ValueError("max_rpc_workers must be positive") + server = grpc.server( + ThreadPoolExecutor(max_workers=max_rpc_workers, thread_name_prefix="reasoner-rpc"), + maximum_concurrent_rpcs=max_rpc_workers, + options=( + ("grpc.max_receive_message_length", 16 * 1024**2), + ("grpc.max_send_message_length", 8 * 1024**2), + ), + ) + reasoner_features_pb2_grpc.add_ReasonerFeatureServiceServicer_to_server(service, server) + bound_port = server.add_insecure_port(address) + if bound_port == 0: + raise RuntimeError(f"Failed to bind Reasoner feature service to {address!r}") + return server, bound_port diff --git a/cosmos_framework/model/generator/reasoner_remote_server_test.py b/cosmos_framework/model/generator/reasoner_remote_server_test.py new file mode 100644 index 000000000..ee5db1c8e --- /dev/null +++ b/cosmos_framework/model/generator/reasoner_remote_server_test.py @@ -0,0 +1,384 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +from __future__ import annotations + +import threading +import time +from collections.abc import Callable, Iterator, Sequence +from contextlib import contextmanager +from typing import Any + +import pytest +import torch + +pytest.importorskip("google.protobuf") +grpc = pytest.importorskip("grpc") + +from cosmos_framework.model.generator.reasoner_features import ( # noqa: E402 + ReasonerFeatureBatch, + ReasonerFeatureIdentity, + ReasonerFeatureRequest, + ReasonerFeatureSignature, +) +from cosmos_framework.model.generator.reasoner_remote import ( # noqa: E402 + REASONER_FEATURE_PROTOCOL_VERSION, + encode_generate_request, +) +from cosmos_framework.model.generator.reasoner_remote_server import ( # noqa: E402 + ReasonerFeatureService, + create_reasoner_grpc_server, +) +from cosmos_framework.model.generator.reasoner_runtime import ReasonerRuntimeInvariantError # noqa: E402 +from cosmos_framework.protos.reasoner_features.v1 import reasoner_features_pb2 as reasoner_pb2 # noqa: E402 +from cosmos_framework.protos.reasoner_features.v1 import reasoner_features_pb2_grpc # noqa: E402 + +pytestmark = [pytest.mark.level(0), pytest.mark.gpus(0)] + + +def _identity() -> ReasonerFeatureIdentity: + return ReasonerFeatureIdentity(reasoner="reasoner", tokenizer="tokenizer", framing="framing") + + +def _signature() -> ReasonerFeatureSignature: + return ReasonerFeatureSignature(num_layers=1, num_kv_heads=1, head_dim=1, dtype=torch.float32) + + +def _request(sample_key: str, token_ids: Sequence[int] = (1,)) -> ReasonerFeatureRequest: + tokens = torch.tensor(token_ids, dtype=torch.int64) + return ReasonerFeatureRequest( + sample_key=sample_key, + token_ids=tokens, + position_ids=torch.arange(tokens.numel(), dtype=torch.int64), + causal_offsets=torch.tensor([0, tokens.numel()], dtype=torch.int64), + fingerprint=f"{sample_key}-fingerprint", + ) + + +def _features(requests: Sequence[ReasonerFeatureRequest]) -> ReasonerFeatureBatch: + offsets = [0] + for request in requests: + offsets.append(offsets[-1] + request.token_ids.numel()) + total_tokens = offsets[-1] + return ReasonerFeatureBatch( + cross_k=(torch.zeros(total_tokens, 1, 1, dtype=torch.float32),), + cross_v=(torch.ones(total_tokens, 1, 1, dtype=torch.float32),), + causal_offsets=torch.tensor(offsets, dtype=torch.int64), + fingerprints=tuple(request.fingerprint for request in requests), + ) + + +class _FakeRuntime: + def __init__( + self, + behavior: Callable[[tuple[ReasonerFeatureRequest, ...]], ReasonerFeatureBatch] | None = None, + ) -> None: + self.identity = _identity() + self.signature = _signature() + self.max_requests = 8 + self.max_total_tokens = 64 + self.device = torch.device("cpu") + self._behavior = behavior + self._calls_lock = threading.Lock() + self._calls: list[str] = [] + + @property + def calls(self) -> list[str]: + with self._calls_lock: + return list(self._calls) + + def execute(self, requests: Sequence[ReasonerFeatureRequest]) -> ReasonerFeatureBatch: + requests = tuple(requests) + with self._calls_lock: + self._calls.append(requests[0].sample_key) + if self._behavior is not None: + return self._behavior(requests) + return _features(requests) + + +@contextmanager +def _running_service( + runtime: _FakeRuntime, + *, + max_rpc_workers: int = 2, +) -> Iterator[tuple[ReasonerFeatureService, Any]]: + service = ReasonerFeatureService( + runtime, # type: ignore[arg-type] + max_queued_tokens=128, + max_chunk_bytes=64, + service_instance_id="test-replica", + ) + server, port = create_reasoner_grpc_server( + service, + address="127.0.0.1:0", + max_rpc_workers=max_rpc_workers, + ) + server.start() + channel = grpc.insecure_channel(f"127.0.0.1:{port}") + stub = reasoner_features_pb2_grpc.ReasonerFeatureServiceStub(channel) + try: + yield service, stub + finally: + channel.close() + server.stop(grace=0).wait(timeout=5) + + +def _consume(call: Any, result: dict[str, Any]) -> None: + try: + result["responses"] = list(call) + except grpc.RpcError as error: + result["code"] = error.code() + result["details"] = error.details() + + +def _wait_until(predicate: Callable[[], bool], *, timeout: float = 3.0) -> None: + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + if predicate(): + return + time.sleep(0.01) + raise AssertionError("condition was not satisfied before timeout") + + +def test_generate_maps_protocol_and_empty_token_validation_statuses() -> None: + runtime = _FakeRuntime() + with _running_service(runtime) as (service, stub): + info = stub.GetInfo( + reasoner_pb2.GetInfoRequest(protocol_version=REASONER_FEATURE_PROTOCOL_VERSION), + timeout=2, + ) + assert info.max_requests == runtime.max_requests + + wrong_version = reasoner_pb2.GenerateRequest(protocol_version=REASONER_FEATURE_PROTOCOL_VERSION + 1) + with pytest.raises(grpc.RpcError) as protocol_error: + list(stub.Generate(wrong_version, timeout=2)) + assert protocol_error.value.code() == grpc.StatusCode.FAILED_PRECONDITION + + empty = encode_generate_request((_request("empty", ()),), runtime.identity) + with pytest.raises(grpc.RpcError) as empty_error: + list(stub.Generate(empty, timeout=2)) + assert empty_error.value.code() == grpc.StatusCode.INVALID_ARGUMENT + + assert runtime.calls == [] + assert service.metrics().rejected_batches == 2 + assert service.metrics().queued_tokens == 0 + + +def test_maximum_concurrent_rpcs_rejects_before_the_executor_queue() -> None: + entered = threading.Event() + release = threading.Event() + + def behavior(requests: tuple[ReasonerFeatureRequest, ...]) -> ReasonerFeatureBatch: + entered.set() + if not release.wait(timeout=5): + raise AssertionError("test did not release the first RPC") + return _features(requests) + + runtime = _FakeRuntime(behavior) + with _running_service(runtime, max_rpc_workers=1) as (service, stub): + first_result: dict[str, Any] = {} + first_call = stub.Generate(encode_generate_request((_request("first"),), runtime.identity), timeout=5) + first_thread = threading.Thread(target=_consume, args=(first_call, first_result)) + first_thread.start() + try: + assert entered.wait(timeout=2) + with pytest.raises(grpc.RpcError) as excess_error: + list(stub.Generate(encode_generate_request((_request("excess"),), runtime.identity), timeout=2)) + assert excess_error.value.code() == grpc.StatusCode.RESOURCE_EXHAUSTED + assert runtime.calls == ["first"] + assert service.metrics().accepted_batches == 1 + finally: + release.set() + first_thread.join(timeout=5) + assert not first_thread.is_alive() + assert "code" not in first_result + + +@pytest.mark.parametrize( + ("cancel_explicitly", "expected_code"), + ((True, grpc.StatusCode.CANCELLED), (False, grpc.StatusCode.DEADLINE_EXCEEDED)), + ids=("cancel", "deadline"), +) +def test_cancelled_or_expired_queued_rpc_never_executes( + cancel_explicitly: bool, + expected_code: Any, +) -> None: + first_entered = threading.Event() + release_first = threading.Event() + + def behavior(requests: tuple[ReasonerFeatureRequest, ...]) -> ReasonerFeatureBatch: + if requests[0].sample_key == "first": + first_entered.set() + if not release_first.wait(timeout=5): + raise AssertionError("test did not release the first RPC") + return _features(requests) + + runtime = _FakeRuntime(behavior) + with _running_service(runtime, max_rpc_workers=2) as (service, stub): + first_result: dict[str, Any] = {} + first_call = stub.Generate(encode_generate_request((_request("first"),), runtime.identity), timeout=5) + first_thread = threading.Thread(target=_consume, args=(first_call, first_result)) + first_thread.start() + + second_result: dict[str, Any] = {} + second_thread: threading.Thread | None = None + try: + assert first_entered.wait(timeout=2) + second_timeout = 5.0 if cancel_explicitly else 1.0 + second_call = stub.Generate( + encode_generate_request((_request("second"),), runtime.identity), + timeout=second_timeout, + ) + second_thread = threading.Thread(target=_consume, args=(second_call, second_result)) + second_thread.start() + _wait_until(lambda: service.metrics().accepted_batches == 2) + if cancel_explicitly: + assert second_call.cancel() + second_thread.join(timeout=3) + assert not second_thread.is_alive() + assert second_result["code"] == expected_code + + # Wait for the server handler, not just the client, to observe the + # cancellation and relinquish its token reservation while the first + # request still owns the execution lock. + _wait_until(lambda: service.metrics().queued_tokens == 1) + assert runtime.calls == ["first"] + finally: + release_first.set() + first_thread.join(timeout=5) + if second_thread is not None: + second_thread.join(timeout=5) + + assert not first_thread.is_alive() + assert "code" not in first_result + assert runtime.calls == ["first"] + _wait_until(lambda: service.metrics().queued_tokens == 0) + + +def test_oom_marks_unhealthy_before_queued_request_can_execute() -> None: + first_entered = threading.Event() + release_oom = threading.Event() + + def behavior(requests: tuple[ReasonerFeatureRequest, ...]) -> ReasonerFeatureBatch: + if requests[0].sample_key == "oom": + first_entered.set() + if not release_oom.wait(timeout=5): + raise AssertionError("test did not release the OOM RPC") + raise torch.cuda.OutOfMemoryError("synthetic CUDA OOM") + return _features(requests) + + runtime = _FakeRuntime(behavior) + with _running_service(runtime, max_rpc_workers=2) as (service, stub): + oom_result: dict[str, Any] = {} + oom_call = stub.Generate(encode_generate_request((_request("oom"),), runtime.identity), timeout=5) + oom_thread = threading.Thread(target=_consume, args=(oom_call, oom_result)) + oom_thread.start() + + queued_result: dict[str, Any] = {} + queued_thread: threading.Thread | None = None + try: + assert first_entered.wait(timeout=2) + queued_call = stub.Generate(encode_generate_request((_request("queued"),), runtime.identity), timeout=5) + queued_thread = threading.Thread(target=_consume, args=(queued_call, queued_result)) + queued_thread.start() + _wait_until(lambda: service.metrics().accepted_batches == 2) + release_oom.set() + oom_thread.join(timeout=3) + queued_thread.join(timeout=3) + finally: + release_oom.set() + oom_thread.join(timeout=5) + if queued_thread is not None: + queued_thread.join(timeout=5) + + assert not oom_thread.is_alive() + assert queued_thread is not None and not queued_thread.is_alive() + assert oom_result["code"] == grpc.StatusCode.RESOURCE_EXHAUSTED + assert queued_result["code"] == grpc.StatusCode.UNAVAILABLE + assert runtime.calls == ["oom"] + metrics = service.metrics() + assert not metrics.healthy + assert metrics.failed_batches == 2 + assert metrics.queued_tokens == 0 + + +def test_fatal_runtime_invariant_marks_unhealthy_before_queued_request_can_execute() -> None: + first_entered = threading.Event() + release_failure = threading.Event() + + def behavior(requests: tuple[ReasonerFeatureRequest, ...]) -> ReasonerFeatureBatch: + if requests[0].sample_key == "fatal": + first_entered.set() + if not release_failure.wait(timeout=5): + raise AssertionError("test did not release the fatal runtime failure") + raise ReasonerRuntimeInvariantError("synthetic output signature mismatch") + return _features(requests) + + runtime = _FakeRuntime(behavior) + with _running_service(runtime, max_rpc_workers=2) as (service, stub): + fatal_result: dict[str, Any] = {} + fatal_call = stub.Generate(encode_generate_request((_request("fatal"),), runtime.identity), timeout=5) + fatal_thread = threading.Thread(target=_consume, args=(fatal_call, fatal_result)) + fatal_thread.start() + + queued_result: dict[str, Any] = {} + queued_thread: threading.Thread | None = None + try: + assert first_entered.wait(timeout=2) + queued_call = stub.Generate(encode_generate_request((_request("queued"),), runtime.identity), timeout=5) + queued_thread = threading.Thread(target=_consume, args=(queued_call, queued_result)) + queued_thread.start() + _wait_until(lambda: service.metrics().accepted_batches == 2) + release_failure.set() + fatal_thread.join(timeout=3) + queued_thread.join(timeout=3) + finally: + release_failure.set() + fatal_thread.join(timeout=5) + if queued_thread is not None: + queued_thread.join(timeout=5) + + assert not fatal_thread.is_alive() + assert queued_thread is not None and not queued_thread.is_alive() + assert fatal_result["code"] == grpc.StatusCode.INTERNAL + assert "marked unhealthy" in fatal_result["details"] + assert queued_result["code"] == grpc.StatusCode.UNAVAILABLE + assert runtime.calls == ["fatal"] + metrics = service.metrics() + assert not metrics.healthy + assert metrics.failed_batches == 2 + assert metrics.queued_tokens == 0 + + info = stub.GetInfo( + reasoner_pb2.GetInfoRequest(protocol_version=REASONER_FEATURE_PROTOCOL_VERSION), + timeout=2, + ) + assert not info.ready + with pytest.raises(grpc.RpcError) as subsequent_error: + list(stub.Generate(encode_generate_request((_request("subsequent"),), runtime.identity), timeout=2)) + assert subsequent_error.value.code() == grpc.StatusCode.UNAVAILABLE + assert runtime.calls == ["fatal"] + + +def test_stream_serialization_failure_is_internal_and_releases_admission( + monkeypatch: pytest.MonkeyPatch, +) -> None: + import cosmos_framework.model.generator.reasoner_remote_server as server_module + + def broken_stream(*_args: object, **_kwargs: object) -> Iterator[Any]: + yield from () + raise ValueError("synthetic stream failure") + + monkeypatch.setattr(server_module, "iter_feature_stream", broken_stream) + runtime = _FakeRuntime() + with _running_service(runtime) as (service, stub): + with pytest.raises(grpc.RpcError) as stream_error: + list(stub.Generate(encode_generate_request((_request("stream"),), runtime.identity), timeout=2)) + + assert stream_error.value.code() == grpc.StatusCode.INTERNAL + assert "synthetic stream failure" in stream_error.value.details() + metrics = service.metrics() + assert metrics.accepted_batches == 1 + assert metrics.completed_batches == 0 + assert metrics.failed_batches == 1 + assert metrics.queued_tokens == 0 diff --git a/cosmos_framework/model/generator/reasoner_remote_test.py b/cosmos_framework/model/generator/reasoner_remote_test.py new file mode 100644 index 000000000..4eb3f03c1 --- /dev/null +++ b/cosmos_framework/model/generator/reasoner_remote_test.py @@ -0,0 +1,442 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +from __future__ import annotations + +import threading +from collections.abc import Sequence +from typing import Any + +import pytest +import torch + +import cosmos_framework.model.generator.reasoner_remote as reasoner_remote_module +from cosmos_framework.model.generator.reasoner_features import ( + ReasonerFeatureBatch, + ReasonerFeatureIdentity, + ReasonerFeatureRequest, + ReasonerFeatureSignature, +) +from cosmos_framework.model.generator.reasoner_remote import ( + REASONER_FEATURE_PROTOCOL_VERSION, + RemoteReasonerFeatureProvider, + RemoteReasonerRPCError, + decode_feature_stream, + decode_tensor_payload, + encode_generate_request, + encode_identity, + encode_signature, + encode_tensor_payload, + iter_feature_stream, +) +from cosmos_framework.protos.reasoner_features.v1 import reasoner_features_pb2 as reasoner_pb2 + +pytestmark = [pytest.mark.level(0), pytest.mark.gpus(0)] + + +def _identity() -> ReasonerFeatureIdentity: + return ReasonerFeatureIdentity( + reasoner="reasoner-digest", + tokenizer="tokenizer-digest", + framing="framing-digest", + ) + + +def _signature() -> ReasonerFeatureSignature: + return ReasonerFeatureSignature( + num_layers=2, + num_kv_heads=2, + head_dim=3, + dtype=torch.bfloat16, + ) + + +def _request(sample_key: str, token_ids: Sequence[int]) -> ReasonerFeatureRequest: + tokens = torch.tensor(token_ids, dtype=torch.int64) + return ReasonerFeatureRequest( + sample_key=sample_key, + token_ids=tokens, + position_ids=torch.arange(tokens.numel(), dtype=torch.int64), + causal_offsets=torch.tensor([0, tokens.numel()], dtype=torch.int64), + fingerprint=f"{sample_key}-fingerprint", + ) + + +def _feature_batch( + requests: Sequence[ReasonerFeatureRequest], + *, + signature: ReasonerFeatureSignature | None = None, +) -> ReasonerFeatureBatch: + signature = signature or _signature() + total_tokens = sum(request.token_ids.numel() for request in requests) + shape = (total_tokens, signature.num_kv_heads, signature.head_dim) + numel = total_tokens * signature.num_kv_heads * signature.head_dim + + def values(offset: int) -> torch.Tensor: + return (torch.arange(numel, dtype=torch.float32) + offset).reshape(shape).to(signature.dtype) + + offsets = [0] + for request in requests: + offsets.append(offsets[-1] + request.token_ids.numel()) + return ReasonerFeatureBatch( + cross_k=tuple(values(100 * layer_idx) for layer_idx in range(signature.num_layers)), + cross_v=tuple(values(1000 + 100 * layer_idx) for layer_idx in range(signature.num_layers)), + causal_offsets=torch.tensor(offsets, dtype=torch.int64), + fingerprints=tuple(request.fingerprint for request in requests), + ) + + +def _assert_feature_batches_equal(actual: ReasonerFeatureBatch, expected: ReasonerFeatureBatch) -> None: + assert actual.fingerprints == expected.fingerprints + assert torch.equal(actual.causal_offsets, expected.causal_offsets) + assert len(actual.cross_k) == len(expected.cross_k) + for actual_k, expected_k, actual_v, expected_v in zip( + actual.cross_k, + expected.cross_k, + actual.cross_v, + expected.cross_v, + ): + assert actual_k.dtype == expected_k.dtype + assert actual_v.dtype == expected_v.dtype + assert torch.equal(actual_k, expected_k) + assert torch.equal(actual_v, expected_v) + + +def _stream( + requests: Sequence[ReasonerFeatureRequest], + *, + max_chunk_bytes: int = 7, +) -> tuple[Any, list[Any], ReasonerFeatureBatch]: + identity = _identity() + signature = _signature() + request = encode_generate_request(requests, identity) + features = _feature_batch(requests, signature=signature) + responses = list( + iter_feature_stream( + features, + request_id=request.request_id, + identity=identity, + signature=signature, + service_instance_id="test-replica", + max_chunk_bytes=max_chunk_bytes, + queue_ns=11, + compute_ns=22, + device_to_host_ns=33, + ) + ) + return request, responses, features + + +def _clone_responses(responses: Sequence[Any]) -> list[Any]: + clones = [] + for response in responses: + clone = reasoner_pb2.GenerateResponse() + clone.CopyFrom(response) + clones.append(clone) + return clones + + +def _service_info( + *, + identity: ReasonerFeatureIdentity | None = None, + signature: ReasonerFeatureSignature | None = None, + max_chunk_bytes: int = 7, +) -> Any: + return reasoner_pb2.GetInfoResponse( + protocol_version=REASONER_FEATURE_PROTOCOL_VERSION, + service_instance_id="fake-replica", + identity=encode_identity(identity or _identity()), + signature=encode_signature(signature or _signature()), + capabilities=( + reasoner_pb2.Capability(name="server_streaming", version=1), + reasoner_pb2.Capability(name="strict_fingerprint", version=1), + reasoner_pb2.Capability(name="ordered_requests", version=1), + reasoner_pb2.Capability(name="single_document", version=1), + ), + max_batch_tokens=64, + max_queued_tokens=128, + max_chunk_bytes=max_chunk_bytes, + ready=True, + max_requests=4, + ) + + +def test_tensor_payload_roundtrip_preserves_noncontiguous_bfloat16_bits() -> None: + source = torch.tensor( + [[1.0, -2.5, 3.25], [4.5, 5.75, -6.0]], + dtype=torch.float32, + ).to(torch.bfloat16) + source = source.t() + expected = source.clone() + + decoded = decode_tensor_payload(encode_tensor_payload(source), name="bf16") + source.zero_() + + assert decoded.dtype == torch.bfloat16 + assert decoded.shape == expected.shape + assert decoded.is_contiguous() + assert torch.equal(decoded, expected) + + +def test_multichunk_feature_stream_roundtrip_preserves_bfloat16_kv_and_timing() -> None: + requests = (_request("first", [1, 2, 3]), _request("second", [4, 5])) + request, responses, expected = _stream(requests) + tensor_chunks = [response.tensor_chunk for response in responses if response.HasField("tensor_chunk")] + + assert len(tensor_chunks) > 2 * _signature().num_layers + assert all(0 < len(chunk.data) <= 7 for chunk in tensor_chunks) + + actual, timing = decode_feature_stream( + responses, + request=request, + expected_identity=_identity(), + expected_signature=_signature(), + max_chunk_bytes=7, + ) + + _assert_feature_batches_equal(actual, expected) + assert timing.queue_ns == 11 + assert timing.compute_ns == 22 + assert timing.device_to_host_ns == 33 + assert timing.serialization_ns >= 0 + + +def test_stream_serialization_timing_excludes_consumer_backpressure(monkeypatch: pytest.MonkeyPatch) -> None: + fake_clock_ns = 0 + monkeypatch.setattr(reasoner_remote_module.time, "perf_counter_ns", lambda: fake_clock_ns) + requests = (_request("sample", [1, 2, 3]),) + request = encode_generate_request(requests, _identity()) + stream = iter_feature_stream( + _feature_batch(requests), + request_id=request.request_id, + identity=_identity(), + signature=_signature(), + service_instance_id="timing-test", + max_chunk_bytes=7, + queue_ns=1, + compute_ns=2, + device_to_host_ns=3, + ) + + responses = [] + for response in stream: + responses.append(response) + fake_clock_ns += 1_000_000_000 + + assert responses[-1].trailer.timing.serialization_ns == 0 + + +def test_feature_stream_rejects_corrupt_chunk_checksum() -> None: + request, responses, _ = _stream((_request("sample", [1, 2, 3]),)) + damaged = _clone_responses(responses) + chunk = next(response.tensor_chunk for response in damaged if response.HasField("tensor_chunk")) + chunk.data = bytes([chunk.data[0] ^ 0xFF]) + chunk.data[1:] + + with pytest.raises(ValueError, match="checksum mismatch"): + decode_feature_stream( + damaged, + request=request, + expected_identity=_identity(), + expected_signature=_signature(), + max_chunk_bytes=7, + ) + + +def test_feature_stream_rejects_noncanonical_tensor_order() -> None: + request, responses, _ = _stream((_request("sample", [1, 2, 3]),)) + damaged = _clone_responses(responses) + first_chunk = next(response.tensor_chunk for response in damaged if response.HasField("tensor_chunk")) + first_chunk.kind = reasoner_pb2.TENSOR_KIND_CROSS_V + + with pytest.raises(ValueError, match="tensor order"): + decode_feature_stream( + damaged, + request=request, + expected_identity=_identity(), + expected_signature=_signature(), + max_chunk_bytes=7, + ) + + +class _RetryingFakeTransport: + def __init__(self, *, info: Any | None = None) -> None: + self.info = info or _service_info() + self.get_info_timeouts: list[float] = [] + self.generate_timeouts: list[float] = [] + self.request_ids: list[str] = [] + self.worker_thread_ids: list[int] = [] + self.second_attempt_entered = threading.Event() + self.release_second_attempt = threading.Event() + self.closed = False + + def get_info(self, *, timeout_s: float) -> Any: + self.get_info_timeouts.append(timeout_s) + return self.info + + def generate(self, request: Any, *, timeout_s: float) -> list[Any]: + self.generate_timeouts.append(timeout_s) + self.request_ids.append(str(request.request_id)) + self.worker_thread_ids.append(threading.get_ident()) + if len(self.request_ids) == 1: + raise RemoteReasonerRPCError("UNAVAILABLE", "transient test failure", retryable=True) + + self.second_attempt_entered.set() + if not self.release_second_attempt.wait(timeout=5): + raise AssertionError("test did not release the fake Reasoner transport") + requests = tuple( + _request(str(item.sample_key), range(int(item.token_ids.shape[0]))) for item in request.requests + ) + features = _feature_batch(requests) + features = ReasonerFeatureBatch( + cross_k=features.cross_k, + cross_v=features.cross_v, + causal_offsets=features.causal_offsets, + fingerprints=tuple(str(item.fingerprint) for item in request.requests), + ) + return list( + iter_feature_stream( + features, + request_id=request.request_id, + identity=_identity(), + signature=_signature(), + service_instance_id="fake-replica", + max_chunk_bytes=7, + queue_ns=1, + compute_ns=2, + device_to_host_ns=3, + ) + ) + + def close(self) -> None: + self.closed = True + + +def test_remote_provider_handshake_async_submit_and_retry() -> None: + transport = _RetryingFakeTransport() + provider = RemoteReasonerFeatureProvider( + "unused.test:1234", + expected_identity=_identity(), + expected_dtype=torch.bfloat16, + connect_timeout_s=1.25, + request_timeout_s=5.0, + request_max_retries=1, + retry_backoff_s=0, + transport=transport, + ) + requests = (_request("first", [8, 9, 10]), _request("second", [11, 12])) + main_thread_id = threading.get_ident() + + try: + future = provider.submit(requests) + assert transport.second_attempt_entered.wait(timeout=2) + assert not future.done() + transport.release_second_attempt.set() + actual = future.result(timeout=5) + finally: + transport.release_second_attempt.set() + provider.close() + + _assert_feature_batches_equal(actual, _feature_batch(requests)) + assert transport.get_info_timeouts == [1.25] + assert len(transport.generate_timeouts) == 2 + assert all(timeout > 0 for timeout in transport.generate_timeouts) + assert transport.generate_timeouts[1] <= transport.generate_timeouts[0] + assert len(set(transport.request_ids)) == 1 + assert all(thread_id != main_thread_id for thread_id in transport.worker_thread_ids) + assert provider.last_server_timing is not None + assert provider.last_server_timing.compute_ns == 2 + assert transport.closed + + +def test_remote_provider_rejects_handshake_identity_mismatch_and_closes_transport() -> None: + wrong_identity = ReasonerFeatureIdentity(reasoner="wrong", tokenizer="tokenizer", framing="framing") + transport = _RetryingFakeTransport(info=_service_info(identity=wrong_identity)) + + with pytest.raises(ValueError, match="identity mismatch"): + RemoteReasonerFeatureProvider( + "unused.test:1234", + expected_identity=_identity(), + expected_dtype=torch.bfloat16, + transport=transport, + ) + + assert transport.closed + + +class _FakeRuntime: + def __init__(self) -> None: + self.identity = _identity() + self.signature = _signature() + self.max_requests = 4 + self.max_total_tokens = 64 + self.device = torch.device("cpu") + + def execute(self, requests: Sequence[ReasonerFeatureRequest]) -> ReasonerFeatureBatch: + return _feature_batch(requests, signature=self.signature) + + +def test_service_holds_token_admission_until_stream_is_closed() -> None: + pytest.importorskip("grpc") + from cosmos_framework.model.generator.reasoner_remote_server import ReasonerFeatureService + + runtime = _FakeRuntime() + service = ReasonerFeatureService( + runtime, # type: ignore[arg-type] + max_queued_tokens=128, + max_chunk_bytes=7, + ) + request = _request("sample", [1, 2, 3]) + stream = service.Generate(encode_generate_request((request,), runtime.identity), context=None) + + first_response = next(stream) + assert first_response.HasField("header") + assert service.metrics().queued_tokens == 3 + + stream.close() + metrics = service.metrics() + assert metrics.queued_tokens == 0 + assert metrics.accepted_batches == 1 + assert metrics.completed_batches == 0 + + +def test_remote_provider_end_to_end_through_in_process_grpc_server() -> None: + pytest.importorskip("grpc") + from cosmos_framework.model.generator.reasoner_remote_server import ( + ReasonerFeatureService, + create_reasoner_grpc_server, + ) + + runtime = _FakeRuntime() + service = ReasonerFeatureService( + runtime, # type: ignore[arg-type] + max_queued_tokens=128, + max_chunk_bytes=7, + service_instance_id="in-process-replica", + ) + server, port = create_reasoner_grpc_server(service, address="127.0.0.1:0", max_rpc_workers=2) + server.start() + provider: RemoteReasonerFeatureProvider | None = None + requests = (_request("first", [1, 2, 3]), _request("second", [4, 5])) + try: + provider = RemoteReasonerFeatureProvider( + f"127.0.0.1:{port}", + expected_identity=runtime.identity, + expected_dtype=runtime.signature.dtype, + connect_timeout_s=5, + request_timeout_s=5, + request_max_retries=0, + ) + actual = provider.submit(requests).result(timeout=5) + finally: + if provider is not None: + provider.close() + server.stop(grace=0).wait(timeout=5) + + _assert_feature_batches_equal(actual, _feature_batch(requests)) + metrics = service.metrics() + assert metrics.accepted_batches == 1 + assert metrics.completed_batches == 1 + assert metrics.rejected_batches == 0 + assert metrics.failed_batches == 0 + assert metrics.queued_tokens == 0 + assert metrics.healthy diff --git a/cosmos_framework/model/generator/reasoner_runtime.py b/cosmos_framework/model/generator/reasoner_runtime.py new file mode 100644 index 000000000..947596d21 --- /dev/null +++ b/cosmos_framework/model/generator/reasoner_runtime.py @@ -0,0 +1,321 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +"""Reusable frozen-Reasoner runtime for offline extraction and remote serving. + +The runtime owns one complete, unsharded Reasoner replica. It deliberately +does not initialize a distributed process group: scale-out is achieved by +starting independent one-process/one-GPU replicas behind an external load +balancer. Calls are serialized because the first implementation executes +variable-length requests one at a time inside +``extract_reasoner_feature_batch``. +""" + +from __future__ import annotations + +import copy +import threading +import time +from collections.abc import Sequence +from contextlib import contextmanager +from dataclasses import dataclass +from pathlib import Path +from typing import Iterator, Literal + +import torch +from omegaconf import open_dict + +from cosmos_framework.checkpoint.reasoner_only import load_reasoner_only_dcp +from cosmos_framework.model.generator.reasoner_features import ( + ReasonerFeatureBatch, + ReasonerFeatureIdentity, + ReasonerFeatureRequest, + ReasonerFeatureSignature, + compute_reasoner_feature_fingerprint, + extract_reasoner_feature_batch, +) +from cosmos_framework.utils.lazy_config import instantiate as lazy_instantiate + + +@dataclass(frozen=True) +class ReasonerRuntimeSpec: + """Immutable startup and admission limits for one Reasoner replica.""" + + checkpoint: str | Path + checkpoint_source: Literal["regular", "ema"] + device: torch.device | str + dtype: torch.dtype + identity: ReasonerFeatureIdentity + max_requests: int = 64 + max_total_tokens: int = 4_096 + + def __post_init__(self) -> None: + if self.checkpoint_source not in {"regular", "ema"}: + raise ValueError(f"Unsupported checkpoint_source={self.checkpoint_source!r}") + if self.dtype not in (torch.bfloat16, torch.float16, torch.float32): + raise TypeError(f"Unsupported Reasoner runtime dtype: {self.dtype}") + for name in ("max_requests", "max_total_tokens"): + value = getattr(self, name) + if not isinstance(value, int) or isinstance(value, bool) or value <= 0: + raise ValueError(f"{name} must be a positive integer, got {value!r}") + + +@dataclass(frozen=True) +class ReasonerRuntimeLoadStats: + construction_seconds: float + checkpoint_load_seconds: float + allocated_bytes: int + peak_allocated_bytes: int + reserved_bytes: int + peak_reserved_bytes: int + + +class ReasonerRuntimeInvariantError(RuntimeError): + """A permanent loaded-runtime/output contract violation requiring restart.""" + + +def prepare_reasoner_model_config(config: object) -> object: + """Copy the Nano LM lazy config and remove Generator/vision construction.""" + + model = getattr(getattr(config, "model"), "config") + model_instance = copy.deepcopy(model.vlm_config.model_instance) + if model_instance is None: + raise ValueError("Reasoner runtime requires model.config.vlm_config.model_instance") + nested_config = model_instance["config"] + if isinstance(nested_config, dict): + nested_config.update( + include_gen_pathway=False, + include_und_pathway=True, + include_visual=False, + ) + else: + with open_dict(nested_config): + nested_config.include_gen_pathway = False + nested_config.include_und_pathway = True + nested_config.include_visual = False + return model_instance + + +@contextmanager +def _temporary_default_dtype(dtype: torch.dtype) -> Iterator[None]: + """Set the process-global construction dtype during single-threaded startup.""" + + previous = torch.get_default_dtype() + torch.set_default_dtype(dtype) + try: + yield + finally: + torch.set_default_dtype(previous) + + +def build_reasoner(config: object, *, device: torch.device, dtype: torch.dtype) -> torch.nn.Module: + """Construct a materialized Reasoner replica directly in its compute dtype. + + Direct construction intentionally avoids meta ``to_empty``: HuggingFace + rotary buffers are non-persistent and absent from DCP, so constructing on + the final device initializes them before the strict checkpoint load. + """ + + model_instance = prepare_reasoner_model_config(config) + with _temporary_default_dtype(dtype), torch.device(device): + reasoner = lazy_instantiate(model_instance) + generation_parameters = [name for name, _ in reasoner.named_parameters() if "moe_gen" in name] + if generation_parameters: + raise RuntimeError(f"Reasoner-only construction retained Generator parameters: {generation_parameters}") + reasoner.requires_grad_(False) + reasoner.eval() + return reasoner + + +def _synchronize(device: torch.device) -> None: + if device.type == "cuda": + torch.cuda.synchronize(device) + + +def _memory_stats(device: torch.device) -> tuple[int, int, int, int]: + if device.type != "cuda": + return 0, 0, 0, 0 + return ( + torch.cuda.memory_allocated(device), + torch.cuda.max_memory_allocated(device), + torch.cuda.memory_reserved(device), + torch.cuda.max_memory_reserved(device), + ) + + +class ReasonerFeatureRuntime: + """Loaded Reasoner plus strict request validation and serialized execution.""" + + def __init__( + self, + reasoner: torch.nn.Module, + *, + identity: ReasonerFeatureIdentity, + max_requests: int = 64, + max_total_tokens: int = 4_096, + load_stats: ReasonerRuntimeLoadStats | None = None, + ) -> None: + if reasoner.training: + raise ValueError("ReasonerFeatureRuntime requires reasoner.eval()") + if any(parameter.requires_grad for parameter in reasoner.parameters()): + raise ValueError("ReasonerFeatureRuntime requires every Reasoner parameter to be frozen") + if max_requests <= 0 or max_total_tokens <= 0: + raise ValueError("Reasoner runtime request and token limits must be positive") + try: + model = reasoner.model + layers = model.layers + embedding = model.embed_tokens + attentions = tuple(layer.self_attn for layer in layers) + first_attention = attentions[0] + except (AttributeError, IndexError) as error: + raise TypeError( + "Expected a *TextForCausalLM wrapper with model.embed_tokens and non-empty model.layers" + ) from error + if not getattr(model, "include_und_pathway", True): + raise ValueError("Reasoner runtime requires include_und_pathway=True") + if getattr(model, "include_gen_pathway", False): + raise ValueError("Reasoner runtime must not retain the Generator pathway") + if not callable(getattr(model, "reasoner_forward", None)): + raise TypeError("Reasoner runtime requires model.reasoner_forward to be callable") + unsupported_layers = [ + layer_idx + for layer_idx, attention in enumerate(attentions) + if getattr(attention, "k_norm_und_for_gen", None) is not None + ] + if unsupported_layers: + raise NotImplementedError( + "UND-only Reasoner execution does not support generator-specific K normalization; " + f"affected layers={unsupported_layers}" + ) + + head_dim = int(getattr(first_attention, "head_dim")) + key_projection = getattr(first_attention, "k_proj") + if key_projection.out_features % head_dim: + raise ValueError("Reasoner key projection width is not divisible by head_dim") + num_kv_heads = int(key_projection.out_features // head_dim) + self.reasoner = reasoner + self.identity = identity + self.max_requests = int(max_requests) + self.max_total_tokens = int(max_total_tokens) + self.signature = ReasonerFeatureSignature( + num_layers=len(layers), + num_kv_heads=num_kv_heads, + head_dim=head_dim, + dtype=embedding.weight.dtype, + ) + self.vocab_size = int(embedding.num_embeddings) + self.device = embedding.weight.device + self.load_stats = load_stats + self._execution_lock = threading.Lock() + + @classmethod + def load(cls, config: object, spec: ReasonerRuntimeSpec) -> ReasonerFeatureRuntime: + """Construct and strictly restore one replica before reporting it ready.""" + + device = torch.device(spec.device) + if device.type == "cuda": + if not torch.cuda.is_available(): + raise RuntimeError("CUDA Reasoner runtime requested, but CUDA is unavailable") + if device.index is None: + device = torch.device("cuda", torch.cuda.current_device()) + if device.index < 0 or device.index >= torch.cuda.device_count(): + raise ValueError(f"Reasoner runtime device {device} is not visible") + torch.cuda.set_device(device) + + _synchronize(device) + construction_started = time.perf_counter() + reasoner = build_reasoner(config, device=device, dtype=spec.dtype) + runtime = cls( + reasoner, + identity=spec.identity, + max_requests=spec.max_requests, + max_total_tokens=spec.max_total_tokens, + ) + _synchronize(device) + construction_seconds = time.perf_counter() - construction_started + + checkpoint_started = time.perf_counter() + load_reasoner_only_dcp(reasoner, spec.checkpoint, source=spec.checkpoint_source) + _synchronize(device) + checkpoint_load_seconds = time.perf_counter() - checkpoint_started + allocated, peak_allocated, reserved, peak_reserved = _memory_stats(device) + runtime.load_stats = ReasonerRuntimeLoadStats( + construction_seconds=construction_seconds, + checkpoint_load_seconds=checkpoint_load_seconds, + allocated_bytes=allocated, + peak_allocated_bytes=peak_allocated, + reserved_bytes=reserved, + peak_reserved_bytes=peak_reserved, + ) + return runtime + + def estimated_response_bytes(self, requests: Sequence[ReasonerFeatureRequest]) -> int: + """Return the canonical K/V payload size before transport framing.""" + + total_tokens = sum(request.token_ids.numel() for request in requests) + return ( + total_tokens + * self.signature.num_layers + * 2 + * self.signature.num_kv_heads + * self.signature.head_dim + * self.signature.dtype.itemsize + ) + + def _validate_requests(self, requests: Sequence[ReasonerFeatureRequest]) -> tuple[ReasonerFeatureRequest, ...]: + requests = tuple(requests) + if not requests: + raise ValueError("At least one Reasoner feature request is required") + if len(requests) > self.max_requests: + raise ValueError(f"Reasoner request count {len(requests)} exceeds limit {self.max_requests}") + total_tokens = sum(request.token_ids.numel() for request in requests) + if total_tokens > self.max_total_tokens: + raise ValueError(f"Reasoner token count {total_tokens} exceeds limit {self.max_total_tokens}") + for request in requests: + if request.token_ids.numel() == 0: + raise ValueError(f"Reasoner request {request.sample_key!r} contains no tokens") + token_ids = request.token_ids.detach().to(device="cpu", dtype=torch.int64) + minimum = int(token_ids.min()) + maximum = int(token_ids.max()) + if minimum < 0 or maximum >= self.vocab_size: + raise ValueError( + f"Reasoner request {request.sample_key!r} has token range [{minimum}, {maximum}] " + f"outside vocabulary [0, {self.vocab_size})" + ) + expected_fingerprint = compute_reasoner_feature_fingerprint( + request.token_ids, + request.position_ids, + request.causal_offsets, + identity=self.identity, + ) + if request.fingerprint != expected_fingerprint: + raise ValueError( + f"Reasoner request fingerprint mismatch for {request.sample_key!r}: " + f"request={request.fingerprint!r}, expected={expected_fingerprint!r}" + ) + return requests + + @torch.inference_mode() + def execute(self, requests: Sequence[ReasonerFeatureRequest]) -> ReasonerFeatureBatch: + """Validate and execute requests in order on the owned Reasoner replica.""" + + validated = self._validate_requests(requests) + with self._execution_lock: + features = extract_reasoner_feature_batch(self.reasoner, validated) + actual = ReasonerFeatureSignature( + num_layers=features.num_layers, + num_kv_heads=features.layer(0).num_kv_heads, + head_dim=features.layer(0).head_dim, + dtype=features.layer(0).cross_k.dtype, + ) + if actual != self.signature: + raise ReasonerRuntimeInvariantError( + f"Reasoner runtime output signature changed: output={actual!r}, ready={self.signature!r}" + ) + expected_fingerprints = tuple(request.fingerprint for request in validated) + if features.fingerprints != expected_fingerprints: + raise ReasonerRuntimeInvariantError( + "Reasoner runtime output fingerprints do not preserve request order: " + f"output={features.fingerprints!r}, expected={expected_fingerprints!r}" + ) + return features diff --git a/cosmos_framework/model/generator/reasoner_runtime_test.py b/cosmos_framework/model/generator/reasoner_runtime_test.py new file mode 100644 index 000000000..3732026fc --- /dev/null +++ b/cosmos_framework/model/generator/reasoner_runtime_test.py @@ -0,0 +1,292 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +from __future__ import annotations + +from collections.abc import Sequence + +import pytest +import torch +from torch import nn + +import cosmos_framework.model.generator.reasoner_runtime as reasoner_runtime_module +from cosmos_framework.model.generator.reasoner_features import ( + ReasonerFeatureBatch, + ReasonerFeatureIdentity, + ReasonerFeatureRequest, + ReasonerFeatureSignature, + compute_reasoner_feature_fingerprint, +) +from cosmos_framework.model.generator.reasoner_runtime import ( + ReasonerFeatureRuntime, + ReasonerRuntimeInvariantError, + ReasonerRuntimeSpec, +) + +pytestmark = [pytest.mark.level(0), pytest.mark.gpus(0)] + + +class _FakeAttention(nn.Module): + def __init__(self) -> None: + super().__init__() + self.head_dim = 2 + self.k_proj = nn.Linear(4, 4, bias=False) + + +class _FakeLayer(nn.Module): + def __init__(self) -> None: + super().__init__() + self.self_attn = _FakeAttention() + + +class _FakeReasonerModel(nn.Module): + def __init__(self) -> None: + super().__init__() + self.include_und_pathway = True + self.include_gen_pathway = False + self.embed_tokens = nn.Embedding(32, 4) + self.layers = nn.ModuleList([_FakeLayer(), _FakeLayer()]) + + def reasoner_forward(self, **_kwargs: object) -> None: + raise AssertionError("fake reasoner_forward should not execute in runtime unit tests") + + +class _FakeReasoner(nn.Module): + def __init__(self) -> None: + super().__init__() + self.model = _FakeReasonerModel() + self.requires_grad_(False) + self.eval() + + +def _identity() -> ReasonerFeatureIdentity: + return ReasonerFeatureIdentity( + reasoner="reasoner-digest", + tokenizer="tokenizer-digest", + framing="framing-digest", + ) + + +def _request( + sample_key: str, + token_ids: Sequence[int], + *, + identity: ReasonerFeatureIdentity, + fingerprint: str | None = None, +) -> ReasonerFeatureRequest: + tokens = torch.tensor(token_ids, dtype=torch.int64) + positions = torch.arange(tokens.numel(), dtype=torch.int64) + offsets = torch.tensor([0, tokens.numel()], dtype=torch.int64) + if fingerprint is None: + fingerprint = compute_reasoner_feature_fingerprint( + tokens, + positions, + offsets, + identity=identity, + ) + return ReasonerFeatureRequest( + sample_key=sample_key, + token_ids=tokens, + position_ids=positions, + causal_offsets=offsets, + fingerprint=fingerprint, + ) + + +def _runtime(*, max_requests: int = 4, max_total_tokens: int = 16) -> ReasonerFeatureRuntime: + return ReasonerFeatureRuntime( + _FakeReasoner(), + identity=_identity(), + max_requests=max_requests, + max_total_tokens=max_total_tokens, + ) + + +def _features_for(requests: Sequence[ReasonerFeatureRequest]) -> ReasonerFeatureBatch: + total_tokens = sum(request.token_ids.numel() for request in requests) + offsets = [0] + for request in requests: + offsets.append(offsets[-1] + request.token_ids.numel()) + return ReasonerFeatureBatch( + cross_k=tuple(torch.full((total_tokens, 2, 2), layer_idx, dtype=torch.float32) for layer_idx in range(2)), + cross_v=tuple(torch.full((total_tokens, 2, 2), layer_idx + 10, dtype=torch.float32) for layer_idx in range(2)), + causal_offsets=torch.tensor(offsets, dtype=torch.int64), + fingerprints=tuple(request.fingerprint for request in requests), + ) + + +def test_runtime_exposes_reasoner_feature_signature() -> None: + runtime = _runtime() + + assert runtime.signature == ReasonerFeatureSignature( + num_layers=2, + num_kv_heads=2, + head_dim=2, + dtype=torch.float32, + ) + assert runtime.vocab_size == 32 + assert runtime.device == torch.device("cpu") + + +def test_runtime_rejects_missing_reasoner_forward_at_startup() -> None: + reasoner = _FakeReasoner() + reasoner.model.reasoner_forward = None # type: ignore[method-assign] + + with pytest.raises(TypeError, match="model.reasoner_forward to be callable"): + ReasonerFeatureRuntime(reasoner, identity=_identity()) + + +def test_runtime_rejects_generator_specific_k_normalization_at_startup() -> None: + reasoner = _FakeReasoner() + reasoner.model.layers[1].self_attn.k_norm_und_for_gen = nn.Identity() + + with pytest.raises(NotImplementedError, match=r"affected layers=\[1\]"): + ReasonerFeatureRuntime(reasoner, identity=_identity()) + + +def test_load_validates_execution_compatibility_before_checkpoint_restore( + monkeypatch: pytest.MonkeyPatch, +) -> None: + reasoner = _FakeReasoner() + reasoner.model.reasoner_forward = None # type: ignore[method-assign] + checkpoint_loader_called = False + + monkeypatch.setattr(reasoner_runtime_module, "build_reasoner", lambda *_args, **_kwargs: reasoner) + + def checkpoint_loader(*_args: object, **_kwargs: object) -> None: + nonlocal checkpoint_loader_called + checkpoint_loader_called = True + + monkeypatch.setattr(reasoner_runtime_module, "load_reasoner_only_dcp", checkpoint_loader) + spec = ReasonerRuntimeSpec( + checkpoint="unused", + checkpoint_source="regular", + device="cpu", + dtype=torch.float32, + identity=_identity(), + ) + + with pytest.raises(TypeError, match="model.reasoner_forward to be callable"): + ReasonerFeatureRuntime.load(object(), spec) + assert not checkpoint_loader_called + + +def test_fingerprint_mismatch_fails_before_model_execution(monkeypatch: pytest.MonkeyPatch) -> None: + runtime = _runtime() + request = _request("sample", [1, 2, 3], identity=runtime.identity, fingerprint="tampered") + extractor_called = False + + def fail_if_called(*_args: object, **_kwargs: object) -> ReasonerFeatureBatch: + nonlocal extractor_called + extractor_called = True + raise AssertionError("extractor must not run for an invalid fingerprint") + + monkeypatch.setattr(reasoner_runtime_module, "extract_reasoner_feature_batch", fail_if_called) + + with pytest.raises(ValueError, match="fingerprint mismatch"): + runtime.execute([request]) + assert not extractor_called + + +def test_token_budget_fails_before_model_execution(monkeypatch: pytest.MonkeyPatch) -> None: + runtime = _runtime(max_total_tokens=2) + request = _request("too-long", [1, 2, 3], identity=runtime.identity) + extractor_called = False + + def fail_if_called(*_args: object, **_kwargs: object) -> ReasonerFeatureBatch: + nonlocal extractor_called + extractor_called = True + raise AssertionError("extractor must not run above the token budget") + + monkeypatch.setattr(reasoner_runtime_module, "extract_reasoner_feature_batch", fail_if_called) + + with pytest.raises(ValueError, match=r"token count 3 exceeds limit 2"): + runtime.execute([request]) + assert not extractor_called + + +def test_estimated_response_bytes_uses_signature_and_total_tokens() -> None: + runtime = _runtime() + requests = ( + _request("first", [1, 2, 3], identity=runtime.identity), + _request("second", [4, 5], identity=runtime.identity), + ) + + expected = 5 * 2 * 2 * 2 * 2 * torch.tensor([], dtype=torch.float32).element_size() + assert runtime.estimated_response_bytes(requests) == expected + + +def test_execute_preserves_request_order(monkeypatch: pytest.MonkeyPatch) -> None: + runtime = _runtime() + requests = ( + _request("second", [4, 5], identity=runtime.identity), + _request("first", [1, 2, 3], identity=runtime.identity), + ) + seen_keys: list[str] = [] + + def extract( + reasoner: nn.Module, + received: Sequence[ReasonerFeatureRequest], + ) -> ReasonerFeatureBatch: + assert reasoner is runtime.reasoner + seen_keys.extend(request.sample_key for request in received) + return _features_for(received) + + monkeypatch.setattr(reasoner_runtime_module, "extract_reasoner_feature_batch", extract) + + features = runtime.execute(requests) + + assert seen_keys == ["second", "first"] + assert features.fingerprints == tuple(request.fingerprint for request in requests) + assert features.causal_offsets.tolist() == [0, 2, 5] + + +def test_execute_rejects_extractor_output_in_a_different_request_order( + monkeypatch: pytest.MonkeyPatch, +) -> None: + runtime = _runtime() + requests = ( + _request("first", [1, 2], identity=runtime.identity), + _request("second", [3, 4], identity=runtime.identity), + ) + + def extract( + _reasoner: nn.Module, + received: Sequence[ReasonerFeatureRequest], + ) -> ReasonerFeatureBatch: + features = _features_for(received) + return ReasonerFeatureBatch( + cross_k=features.cross_k, + cross_v=features.cross_v, + causal_offsets=features.causal_offsets, + fingerprints=tuple(reversed(features.fingerprints)), + ) + + monkeypatch.setattr(reasoner_runtime_module, "extract_reasoner_feature_batch", extract) + + with pytest.raises(ReasonerRuntimeInvariantError, match="do not preserve request order"): + runtime.execute(requests) + + +def test_execute_raises_fatal_invariant_error_for_output_signature_change( + monkeypatch: pytest.MonkeyPatch, +) -> None: + runtime = _runtime() + requests = (_request("sample", [1, 2], identity=runtime.identity),) + + def extract( + _reasoner: nn.Module, + received: Sequence[ReasonerFeatureRequest], + ) -> ReasonerFeatureBatch: + features = _features_for(received) + return ReasonerFeatureBatch( + cross_k=tuple(tensor.to(torch.bfloat16) for tensor in features.cross_k), + cross_v=tuple(tensor.to(torch.bfloat16) for tensor in features.cross_v), + causal_offsets=features.causal_offsets, + fingerprints=features.fingerprints, + ) + + monkeypatch.setattr(reasoner_runtime_module, "extract_reasoner_feature_batch", extract) + + with pytest.raises(ReasonerRuntimeInvariantError, match="output signature changed"): + runtime.execute(requests) diff --git a/cosmos_framework/model/model_lifecycle_test.py b/cosmos_framework/model/model_lifecycle_test.py new file mode 100644 index 000000000..969f52526 --- /dev/null +++ b/cosmos_framework/model/model_lifecycle_test.py @@ -0,0 +1,123 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +from __future__ import annotations + +from types import SimpleNamespace +from typing import Any + +import pytest + +from cosmos_framework.model._base import ImaginaireModel, close_model +from cosmos_framework.model.generator.omni_mot_model import OmniMoTModel +from cosmos_framework.trainer import ImaginaireTrainer + +pytestmark = [pytest.mark.level(0), pytest.mark.gpus(0)] + + +class _RecordingModel(ImaginaireModel): + def __init__(self, *, close_error: BaseException | None = None) -> None: + super().__init__() + self.close_error = close_error + self.close_calls = 0 + + def close(self) -> None: + self.close_calls += 1 + if self.close_error is not None: + raise self.close_error + + +class _RecordingTrainer(ImaginaireTrainer): + def __init__(self, *, train_error: BaseException | None = None) -> None: + self.train_error = train_error + self.train_calls = 0 + + def _train( + self, + model: ImaginaireModel, + dataloader_train: Any, + dataloader_val: Any, + ) -> None: + del model, dataloader_train, dataloader_val + self.train_calls += 1 + if self.train_error is not None: + raise self.train_error + + +class _RecordingProvider: + def __init__(self) -> None: + self.close_calls = 0 + + def close(self) -> None: + self.close_calls += 1 + + +def test_close_model_preserves_primary_error_when_cleanup_also_fails() -> None: + primary_error = RuntimeError("primary failure") + close_error = ValueError("cleanup failure") + model = _RecordingModel(close_error=close_error) + + close_model(model, primary_error=primary_error) + + assert model.close_calls == 1 + if hasattr(primary_error, "add_note"): + assert any("ValueError: cleanup failure" in note for note in getattr(primary_error, "__notes__", ())) + + with pytest.raises(ValueError, match="cleanup failure") as raised: + close_model(_RecordingModel(close_error=close_error)) + assert raised.value is close_error + + +def test_trainer_closes_model_after_successful_training() -> None: + trainer = _RecordingTrainer() + model = _RecordingModel() + dataloader: Any = object() + + trainer.train(model, dataloader, dataloader) + + assert trainer.train_calls == 1 + assert model.close_calls == 1 + + +def test_trainer_closes_model_and_preserves_training_failure() -> None: + train_error = RuntimeError("training failure") + trainer = _RecordingTrainer(train_error=train_error) + model = _RecordingModel() + dataloader: Any = object() + + with pytest.raises(RuntimeError, match="training failure") as raised: + trainer.train(model, dataloader, dataloader) + + assert raised.value is train_error + assert trainer.train_calls == 1 + assert model.close_calls == 1 + + +def test_omni_setup_failure_closes_provider_once_and_close_is_idempotent( + monkeypatch: pytest.MonkeyPatch, +) -> None: + provider = _RecordingProvider() + created_models: list[OmniMoTModel] = [] + + def create_provider(model: OmniMoTModel) -> _RecordingProvider: + created_models.append(model) + return provider + + def fail_setup(_model: OmniMoTModel) -> None: + raise RuntimeError("setup failure") + + monkeypatch.setattr(OmniMoTModel, "_create_reasoner_feature_provider", create_provider) + monkeypatch.setattr(OmniMoTModel, "set_precision", fail_setup) + config = SimpleNamespace(reasoner_conditioning={"backend": "joint"}) + + with pytest.raises(RuntimeError, match="setup failure"): + OmniMoTModel(config) # type: ignore[arg-type] + + assert len(created_models) == 1 + model = created_models[0] + assert provider.close_calls == 1 + assert model.reasoner_feature_provider is None + + model.close() + model.close() + assert provider.close_calls == 1 diff --git a/cosmos_framework/protos/__init__.py b/cosmos_framework/protos/__init__.py new file mode 100644 index 000000000..7429a21ff --- /dev/null +++ b/cosmos_framework/protos/__init__.py @@ -0,0 +1,4 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +"""Versioned wire protocols used by Cosmos Framework services.""" diff --git a/cosmos_framework/protos/reasoner_features/__init__.py b/cosmos_framework/protos/reasoner_features/__init__.py new file mode 100644 index 000000000..0d87f02e7 --- /dev/null +++ b/cosmos_framework/protos/reasoner_features/__init__.py @@ -0,0 +1,4 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +"""Reasoner feature-service protocols.""" diff --git a/cosmos_framework/protos/reasoner_features/v1/__init__.py b/cosmos_framework/protos/reasoner_features/v1/__init__.py new file mode 100644 index 000000000..6f922dcdf --- /dev/null +++ b/cosmos_framework/protos/reasoner_features/v1/__init__.py @@ -0,0 +1,4 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +"""Version 1 of the Reasoner feature-service wire protocol.""" diff --git a/cosmos_framework/protos/reasoner_features/v1/reasoner_features.proto b/cosmos_framework/protos/reasoner_features/v1/reasoner_features.proto new file mode 100644 index 000000000..c850b0fab --- /dev/null +++ b/cosmos_framework/protos/reasoner_features/v1/reasoner_features.proto @@ -0,0 +1,149 @@ +// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: OpenMDW-1.1 + +syntax = "proto3"; + +package cosmos.reasoner_features.v1; + +// DType describes the exact little-endian representation stored in a tensor's +// bytes field. Implementations must reject an unspecified or unsupported dtype. +enum DType { + DTYPE_UNSPECIFIED = 0; + DTYPE_INT32 = 1; + DTYPE_INT64 = 2; + DTYPE_FLOAT16 = 3; + DTYPE_BFLOAT16 = 4; + DTYPE_FLOAT32 = 5; + DTYPE_FLOAT64 = 6; +} + +enum TensorKind { + TENSOR_KIND_UNSPECIFIED = 0; + TENSOR_KIND_CROSS_K = 1; + TENSOR_KIND_CROSS_V = 2; +} + +// TensorPayload is used for the comparatively small request tensors and +// response offsets. Large Reasoner K/V tensors use TensorChunk below. +message TensorPayload { + DType dtype = 1; + repeated uint64 shape = 2; + bytes data = 3; +} + +message CacheIdentity { + string reasoner = 1; + string tokenizer = 2; + string framing = 3; +} + +message FeatureSignature { + DType dtype = 1; + uint32 num_layers = 2; + uint32 num_kv_heads = 3; + uint32 head_dim = 4; +} + +// Capability names are stable, lower-case identifiers. Version lets a service +// advertise later extensions without changing the GetInfo envelope. +message Capability { + string name = 1; + uint32 version = 2; +} + +message GetInfoRequest { + uint32 protocol_version = 1; +} + +message GetInfoResponse { + uint32 protocol_version = 1; + string service_instance_id = 2; + CacheIdentity identity = 3; + FeatureSignature signature = 4; + repeated Capability capabilities = 5; + uint64 max_batch_tokens = 6; + uint64 max_queued_tokens = 7; + uint32 max_chunk_bytes = 8; + bool ready = 9; + uint32 max_requests = 10; +} + +// One logical Reasoner input. causal_offsets partitions this request's token +// stream into independent causal documents. It is [0, S] for an ordinary +// sample, and can contain more boundaries for per-view captions. The service +// must preserve document isolation during Reasoner execution. +message ReasonerFeatureRequest { + string sample_key = 1; + TensorPayload token_ids = 2; + TensorPayload position_ids = 3; + TensorPayload causal_offsets = 4; + string fingerprint = 5; +} + +message GenerateRequest { + uint32 protocol_version = 1; + // Stable across an idempotent retry of exactly the same ordered payload. + string request_id = 2; + CacheIdentity identity = 3; + repeated ReasonerFeatureRequest requests = 4; +} + +message GenerateHeader { + uint32 protocol_version = 1; + string request_id = 2; + CacheIdentity identity = 3; + FeatureSignature signature = 4; + // Ordered exactly like GenerateRequest.requests. + repeated string fingerprints = 5; + // Partitions the returned unpadded K/V stream by request, not by the causal + // documents inside each request. + TensorPayload causal_offsets = 6; + uint64 total_tensor_bytes = 7; + string service_instance_id = 8; +} + +// Chunks are streamed in canonical order: for every layer from zero upward, +// all CROSS_K chunks followed by all CROSS_V chunks. chunk_index is global to +// the response and byte_offset is local to the described tensor. sha256 is the +// raw 32-byte digest of data. +message TensorChunk { + uint64 chunk_index = 1; + TensorKind kind = 2; + uint32 layer_index = 3; + DType dtype = 4; + repeated uint64 shape = 5; + uint64 byte_offset = 6; + uint64 total_bytes = 7; + bytes data = 8; + bytes sha256 = 9; +} + +message ServerTiming { + uint64 queue_ns = 1; + uint64 compute_ns = 2; + uint64 device_to_host_ns = 3; + uint64 serialization_ns = 4; +} + +message GenerateTrailer { + uint64 total_chunks = 1; + uint64 total_tensor_bytes = 2; + // SHA-256 over raw tensors in canonical K0 || V0 || K1 || V1 ... order. + bytes response_sha256 = 3; + ServerTiming timing = 4; + uint64 cache_hits = 5; + uint64 cache_misses = 6; +} + +message GenerateResponse { + oneof body { + GenerateHeader header = 1; + TensorChunk tensor_chunk = 2; + GenerateTrailer trailer = 3; + } +} + +service ReasonerFeatureService { + rpc GetInfo(GetInfoRequest) returns (GetInfoResponse); + rpc Generate(GenerateRequest) returns (stream GenerateResponse); +} diff --git a/cosmos_framework/protos/reasoner_features/v1/reasoner_features_pb2.py b/cosmos_framework/protos/reasoner_features/v1/reasoner_features_pb2.py new file mode 100644 index 000000000..8739d1214 --- /dev/null +++ b/cosmos_framework/protos/reasoner_features/v1/reasoner_features_pb2.py @@ -0,0 +1,67 @@ +# -*- coding: utf-8 -*- +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 +# Generated by the protocol buffer compiler. DO NOT EDIT! +# NO CHECKED-IN PROTOBUF GENCODE +# source: cosmos_framework/protos/reasoner_features/v1/reasoner_features.proto +# Protobuf Python Version: 6.31.1 +"""Generated protocol buffer code.""" + +from google.protobuf import descriptor as _descriptor +from google.protobuf import descriptor_pool as _descriptor_pool +from google.protobuf import runtime_version as _runtime_version +from google.protobuf import symbol_database as _symbol_database +from google.protobuf.internal import builder as _builder + +_runtime_version.ValidateProtobufRuntimeVersion( + _runtime_version.Domain.PUBLIC, 6, 31, 1, "", "cosmos_framework/protos/reasoner_features/v1/reasoner_features.proto" +) +# @@protoc_insertion_point(imports) + +_sym_db = _symbol_database.Default() + + +DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile( + b'\nDcosmos_framework/protos/reasoner_features/v1/reasoner_features.proto\x12\x1b\x63osmos.reasoner_features.v1"_\n\rTensorPayload\x12\x31\n\x05\x64type\x18\x01 \x01(\x0e\x32".cosmos.reasoner_features.v1.DType\x12\r\n\x05shape\x18\x02 \x03(\x04\x12\x0c\n\x04\x64\x61ta\x18\x03 \x01(\x0c"E\n\rCacheIdentity\x12\x10\n\x08reasoner\x18\x01 \x01(\t\x12\x11\n\ttokenizer\x18\x02 \x01(\t\x12\x0f\n\x07\x66raming\x18\x03 \x01(\t"\x81\x01\n\x10\x46\x65\x61tureSignature\x12\x31\n\x05\x64type\x18\x01 \x01(\x0e\x32".cosmos.reasoner_features.v1.DType\x12\x12\n\nnum_layers\x18\x02 \x01(\r\x12\x14\n\x0cnum_kv_heads\x18\x03 \x01(\r\x12\x10\n\x08head_dim\x18\x04 \x01(\r"+\n\nCapability\x12\x0c\n\x04name\x18\x01 \x01(\t\x12\x0f\n\x07version\x18\x02 \x01(\r"*\n\x0eGetInfoRequest\x12\x18\n\x10protocol_version\x18\x01 \x01(\r"\xfa\x02\n\x0fGetInfoResponse\x12\x18\n\x10protocol_version\x18\x01 \x01(\r\x12\x1b\n\x13service_instance_id\x18\x02 \x01(\t\x12<\n\x08identity\x18\x03 \x01(\x0b\x32*.cosmos.reasoner_features.v1.CacheIdentity\x12@\n\tsignature\x18\x04 \x01(\x0b\x32-.cosmos.reasoner_features.v1.FeatureSignature\x12=\n\x0c\x63\x61pabilities\x18\x05 \x03(\x0b\x32\'.cosmos.reasoner_features.v1.Capability\x12\x18\n\x10max_batch_tokens\x18\x06 \x01(\x04\x12\x19\n\x11max_queued_tokens\x18\x07 \x01(\x04\x12\x17\n\x0fmax_chunk_bytes\x18\x08 \x01(\r\x12\r\n\x05ready\x18\t \x01(\x08\x12\x14\n\x0cmax_requests\x18\n \x01(\r"\x86\x02\n\x16ReasonerFeatureRequest\x12\x12\n\nsample_key\x18\x01 \x01(\t\x12=\n\ttoken_ids\x18\x02 \x01(\x0b\x32*.cosmos.reasoner_features.v1.TensorPayload\x12@\n\x0cposition_ids\x18\x03 \x01(\x0b\x32*.cosmos.reasoner_features.v1.TensorPayload\x12\x42\n\x0e\x63\x61usal_offsets\x18\x04 \x01(\x0b\x32*.cosmos.reasoner_features.v1.TensorPayload\x12\x13\n\x0b\x66ingerprint\x18\x05 \x01(\t"\xc4\x01\n\x0fGenerateRequest\x12\x18\n\x10protocol_version\x18\x01 \x01(\r\x12\x12\n\nrequest_id\x18\x02 \x01(\t\x12<\n\x08identity\x18\x03 \x01(\x0b\x32*.cosmos.reasoner_features.v1.CacheIdentity\x12\x45\n\x08requests\x18\x04 \x03(\x0b\x32\x33.cosmos.reasoner_features.v1.ReasonerFeatureRequest"\xd1\x02\n\x0eGenerateHeader\x12\x18\n\x10protocol_version\x18\x01 \x01(\r\x12\x12\n\nrequest_id\x18\x02 \x01(\t\x12<\n\x08identity\x18\x03 \x01(\x0b\x32*.cosmos.reasoner_features.v1.CacheIdentity\x12@\n\tsignature\x18\x04 \x01(\x0b\x32-.cosmos.reasoner_features.v1.FeatureSignature\x12\x14\n\x0c\x66ingerprints\x18\x05 \x03(\t\x12\x42\n\x0e\x63\x61usal_offsets\x18\x06 \x01(\x0b\x32*.cosmos.reasoner_features.v1.TensorPayload\x12\x1a\n\x12total_tensor_bytes\x18\x07 \x01(\x04\x12\x1b\n\x13service_instance_id\x18\x08 \x01(\t"\xf8\x01\n\x0bTensorChunk\x12\x13\n\x0b\x63hunk_index\x18\x01 \x01(\x04\x12\x35\n\x04kind\x18\x02 \x01(\x0e\x32\'.cosmos.reasoner_features.v1.TensorKind\x12\x13\n\x0blayer_index\x18\x03 \x01(\r\x12\x31\n\x05\x64type\x18\x04 \x01(\x0e\x32".cosmos.reasoner_features.v1.DType\x12\r\n\x05shape\x18\x05 \x03(\x04\x12\x13\n\x0b\x62yte_offset\x18\x06 \x01(\x04\x12\x13\n\x0btotal_bytes\x18\x07 \x01(\x04\x12\x0c\n\x04\x64\x61ta\x18\x08 \x01(\x0c\x12\x0e\n\x06sha256\x18\t \x01(\x0c"i\n\x0cServerTiming\x12\x10\n\x08queue_ns\x18\x01 \x01(\x04\x12\x12\n\ncompute_ns\x18\x02 \x01(\x04\x12\x19\n\x11\x64\x65vice_to_host_ns\x18\x03 \x01(\x04\x12\x18\n\x10serialization_ns\x18\x04 \x01(\x04"\xc1\x01\n\x0fGenerateTrailer\x12\x14\n\x0ctotal_chunks\x18\x01 \x01(\x04\x12\x1a\n\x12total_tensor_bytes\x18\x02 \x01(\x04\x12\x17\n\x0fresponse_sha256\x18\x03 \x01(\x0c\x12\x39\n\x06timing\x18\x04 \x01(\x0b\x32).cosmos.reasoner_features.v1.ServerTiming\x12\x12\n\ncache_hits\x18\x05 \x01(\x04\x12\x14\n\x0c\x63\x61\x63he_misses\x18\x06 \x01(\x04"\xdc\x01\n\x10GenerateResponse\x12=\n\x06header\x18\x01 \x01(\x0b\x32+.cosmos.reasoner_features.v1.GenerateHeaderH\x00\x12@\n\x0ctensor_chunk\x18\x02 \x01(\x0b\x32(.cosmos.reasoner_features.v1.TensorChunkH\x00\x12?\n\x07trailer\x18\x03 \x01(\x0b\x32,.cosmos.reasoner_features.v1.GenerateTrailerH\x00\x42\x06\n\x04\x62ody*\x8d\x01\n\x05\x44Type\x12\x15\n\x11\x44TYPE_UNSPECIFIED\x10\x00\x12\x0f\n\x0b\x44TYPE_INT32\x10\x01\x12\x0f\n\x0b\x44TYPE_INT64\x10\x02\x12\x11\n\rDTYPE_FLOAT16\x10\x03\x12\x12\n\x0e\x44TYPE_BFLOAT16\x10\x04\x12\x11\n\rDTYPE_FLOAT32\x10\x05\x12\x11\n\rDTYPE_FLOAT64\x10\x06*[\n\nTensorKind\x12\x1b\n\x17TENSOR_KIND_UNSPECIFIED\x10\x00\x12\x17\n\x13TENSOR_KIND_CROSS_K\x10\x01\x12\x17\n\x13TENSOR_KIND_CROSS_V\x10\x02\x32\xe9\x01\n\x16ReasonerFeatureService\x12\x64\n\x07GetInfo\x12+.cosmos.reasoner_features.v1.GetInfoRequest\x1a,.cosmos.reasoner_features.v1.GetInfoResponse\x12i\n\x08Generate\x12,.cosmos.reasoner_features.v1.GenerateRequest\x1a-.cosmos.reasoner_features.v1.GenerateResponse0\x01\x62\x06proto3' +) + +_globals = globals() +_builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals) +_builder.BuildTopDescriptorsAndMessages( + DESCRIPTOR, "cosmos_framework.protos.reasoner_features.v1.reasoner_features_pb2", _globals +) +if not _descriptor._USE_C_DESCRIPTORS: + DESCRIPTOR._loaded_options = None + _globals["_DTYPE"]._serialized_start = 2453 + _globals["_DTYPE"]._serialized_end = 2594 + _globals["_TENSORKIND"]._serialized_start = 2596 + _globals["_TENSORKIND"]._serialized_end = 2687 + _globals["_TENSORPAYLOAD"]._serialized_start = 101 + _globals["_TENSORPAYLOAD"]._serialized_end = 196 + _globals["_CACHEIDENTITY"]._serialized_start = 198 + _globals["_CACHEIDENTITY"]._serialized_end = 267 + _globals["_FEATURESIGNATURE"]._serialized_start = 270 + _globals["_FEATURESIGNATURE"]._serialized_end = 399 + _globals["_CAPABILITY"]._serialized_start = 401 + _globals["_CAPABILITY"]._serialized_end = 444 + _globals["_GETINFOREQUEST"]._serialized_start = 446 + _globals["_GETINFOREQUEST"]._serialized_end = 488 + _globals["_GETINFORESPONSE"]._serialized_start = 491 + _globals["_GETINFORESPONSE"]._serialized_end = 869 + _globals["_REASONERFEATUREREQUEST"]._serialized_start = 872 + _globals["_REASONERFEATUREREQUEST"]._serialized_end = 1134 + _globals["_GENERATEREQUEST"]._serialized_start = 1137 + _globals["_GENERATEREQUEST"]._serialized_end = 1333 + _globals["_GENERATEHEADER"]._serialized_start = 1336 + _globals["_GENERATEHEADER"]._serialized_end = 1673 + _globals["_TENSORCHUNK"]._serialized_start = 1676 + _globals["_TENSORCHUNK"]._serialized_end = 1924 + _globals["_SERVERTIMING"]._serialized_start = 1926 + _globals["_SERVERTIMING"]._serialized_end = 2031 + _globals["_GENERATETRAILER"]._serialized_start = 2034 + _globals["_GENERATETRAILER"]._serialized_end = 2227 + _globals["_GENERATERESPONSE"]._serialized_start = 2230 + _globals["_GENERATERESPONSE"]._serialized_end = 2450 + _globals["_REASONERFEATURESERVICE"]._serialized_start = 2690 + _globals["_REASONERFEATURESERVICE"]._serialized_end = 2923 +# @@protoc_insertion_point(module_scope) diff --git a/cosmos_framework/protos/reasoner_features/v1/reasoner_features_pb2.pyi b/cosmos_framework/protos/reasoner_features/v1/reasoner_features_pb2.pyi new file mode 100644 index 000000000..4d09b2d4c --- /dev/null +++ b/cosmos_framework/protos/reasoner_features/v1/reasoner_features_pb2.pyi @@ -0,0 +1,316 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +from collections.abc import Iterable as _Iterable +from collections.abc import Mapping as _Mapping +from typing import ClassVar as _ClassVar +from typing import Optional as _Optional +from typing import Union as _Union + +from google.protobuf import descriptor as _descriptor +from google.protobuf import message as _message +from google.protobuf.internal import containers as _containers +from google.protobuf.internal import enum_type_wrapper as _enum_type_wrapper + +DESCRIPTOR: _descriptor.FileDescriptor + +class DType(int, metaclass=_enum_type_wrapper.EnumTypeWrapper): + __slots__ = () + DTYPE_UNSPECIFIED: _ClassVar[DType] + DTYPE_INT32: _ClassVar[DType] + DTYPE_INT64: _ClassVar[DType] + DTYPE_FLOAT16: _ClassVar[DType] + DTYPE_BFLOAT16: _ClassVar[DType] + DTYPE_FLOAT32: _ClassVar[DType] + DTYPE_FLOAT64: _ClassVar[DType] + +class TensorKind(int, metaclass=_enum_type_wrapper.EnumTypeWrapper): + __slots__ = () + TENSOR_KIND_UNSPECIFIED: _ClassVar[TensorKind] + TENSOR_KIND_CROSS_K: _ClassVar[TensorKind] + TENSOR_KIND_CROSS_V: _ClassVar[TensorKind] + +DTYPE_UNSPECIFIED: DType +DTYPE_INT32: DType +DTYPE_INT64: DType +DTYPE_FLOAT16: DType +DTYPE_BFLOAT16: DType +DTYPE_FLOAT32: DType +DTYPE_FLOAT64: DType +TENSOR_KIND_UNSPECIFIED: TensorKind +TENSOR_KIND_CROSS_K: TensorKind +TENSOR_KIND_CROSS_V: TensorKind + +class TensorPayload(_message.Message): + __slots__ = ("dtype", "shape", "data") + DTYPE_FIELD_NUMBER: _ClassVar[int] + SHAPE_FIELD_NUMBER: _ClassVar[int] + DATA_FIELD_NUMBER: _ClassVar[int] + dtype: DType + shape: _containers.RepeatedScalarFieldContainer[int] + data: bytes + def __init__( + self, + dtype: _Optional[_Union[DType, str]] = ..., + shape: _Optional[_Iterable[int]] = ..., + data: _Optional[bytes] = ..., + ) -> None: ... + +class CacheIdentity(_message.Message): + __slots__ = ("reasoner", "tokenizer", "framing") + REASONER_FIELD_NUMBER: _ClassVar[int] + TOKENIZER_FIELD_NUMBER: _ClassVar[int] + FRAMING_FIELD_NUMBER: _ClassVar[int] + reasoner: str + tokenizer: str + framing: str + def __init__( + self, reasoner: _Optional[str] = ..., tokenizer: _Optional[str] = ..., framing: _Optional[str] = ... + ) -> None: ... + +class FeatureSignature(_message.Message): + __slots__ = ("dtype", "num_layers", "num_kv_heads", "head_dim") + DTYPE_FIELD_NUMBER: _ClassVar[int] + NUM_LAYERS_FIELD_NUMBER: _ClassVar[int] + NUM_KV_HEADS_FIELD_NUMBER: _ClassVar[int] + HEAD_DIM_FIELD_NUMBER: _ClassVar[int] + dtype: DType + num_layers: int + num_kv_heads: int + head_dim: int + def __init__( + self, + dtype: _Optional[_Union[DType, str]] = ..., + num_layers: _Optional[int] = ..., + num_kv_heads: _Optional[int] = ..., + head_dim: _Optional[int] = ..., + ) -> None: ... + +class Capability(_message.Message): + __slots__ = ("name", "version") + NAME_FIELD_NUMBER: _ClassVar[int] + VERSION_FIELD_NUMBER: _ClassVar[int] + name: str + version: int + def __init__(self, name: _Optional[str] = ..., version: _Optional[int] = ...) -> None: ... + +class GetInfoRequest(_message.Message): + __slots__ = ("protocol_version",) + PROTOCOL_VERSION_FIELD_NUMBER: _ClassVar[int] + protocol_version: int + def __init__(self, protocol_version: _Optional[int] = ...) -> None: ... + +class GetInfoResponse(_message.Message): + __slots__ = ( + "protocol_version", + "service_instance_id", + "identity", + "signature", + "capabilities", + "max_batch_tokens", + "max_queued_tokens", + "max_chunk_bytes", + "ready", + "max_requests", + ) + PROTOCOL_VERSION_FIELD_NUMBER: _ClassVar[int] + SERVICE_INSTANCE_ID_FIELD_NUMBER: _ClassVar[int] + IDENTITY_FIELD_NUMBER: _ClassVar[int] + SIGNATURE_FIELD_NUMBER: _ClassVar[int] + CAPABILITIES_FIELD_NUMBER: _ClassVar[int] + MAX_BATCH_TOKENS_FIELD_NUMBER: _ClassVar[int] + MAX_QUEUED_TOKENS_FIELD_NUMBER: _ClassVar[int] + MAX_CHUNK_BYTES_FIELD_NUMBER: _ClassVar[int] + READY_FIELD_NUMBER: _ClassVar[int] + MAX_REQUESTS_FIELD_NUMBER: _ClassVar[int] + protocol_version: int + service_instance_id: str + identity: CacheIdentity + signature: FeatureSignature + capabilities: _containers.RepeatedCompositeFieldContainer[Capability] + max_batch_tokens: int + max_queued_tokens: int + max_chunk_bytes: int + ready: bool + max_requests: int + def __init__( + self, + protocol_version: _Optional[int] = ..., + service_instance_id: _Optional[str] = ..., + identity: _Optional[_Union[CacheIdentity, _Mapping]] = ..., + signature: _Optional[_Union[FeatureSignature, _Mapping]] = ..., + capabilities: _Optional[_Iterable[_Union[Capability, _Mapping]]] = ..., + max_batch_tokens: _Optional[int] = ..., + max_queued_tokens: _Optional[int] = ..., + max_chunk_bytes: _Optional[int] = ..., + ready: bool = ..., + max_requests: _Optional[int] = ..., + ) -> None: ... + +class ReasonerFeatureRequest(_message.Message): + __slots__ = ("sample_key", "token_ids", "position_ids", "causal_offsets", "fingerprint") + SAMPLE_KEY_FIELD_NUMBER: _ClassVar[int] + TOKEN_IDS_FIELD_NUMBER: _ClassVar[int] + POSITION_IDS_FIELD_NUMBER: _ClassVar[int] + CAUSAL_OFFSETS_FIELD_NUMBER: _ClassVar[int] + FINGERPRINT_FIELD_NUMBER: _ClassVar[int] + sample_key: str + token_ids: TensorPayload + position_ids: TensorPayload + causal_offsets: TensorPayload + fingerprint: str + def __init__( + self, + sample_key: _Optional[str] = ..., + token_ids: _Optional[_Union[TensorPayload, _Mapping]] = ..., + position_ids: _Optional[_Union[TensorPayload, _Mapping]] = ..., + causal_offsets: _Optional[_Union[TensorPayload, _Mapping]] = ..., + fingerprint: _Optional[str] = ..., + ) -> None: ... + +class GenerateRequest(_message.Message): + __slots__ = ("protocol_version", "request_id", "identity", "requests") + PROTOCOL_VERSION_FIELD_NUMBER: _ClassVar[int] + REQUEST_ID_FIELD_NUMBER: _ClassVar[int] + IDENTITY_FIELD_NUMBER: _ClassVar[int] + REQUESTS_FIELD_NUMBER: _ClassVar[int] + protocol_version: int + request_id: str + identity: CacheIdentity + requests: _containers.RepeatedCompositeFieldContainer[ReasonerFeatureRequest] + def __init__( + self, + protocol_version: _Optional[int] = ..., + request_id: _Optional[str] = ..., + identity: _Optional[_Union[CacheIdentity, _Mapping]] = ..., + requests: _Optional[_Iterable[_Union[ReasonerFeatureRequest, _Mapping]]] = ..., + ) -> None: ... + +class GenerateHeader(_message.Message): + __slots__ = ( + "protocol_version", + "request_id", + "identity", + "signature", + "fingerprints", + "causal_offsets", + "total_tensor_bytes", + "service_instance_id", + ) + PROTOCOL_VERSION_FIELD_NUMBER: _ClassVar[int] + REQUEST_ID_FIELD_NUMBER: _ClassVar[int] + IDENTITY_FIELD_NUMBER: _ClassVar[int] + SIGNATURE_FIELD_NUMBER: _ClassVar[int] + FINGERPRINTS_FIELD_NUMBER: _ClassVar[int] + CAUSAL_OFFSETS_FIELD_NUMBER: _ClassVar[int] + TOTAL_TENSOR_BYTES_FIELD_NUMBER: _ClassVar[int] + SERVICE_INSTANCE_ID_FIELD_NUMBER: _ClassVar[int] + protocol_version: int + request_id: str + identity: CacheIdentity + signature: FeatureSignature + fingerprints: _containers.RepeatedScalarFieldContainer[str] + causal_offsets: TensorPayload + total_tensor_bytes: int + service_instance_id: str + def __init__( + self, + protocol_version: _Optional[int] = ..., + request_id: _Optional[str] = ..., + identity: _Optional[_Union[CacheIdentity, _Mapping]] = ..., + signature: _Optional[_Union[FeatureSignature, _Mapping]] = ..., + fingerprints: _Optional[_Iterable[str]] = ..., + causal_offsets: _Optional[_Union[TensorPayload, _Mapping]] = ..., + total_tensor_bytes: _Optional[int] = ..., + service_instance_id: _Optional[str] = ..., + ) -> None: ... + +class TensorChunk(_message.Message): + __slots__ = ("chunk_index", "kind", "layer_index", "dtype", "shape", "byte_offset", "total_bytes", "data", "sha256") + CHUNK_INDEX_FIELD_NUMBER: _ClassVar[int] + KIND_FIELD_NUMBER: _ClassVar[int] + LAYER_INDEX_FIELD_NUMBER: _ClassVar[int] + DTYPE_FIELD_NUMBER: _ClassVar[int] + SHAPE_FIELD_NUMBER: _ClassVar[int] + BYTE_OFFSET_FIELD_NUMBER: _ClassVar[int] + TOTAL_BYTES_FIELD_NUMBER: _ClassVar[int] + DATA_FIELD_NUMBER: _ClassVar[int] + SHA256_FIELD_NUMBER: _ClassVar[int] + chunk_index: int + kind: TensorKind + layer_index: int + dtype: DType + shape: _containers.RepeatedScalarFieldContainer[int] + byte_offset: int + total_bytes: int + data: bytes + sha256: bytes + def __init__( + self, + chunk_index: _Optional[int] = ..., + kind: _Optional[_Union[TensorKind, str]] = ..., + layer_index: _Optional[int] = ..., + dtype: _Optional[_Union[DType, str]] = ..., + shape: _Optional[_Iterable[int]] = ..., + byte_offset: _Optional[int] = ..., + total_bytes: _Optional[int] = ..., + data: _Optional[bytes] = ..., + sha256: _Optional[bytes] = ..., + ) -> None: ... + +class ServerTiming(_message.Message): + __slots__ = ("queue_ns", "compute_ns", "device_to_host_ns", "serialization_ns") + QUEUE_NS_FIELD_NUMBER: _ClassVar[int] + COMPUTE_NS_FIELD_NUMBER: _ClassVar[int] + DEVICE_TO_HOST_NS_FIELD_NUMBER: _ClassVar[int] + SERIALIZATION_NS_FIELD_NUMBER: _ClassVar[int] + queue_ns: int + compute_ns: int + device_to_host_ns: int + serialization_ns: int + def __init__( + self, + queue_ns: _Optional[int] = ..., + compute_ns: _Optional[int] = ..., + device_to_host_ns: _Optional[int] = ..., + serialization_ns: _Optional[int] = ..., + ) -> None: ... + +class GenerateTrailer(_message.Message): + __slots__ = ("total_chunks", "total_tensor_bytes", "response_sha256", "timing", "cache_hits", "cache_misses") + TOTAL_CHUNKS_FIELD_NUMBER: _ClassVar[int] + TOTAL_TENSOR_BYTES_FIELD_NUMBER: _ClassVar[int] + RESPONSE_SHA256_FIELD_NUMBER: _ClassVar[int] + TIMING_FIELD_NUMBER: _ClassVar[int] + CACHE_HITS_FIELD_NUMBER: _ClassVar[int] + CACHE_MISSES_FIELD_NUMBER: _ClassVar[int] + total_chunks: int + total_tensor_bytes: int + response_sha256: bytes + timing: ServerTiming + cache_hits: int + cache_misses: int + def __init__( + self, + total_chunks: _Optional[int] = ..., + total_tensor_bytes: _Optional[int] = ..., + response_sha256: _Optional[bytes] = ..., + timing: _Optional[_Union[ServerTiming, _Mapping]] = ..., + cache_hits: _Optional[int] = ..., + cache_misses: _Optional[int] = ..., + ) -> None: ... + +class GenerateResponse(_message.Message): + __slots__ = ("header", "tensor_chunk", "trailer") + HEADER_FIELD_NUMBER: _ClassVar[int] + TENSOR_CHUNK_FIELD_NUMBER: _ClassVar[int] + TRAILER_FIELD_NUMBER: _ClassVar[int] + header: GenerateHeader + tensor_chunk: TensorChunk + trailer: GenerateTrailer + def __init__( + self, + header: _Optional[_Union[GenerateHeader, _Mapping]] = ..., + tensor_chunk: _Optional[_Union[TensorChunk, _Mapping]] = ..., + trailer: _Optional[_Union[GenerateTrailer, _Mapping]] = ..., + ) -> None: ... diff --git a/cosmos_framework/protos/reasoner_features/v1/reasoner_features_pb2_grpc.py b/cosmos_framework/protos/reasoner_features/v1/reasoner_features_pb2_grpc.py new file mode 100644 index 000000000..a5fd64868 --- /dev/null +++ b/cosmos_framework/protos/reasoner_features/v1/reasoner_features_pb2_grpc.py @@ -0,0 +1,155 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +# Generated by the gRPC Python protocol compiler plugin. DO NOT EDIT! +"""Client and server classes corresponding to protobuf-defined services.""" + +import grpc + +from cosmos_framework.protos.reasoner_features.v1 import ( + reasoner_features_pb2 as cosmos__framework_dot_protos_dot_reasoner__features_dot_v1_dot_reasoner__features__pb2, +) + +GRPC_GENERATED_VERSION = "1.78.0" +GRPC_VERSION = grpc.__version__ +_version_not_supported = False + +try: + from grpc._utilities import first_version_is_lower + + _version_not_supported = first_version_is_lower(GRPC_VERSION, GRPC_GENERATED_VERSION) +except ImportError: + _version_not_supported = True + +if _version_not_supported: + raise RuntimeError( + f"The grpc package installed is at version {GRPC_VERSION}," + + " but the generated code in cosmos_framework/protos/reasoner_features/v1/reasoner_features_pb2_grpc.py depends on" + + f" grpcio>={GRPC_GENERATED_VERSION}." + + f" Please upgrade your grpc module to grpcio>={GRPC_GENERATED_VERSION}" + + f" or downgrade your generated code using grpcio-tools<={GRPC_VERSION}." + ) + + +class ReasonerFeatureServiceStub(object): + """Missing associated documentation comment in .proto file.""" + + def __init__(self, channel): + """Constructor. + + Args: + channel: A grpc.Channel. + """ + self.GetInfo = channel.unary_unary( + "/cosmos.reasoner_features.v1.ReasonerFeatureService/GetInfo", + request_serializer=cosmos__framework_dot_protos_dot_reasoner__features_dot_v1_dot_reasoner__features__pb2.GetInfoRequest.SerializeToString, + response_deserializer=cosmos__framework_dot_protos_dot_reasoner__features_dot_v1_dot_reasoner__features__pb2.GetInfoResponse.FromString, + _registered_method=True, + ) + self.Generate = channel.unary_stream( + "/cosmos.reasoner_features.v1.ReasonerFeatureService/Generate", + request_serializer=cosmos__framework_dot_protos_dot_reasoner__features_dot_v1_dot_reasoner__features__pb2.GenerateRequest.SerializeToString, + response_deserializer=cosmos__framework_dot_protos_dot_reasoner__features_dot_v1_dot_reasoner__features__pb2.GenerateResponse.FromString, + _registered_method=True, + ) + + +class ReasonerFeatureServiceServicer(object): + """Missing associated documentation comment in .proto file.""" + + def GetInfo(self, request, context): + """Missing associated documentation comment in .proto file.""" + context.set_code(grpc.StatusCode.UNIMPLEMENTED) + context.set_details("Method not implemented!") + raise NotImplementedError("Method not implemented!") + + def Generate(self, request, context): + """Missing associated documentation comment in .proto file.""" + context.set_code(grpc.StatusCode.UNIMPLEMENTED) + context.set_details("Method not implemented!") + raise NotImplementedError("Method not implemented!") + + +def add_ReasonerFeatureServiceServicer_to_server(servicer, server): + rpc_method_handlers = { + "GetInfo": grpc.unary_unary_rpc_method_handler( + servicer.GetInfo, + request_deserializer=cosmos__framework_dot_protos_dot_reasoner__features_dot_v1_dot_reasoner__features__pb2.GetInfoRequest.FromString, + response_serializer=cosmos__framework_dot_protos_dot_reasoner__features_dot_v1_dot_reasoner__features__pb2.GetInfoResponse.SerializeToString, + ), + "Generate": grpc.unary_stream_rpc_method_handler( + servicer.Generate, + request_deserializer=cosmos__framework_dot_protos_dot_reasoner__features_dot_v1_dot_reasoner__features__pb2.GenerateRequest.FromString, + response_serializer=cosmos__framework_dot_protos_dot_reasoner__features_dot_v1_dot_reasoner__features__pb2.GenerateResponse.SerializeToString, + ), + } + generic_handler = grpc.method_handlers_generic_handler( + "cosmos.reasoner_features.v1.ReasonerFeatureService", rpc_method_handlers + ) + server.add_generic_rpc_handlers((generic_handler,)) + server.add_registered_method_handlers("cosmos.reasoner_features.v1.ReasonerFeatureService", rpc_method_handlers) + + +# This class is part of an EXPERIMENTAL API. +class ReasonerFeatureService(object): + """Missing associated documentation comment in .proto file.""" + + @staticmethod + def GetInfo( + request, + target, + options=(), + channel_credentials=None, + call_credentials=None, + insecure=False, + compression=None, + wait_for_ready=None, + timeout=None, + metadata=None, + ): + return grpc.experimental.unary_unary( + request, + target, + "/cosmos.reasoner_features.v1.ReasonerFeatureService/GetInfo", + cosmos__framework_dot_protos_dot_reasoner__features_dot_v1_dot_reasoner__features__pb2.GetInfoRequest.SerializeToString, + cosmos__framework_dot_protos_dot_reasoner__features_dot_v1_dot_reasoner__features__pb2.GetInfoResponse.FromString, + options, + channel_credentials, + insecure, + call_credentials, + compression, + wait_for_ready, + timeout, + metadata, + _registered_method=True, + ) + + @staticmethod + def Generate( + request, + target, + options=(), + channel_credentials=None, + call_credentials=None, + insecure=False, + compression=None, + wait_for_ready=None, + timeout=None, + metadata=None, + ): + return grpc.experimental.unary_stream( + request, + target, + "/cosmos.reasoner_features.v1.ReasonerFeatureService/Generate", + cosmos__framework_dot_protos_dot_reasoner__features_dot_v1_dot_reasoner__features__pb2.GenerateRequest.SerializeToString, + cosmos__framework_dot_protos_dot_reasoner__features_dot_v1_dot_reasoner__features__pb2.GenerateResponse.FromString, + options, + channel_credentials, + insecure, + call_credentials, + compression, + wait_for_ready, + timeout, + metadata, + _registered_method=True, + ) diff --git a/cosmos_framework/scripts/_train.py b/cosmos_framework/scripts/_train.py index 6532f40a7..8a43cf8ea 100644 --- a/cosmos_framework/scripts/_train.py +++ b/cosmos_framework/scripts/_train.py @@ -30,15 +30,16 @@ structure_config, ) from cosmos_framework.inference.common.init import init_output_dir, is_rank0 -from cosmos_framework.utils.flags import SMOKE +from cosmos_framework.model._base import close_model from cosmos_framework.trainer import ImaginaireTrainer from cosmos_framework.utils import log +from cosmos_framework.utils.flags import SMOKE if TYPE_CHECKING: from torch.utils.data import DataLoader - from cosmos_framework.utils.config import Config from cosmos_framework.model.generator.omni_mot_model import OmniMoTModel + from cosmos_framework.utils.config import Config def _validate_config_file(v: Path) -> Path: @@ -145,18 +146,25 @@ def train(args: Args) -> None: # Trainer init sets the rank-local CUDA device before tokenizers allocate weights. trainer: "ImaginaireTrainer" = config.trainer.type(config) model: "OmniMoTModel" = hydra.utils.instantiate(config.model) - dataloader_train: "DataLoader" = hydra.utils.instantiate(config.dataloader_train) - dataloader_val: "DataLoader" = hydra.utils.instantiate(config.dataloader_val) - - if args.dry_run: - return - - # Start training - trainer.train( - model=model, - dataloader_train=dataloader_train, - dataloader_val=dataloader_val, - ) + try: + dataloader_train: "DataLoader" = hydra.utils.instantiate(config.dataloader_train) + dataloader_val: "DataLoader" = hydra.utils.instantiate(config.dataloader_val) + + if args.dry_run: + close_model(model) + return + + # Start training + trainer.train( + model=model, + dataloader_train=dataloader_train, + dataloader_val=dataloader_val, + ) + except BaseException as error: + close_model(model, primary_error=error) + raise + else: + close_model(model) def main() -> None: diff --git a/cosmos_framework/scripts/extract_reasoner_features.py b/cosmos_framework/scripts/extract_reasoner_features.py index 188965796..e5e444778 100644 --- a/cosmos_framework/scripts/extract_reasoner_features.py +++ b/cosmos_framework/scripts/extract_reasoner_features.py @@ -34,15 +34,12 @@ import time import traceback from collections.abc import Iterator, Mapping, Sequence -from contextlib import contextmanager from dataclasses import asdict, dataclass from pathlib import Path from typing import Any import torch -from omegaconf import open_dict -from cosmos_framework.checkpoint.reasoner_only import load_reasoner_only_dcp from cosmos_framework.configs.toml_config.sft_config import load_experiment_from_toml from cosmos_framework.data.generator.local_datasets.sft_dataset import ( SFTDataset, @@ -52,7 +49,7 @@ from cosmos_framework.data.generator.local_datasets.sft_reasoner_documents import ( iter_sft_reasoner_documents, ) -from cosmos_framework.data.generator.sequence_packing.modalities import add_special_tokens +from cosmos_framework.data.generator.sequence_packing.modalities import add_special_tokens, compute_text_split_length from cosmos_framework.model.generator.reasoner_feature_cache import ( IncrementalReasonerFeatureCacheWriter, OfflineReasonerFeatureProvider, @@ -63,7 +60,11 @@ ) from cosmos_framework.model.generator.reasoner_features import ( ReasonerFeatureRequest, - extract_reasoner_feature_batch, +) +from cosmos_framework.model.generator.reasoner_runtime import ( + ReasonerFeatureRuntime, + ReasonerRuntimeSpec, + prepare_reasoner_model_config, ) from cosmos_framework.utils.lazy_config import instantiate as lazy_instantiate @@ -165,54 +166,8 @@ def _validate_extraction_config(config: object) -> None: raise ValueError(f"Reasoner extraction MVP supports the Nano Qwen3VLTextForCausalLM target, got {target_name}") -def _prepare_reasoner_model_config(config: object) -> object: - """Return a lazy Nano LM config that cannot instantiate the Generator or ViT.""" - - model = getattr(getattr(config, "model"), "config") - model_instance = copy.deepcopy(model.vlm_config.model_instance) - nested_config = model_instance["config"] - if isinstance(nested_config, dict): - nested_config.update( - include_gen_pathway=False, - include_und_pathway=True, - include_visual=False, - ) - else: - with open_dict(nested_config): - nested_config.include_gen_pathway = False - nested_config.include_und_pathway = True - nested_config.include_visual = False - return model_instance - - -@contextmanager -def _temporary_default_dtype(dtype: torch.dtype) -> Iterator[None]: - previous = torch.get_default_dtype() - torch.set_default_dtype(dtype) - try: - yield - finally: - torch.set_default_dtype(previous) - - -def _build_reasoner(config: object, *, device: torch.device, dtype: torch.dtype) -> torch.nn.Module: - """Construct a materialized Reasoner replica directly in its compute dtype. - - Direct construction intentionally avoids meta ``to_empty`` here: HuggingFace - rotary buffers are non-persistent and therefore absent from DCP. Constructing - on the final device initializes those buffers correctly before parameters are - overwritten by the strict Reasoner-only checkpoint load. - """ - - model_instance = _prepare_reasoner_model_config(config) - with _temporary_default_dtype(dtype), torch.device(device): - reasoner = lazy_instantiate(model_instance) - generation_parameters = [name for name, _ in reasoner.named_parameters() if "moe_gen" in name] - if generation_parameters: - raise RuntimeError(f"Reasoner-only construction retained Generator parameters: {generation_parameters}") - reasoner.requires_grad_(False) - reasoner.eval() - return reasoner +# Backward-compatible private alias retained for tests and downstream scripts. +_prepare_reasoner_model_config = prepare_reasoner_model_config def _special_tokens(dataset: SFTDataset) -> dict[str, int]: @@ -223,6 +178,15 @@ def _special_tokens(dataset: SFTDataset) -> dict[str, int]: return {**special_tokens, "eos_token_id": int(eos_token_id)} +def _extraction_max_total_tokens(dataset: SFTDataset, special_tokens: Mapping[str, int]) -> int: + """Derive the largest framed request from the dataset's caption limit.""" + + max_caption_tokens = dataset.max_caption_tokens + if not isinstance(max_caption_tokens, int) or isinstance(max_caption_tokens, bool) or max_caption_tokens <= 0: + raise ValueError(f"SFT max_caption_tokens must be a positive integer, got {max_caption_tokens!r}") + return compute_text_split_length(max_caption_tokens, dict(special_tokens), has_generation=True) + + def _rank_dataset(dataset: SFTDataset, *, rank: int, world_size: int) -> SFTDataset: """Shallow-copy a dataset and assign it a deterministic video-level slice.""" @@ -245,6 +209,7 @@ def _iter_rank_requests( world_size: int, use_float_positions: bool, max_documents_per_rank: int | None, + special_tokens: Mapping[str, int] | None = None, ) -> Iterator[ReasonerFeatureRequest]: """Yield framed requests for one deterministic metadata slice. @@ -253,7 +218,7 @@ def _iter_rank_requests( is published. Other requests stay attached to their diagnostic sample key. """ - tokens = _special_tokens(dataset) + tokens = dict(special_tokens) if special_tokens is not None else _special_tokens(dataset) produced = 0 if dataset.cfg_dropout_rate > 0 and not dataset.cfg_dropout_keep_metadata: null_text_ids, _ = dataset._tokenize_caption("") @@ -341,7 +306,7 @@ def _initialize_distributed(coordination_run_id: str | None = None) -> Distribut def _extract_rank( *, - reasoner: torch.nn.Module, + runtime: ReasonerFeatureRuntime, requests: Iterator[ReasonerFeatureRequest], writer: IncrementalReasonerFeatureCacheWriter, context: DistributedContext, @@ -354,7 +319,7 @@ def _extract_rank( if writer.contains(request.sample_key, request.fingerprint): stats.cache_hits += 1 continue - features = extract_reasoner_feature_batch(reasoner, [request]) + features = runtime.execute([request]) stats.extracted_tokens += request.token_ids.numel() if not writer.append(ReasonerFeatureCacheEntry(request.sample_key, features)): raise RuntimeError(f"Writer unexpectedly rejected newly extracted record {request.sample_key!r}") @@ -524,6 +489,8 @@ def _run(args: argparse.Namespace, context: DistributedContext) -> Path: return manifest dataset = _instantiate_sft_dataset(config) + special_tokens = _special_tokens(dataset) + max_total_tokens = _extraction_max_total_tokens(dataset, special_tokens) writer = IncrementalReasonerFeatureCacheWriter.resume( output, identity=identity, @@ -540,28 +507,33 @@ def _run(args: argparse.Namespace, context: DistributedContext) -> Path: world_size=context.world_size, use_float_positions=use_float_positions, max_documents_per_rank=args.max_documents_per_rank, + special_tokens=special_tokens, ) torch.cuda.reset_peak_memory_stats(context.device) - torch.cuda.synchronize(context.device) - reasoner_construction_started = time.perf_counter() - reasoner = _build_reasoner(config, device=context.device, dtype=dtype) - torch.cuda.synchronize(context.device) - reasoner_construction_seconds = time.perf_counter() - reasoner_construction_started - torch.cuda.synchronize(context.device) - checkpoint_load_started = time.perf_counter() - load_reasoner_only_dcp(reasoner, args.checkpoint, source=args.checkpoint_source) - torch.cuda.synchronize(context.device) - checkpoint_load_seconds = time.perf_counter() - checkpoint_load_started + runtime = ReasonerFeatureRuntime.load( + config, + ReasonerRuntimeSpec( + checkpoint=args.checkpoint, + checkpoint_source=args.checkpoint_source, + device=context.device, + dtype=dtype, + identity=identity, + max_requests=1, + max_total_tokens=max_total_tokens, + ), + ) publishable = args.max_documents_per_rank is None stats = _extract_rank( - reasoner=reasoner, + runtime=runtime, requests=requests, writer=writer, context=context, ) - stats.reasoner_construction_seconds = reasoner_construction_seconds - stats.checkpoint_load_seconds = checkpoint_load_seconds + if runtime.load_stats is None: + raise RuntimeError("Loaded Reasoner runtime did not report startup statistics") + stats.reasoner_construction_seconds = runtime.load_stats.construction_seconds + stats.checkpoint_load_seconds = runtime.load_stats.checkpoint_load_seconds _commit_rank_result(writer, stats, context=context, publishable=publishable) if not publishable: diff --git a/cosmos_framework/scripts/extract_reasoner_features_test.py b/cosmos_framework/scripts/extract_reasoner_features_test.py index 2cd941ea9..539837ffd 100644 --- a/cosmos_framework/scripts/extract_reasoner_features_test.py +++ b/cosmos_framework/scripts/extract_reasoner_features_test.py @@ -19,6 +19,7 @@ ExtractionStats, _atomic_write_json, _commit_rank_result, + _extraction_max_total_tokens, _find_sft_dataset_config, _initialize_distributed, _iter_rank_requests, @@ -98,6 +99,31 @@ def test_rank_requests_frame_tokens_and_deduplicate_standard_cfg_null(monkeypatc assert all(request.position_ids.dtype == torch.float32 for request in requests) +@pytest.mark.level(0) +@pytest.mark.gpus(0) +def test_extraction_token_limit_is_derived_from_dataset_and_framing() -> None: + dataset = SimpleNamespace(max_caption_tokens=8_192) + + assert ( + _extraction_max_total_tokens( # type: ignore[arg-type] + dataset, + {"eos_token_id": 1, "start_of_generation": 2}, + ) + == 8_194 + ) + assert ( + _extraction_max_total_tokens( # type: ignore[arg-type] + dataset, + {"bos_token_id": 0, "eos_token_id": 1, "start_of_generation": 2}, + ) + == 8_195 + ) + + dataset.max_caption_tokens = 0 + with pytest.raises(ValueError, match="max_caption_tokens must be a positive integer"): + _extraction_max_total_tokens(dataset, {}) # type: ignore[arg-type] + + def _config() -> SimpleNamespace: model_instance = { "_target_": "cosmos_framework.model.generator.mot.unified_mot.Qwen3VLTextForCausalLM", diff --git a/cosmos_framework/scripts/serve_reasoner_features.py b/cosmos_framework/scripts/serve_reasoner_features.py new file mode 100644 index 000000000..6ac621220 --- /dev/null +++ b/cosmos_framework/scripts/serve_reasoner_features.py @@ -0,0 +1,186 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +"""Start one independent GPU Reasoner feature-service replica. + +Start one process per service GPU. The processes do not create a PyTorch +distributed process group and must remain outside the Generator FSDP world. +Use an external TCP/gRPC load balancer when multiple replicas serve training +ranks. +""" + +from __future__ import annotations + +import argparse +import json +import signal +import threading +from dataclasses import asdict +from pathlib import Path +from typing import Sequence + +import torch + +from cosmos_framework.configs.toml_config.sft_config import load_experiment_from_toml +from cosmos_framework.model.generator.reasoner_features import ReasonerFeatureIdentity +from cosmos_framework.model.generator.reasoner_remote_server import ( + ReasonerFeatureService, + create_reasoner_grpc_server, +) +from cosmos_framework.model.generator.reasoner_runtime import ReasonerFeatureRuntime, ReasonerRuntimeSpec + +# The generic SFT loader resolves the complete training config, including +# dataloaders, the video VAE and the trainer's checkpoint input. None of those +# objects are used by the Reasoner-only runtime, but shipped recipes commonly +# populate them with ``${oc.env:...}`` expressions. Replace the unused +# subtrees before OmegaConf resolves the config so a service replica does not +# require training-only environment variables such as DATASET_PATH, +# WAN_VAE_PATH or BASE_CHECKPOINT_PATH. +_UNUSED_REASONER_SERVICE_PATH = "/unused/by-reasoner-service" +_REASONER_SERVICE_CONFIG_OVERRIDES = ( + "dataloader_train=null", + "dataloader_val=null", + f"checkpoint.load_path={_UNUSED_REASONER_SERVICE_PATH}", + f"model.config.tokenizer.vae_path={_UNUSED_REASONER_SERVICE_PATH}", +) + + +def _parse_args(argv: Sequence[str] | None = None) -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) + parser.add_argument("--sft-toml", required=True, help="Nano SFT recipe used to construct the Reasoner") + parser.add_argument("--checkpoint", required=True, help="Local DCP model component or iteration directory") + parser.add_argument("--checkpoint-source", choices=("regular", "ema"), default="regular") + parser.add_argument("--reasoner-fingerprint", required=True) + parser.add_argument("--tokenizer-fingerprint", required=True) + parser.add_argument("--framing-fingerprint", required=True) + parser.add_argument("--device", default="cuda", help="One visible service device, for example cuda or cuda:0") + parser.add_argument( + "--host", + default="127.0.0.1", + help="Bind address (default: loopback; use a non-loopback address only on a trusted private network)", + ) + parser.add_argument("--port", type=int, default=50051) + parser.add_argument("--max-requests", type=int, default=64) + parser.add_argument("--max-batch-tokens", type=int, default=4_096) + parser.add_argument("--max-queued-tokens", type=int, default=8_192) + parser.add_argument("--max-chunk-bytes", type=int, default=2 * 1024**2) + parser.add_argument("--max-rpc-workers", type=int, default=16) + parser.add_argument("--shutdown-grace-seconds", type=float, default=30.0) + parser.add_argument( + "overrides", + nargs=argparse.REMAINDER, + help="Hydra overrides applied after TOML; prefix the list with --", + ) + args = parser.parse_args(argv) + args.overrides = [item for item in args.overrides if item != "--"] + if not 1 <= args.port <= 65_535: + parser.error("--port must be in [1, 65535]") + for name in ("max_requests", "max_batch_tokens", "max_queued_tokens", "max_chunk_bytes", "max_rpc_workers"): + if getattr(args, name) <= 0: + parser.error(f"--{name.replace('_', '-')} must be positive") + if args.max_queued_tokens < args.max_batch_tokens: + parser.error("--max-queued-tokens must be at least --max-batch-tokens") + if args.max_chunk_bytes > 4 * 1024**2: + parser.error("--max-chunk-bytes must not exceed 4 MiB") + if args.shutdown_grace_seconds < 0: + parser.error("--shutdown-grace-seconds must be non-negative") + return args + + +def _load_reasoner_service_config(args: argparse.Namespace) -> object: + """Load only service-relevant settings from a complete SFT recipe. + + Internal pruning overrides intentionally follow user overrides. The + corresponding training subtrees cannot affect Reasoner construction, and + keeping the pruning last prevents an accidental CLI override from + reintroducing an unrelated environment interpolation. + """ + + return load_experiment_from_toml( + Path(args.sft_toml), + extra_overrides=[*args.overrides, *_REASONER_SERVICE_CONFIG_OVERRIDES], + ) + + +def _run(args: argparse.Namespace) -> None: + if not torch.cuda.is_available(): + raise RuntimeError("Reasoner feature serving requires CUDA") + device = torch.device(args.device) + if device.type != "cuda": + raise ValueError("Reasoner feature serving requires a CUDA --device") + config = _load_reasoner_service_config(args) + identity = ReasonerFeatureIdentity( + reasoner=args.reasoner_fingerprint, + tokenizer=args.tokenizer_fingerprint, + framing=args.framing_fingerprint, + ) + if device.index is not None: + torch.cuda.set_device(device) + torch.cuda.reset_peak_memory_stats(device) + runtime = ReasonerFeatureRuntime.load( + config, + ReasonerRuntimeSpec( + checkpoint=args.checkpoint, + checkpoint_source=args.checkpoint_source, + device=device, + dtype=torch.bfloat16, + identity=identity, + max_requests=args.max_requests, + max_total_tokens=args.max_batch_tokens, + ), + ) + service = ReasonerFeatureService( + runtime, + max_queued_tokens=args.max_queued_tokens, + max_chunk_bytes=args.max_chunk_bytes, + ) + server, bound_port = create_reasoner_grpc_server( + service, + address=f"{args.host}:{args.port}", + max_rpc_workers=args.max_rpc_workers, + ) + server.start() + startup = { + "status": "ready", + "address": f"{args.host}:{bound_port}", + "service_instance_id": service.service_instance_id, + "identity": asdict(identity), + "signature": { + "num_layers": runtime.signature.num_layers, + "num_kv_heads": runtime.signature.num_kv_heads, + "head_dim": runtime.signature.head_dim, + "dtype": str(runtime.signature.dtype), + }, + "limits": { + "max_requests": runtime.max_requests, + "max_batch_tokens": runtime.max_total_tokens, + "max_queued_tokens": args.max_queued_tokens, + "max_chunk_bytes": args.max_chunk_bytes, + }, + "load_stats": None if runtime.load_stats is None else asdict(runtime.load_stats), + } + print(json.dumps(startup, indent=2, sort_keys=True), flush=True) + stop_requested = threading.Event() + + def request_stop(_signum: int, _frame: object) -> None: + stop_requested.set() + + previous_sigterm = signal.signal(signal.SIGTERM, request_stop) + try: + while not stop_requested.wait(timeout=1.0): + if not server.wait_for_termination(timeout=0): + break + except KeyboardInterrupt: + pass + finally: + signal.signal(signal.SIGTERM, previous_sigterm) + server.stop(args.shutdown_grace_seconds).wait() + + +def main(argv: Sequence[str] | None = None) -> int: + _run(_parse_args(argv)) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/cosmos_framework/scripts/serve_reasoner_features_test.py b/cosmos_framework/scripts/serve_reasoner_features_test.py new file mode 100644 index 000000000..c38d43333 --- /dev/null +++ b/cosmos_framework/scripts/serve_reasoner_features_test.py @@ -0,0 +1,49 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +from __future__ import annotations + +from pathlib import Path + +import pytest + +from cosmos_framework.model.generator.reasoner_runtime import prepare_reasoner_model_config +from cosmos_framework.scripts.serve_reasoner_features import ( + _UNUSED_REASONER_SERVICE_PATH, + _load_reasoner_service_config, + _parse_args, +) + +_RECIPE = Path(__file__).parents[2] / "examples" / "toml" / "sft_config" / "vision_sft_nano.toml" + + +def test_service_config_does_not_require_training_only_environment(monkeypatch: pytest.MonkeyPatch) -> None: + for name in ("DATASET_PATH", "WAN_VAE_PATH", "BASE_CHECKPOINT_PATH"): + monkeypatch.delenv(name, raising=False) + args = _parse_args( + [ + "--sft-toml", + str(_RECIPE), + "--checkpoint", + "/checkpoints/Cosmos3-Nano", + "--reasoner-fingerprint", + "reasoner", + "--tokenizer-fingerprint", + "tokenizer", + "--framing-fingerprint", + "framing", + "--", + "model.config.vlm_config.model_instance.config.qk_norm_for_text=false", + ] + ) + + assert args.host == "127.0.0.1" + config = _load_reasoner_service_config(args) + + model_instance = prepare_reasoner_model_config(config) + assert str(model_instance["_target_"]).endswith("Qwen3VLTextForCausalLM") + assert model_instance["config"]["qk_norm_for_text"] is False + assert config.dataloader_train is None + assert config.dataloader_val is None + assert config.checkpoint.load_path == _UNUSED_REASONER_SERVICE_PATH + assert config.model.config.tokenizer.vae_path == _UNUSED_REASONER_SERVICE_PATH diff --git a/cosmos_framework/scripts/train.py b/cosmos_framework/scripts/train.py index 7e1375033..850882fd0 100644 --- a/cosmos_framework/scripts/train.py +++ b/cosmos_framework/scripts/train.py @@ -31,15 +31,15 @@ import torch from loguru import logger as logging -from cosmos_framework.utils.config import Config -from cosmos_framework.utils.lazy_config import LazyConfig, instantiate -from cosmos_framework.utils.serialization import to_yaml +from cosmos_framework.configs.toml_config.sft_config import load_experiment_from_toml +from cosmos_framework.model._base import close_model from cosmos_framework.utils import distributed +from cosmos_framework.utils.config import Config from cosmos_framework.utils.context_managers import data_loader_init, distributed_init, model_init from cosmos_framework.utils.launch import log_reproducible_setup +from cosmos_framework.utils.lazy_config import LazyConfig, instantiate +from cosmos_framework.utils.serialization import to_yaml from cosmos_framework.utils.training_telemetry import telemetry -from cosmos_framework.configs.toml_config.sft_config import load_experiment_from_toml - # --------------------------------------------------------------------------- # --deterministic: mirrors launch_vfm.sh determinism settings. @@ -211,18 +211,23 @@ def launch(config: Config, args: argparse.Namespace) -> None: with model_init(): model = instantiate(config.model) - - # Create the dataloaders. - with data_loader_init(): - dataloader_train = instantiate(config.dataloader_train) - dataloader_val = instantiate(config.dataloader_val) - - # Start training - trainer.train( - model, - dataloader_train, - dataloader_val, - ) + try: + # Create the dataloaders. + with data_loader_init(): + dataloader_train = instantiate(config.dataloader_train) + dataloader_val = instantiate(config.dataloader_val) + + # Start training + trainer.train( + model, + dataloader_train, + dataloader_val, + ) + except BaseException as error: + close_model(model, primary_error=error) + raise + else: + close_model(model) if __name__ == "__main__": diff --git a/cosmos_framework/trainer/__init__.py b/cosmos_framework/trainer/__init__.py index 94aea9449..e33a6f7fd 100644 --- a/cosmos_framework/trainer/__init__.py +++ b/cosmos_framework/trainer/__init__.py @@ -13,9 +13,13 @@ import torch.distributed as dist import torch.utils.data -from cosmos_framework.utils.flags import INTERNAL from cosmos_framework.utils.context_managers import distributed_init -from cosmos_framework.utils.profiling import maybe_enable_memory_snapshot, maybe_enable_nsys_profiling, maybe_enable_profiling +from cosmos_framework.utils.flags import INTERNAL +from cosmos_framework.utils.profiling import ( + maybe_enable_memory_snapshot, + maybe_enable_nsys_profiling, + maybe_enable_profiling, +) try: from megatron.core import parallel_state @@ -25,14 +29,13 @@ USE_MEGATRON = False -from cosmos_framework.utils.lazy_config import LazyConfig, instantiate -from cosmos_framework.model._base import ImaginaireModel +from cosmos_framework.model._base import ImaginaireModel, close_model from cosmos_framework.utils import callback, distributed, ema, log, misc from cosmos_framework.utils.checkpointer import Checkpointer +from cosmos_framework.utils.lazy_config import LazyConfig, instantiate from cosmos_framework.utils.misc import StragglerDetectorV2 - @dataclass class ContextParallelDataWindow: """Caches one dataloader batch across a ``cp_size``-step CP data window. @@ -267,6 +270,22 @@ def train( model: ImaginaireModel, dataloader_train: torch.utils.data.DataLoader, dataloader_val: torch.utils.data.DataLoader, + ) -> None: + """Run training and always release model-owned external resources.""" + + try: + self._train(model, dataloader_train, dataloader_val) + except BaseException as error: + close_model(model, primary_error=error) + raise + else: + close_model(model) + + def _train( + self, + model: ImaginaireModel, + dataloader_train: torch.utils.data.DataLoader, + dataloader_val: torch.utils.data.DataLoader, ) -> None: """The training function. diff --git a/docs/nano_sft_decoupled_reasoner.md b/docs/nano_sft_decoupled_reasoner.md index a93794b71..75ddfdc6e 100644 --- a/docs/nano_sft_decoupled_reasoner.md +++ b/docs/nano_sft_decoupled_reasoner.md @@ -1,7 +1,8 @@ # Decoupling the frozen Reasoner from Nano generator SFT -Status: Phase 1 and the Phase 2 offline extraction MVP are implemented; an -eight-H20 short-run FSDP/EMA A/B is complete, while full-corpus and +Status: Phase 1, the Phase 2 offline extraction MVP, and the Phase 3 remote +transport MVP are implemented; an eight-H20 short-run offline FSDP/EMA A/B is +complete, while real-GPU remote capacity tests, full-corpus extraction, and compile-enabled production validation remain pending Target recipe: `vision_sft_nano` and later Nano multiview SFT variants @@ -31,15 +32,19 @@ Implemented in this branch: - strict external-backend fingerprints plus cache/model layer, KV-head, and head-dimension validation before FSDP materialization; - offline-provider submission before noising and synchronized failure handling - before the FSDP forward; and + before the FSDP forward; +- a shared Reasoner-only runtime, versioned gRPC protocol, asynchronous remote + provider, bounded single-GPU service, layer/K/V chunk streaming, end-to-end + checksums, strict startup handshake, deadlines, and transient retries; and - CPU contract/storage tests plus a tiny H20 CUDA test comparing all generator outputs and gradients between the full cached model and the structurally pruned model. Not implemented yet: -- remote and read-through clients/services (the backend names and `endpoint` - field are reserved, but selecting either backend currently fails explicitly); +- the `read_through` cache/service composition and persistent miss writes; +- token-bucket dynamic batching inside a Reasoner replica (the remote MVP has a + bounded queue but deliberately executes variable-length requests serially); - layerwise H2D staging (`layerwise_h2d=true` is rejected); - automatic content-digest derivation for Reasoner/tokenizer/dataset artifacts; the CLI currently requires the three pinned fingerprints explicitly; and @@ -56,11 +61,11 @@ cross-attention K/V tensors. The configuration contract names these backends: 2. `inline` captures Reasoner K/V locally, then runs a second cached GEN-only pass; it is the numerical-reference path. 3. `offline` reads precomputed immutable Reasoner K/V shards and is implemented. -4. `remote` will obtain the same tensors asynchronously from dedicated Reasoner workers. +4. `remote` obtains the same tensors asynchronously from dedicated Reasoner workers. 5. `read_through` will check the offline cache first and send misses to remote workers. -Only `joint`, `inline`, and `offline` are executable today. `remote` and -`read_through` intentionally raise `NotImplementedError` during provider creation. +`joint`, `inline`, `offline`, and `remote` are executable today. `read_through` +intentionally raises `NotImplementedError` during provider creation. For a fixed SFT corpus, `offline` should be the default. It removes the Reasoner from every training rank, is deterministic, and turns Reasoner work into a one-time dataset @@ -144,12 +149,12 @@ bytes per UND token ``` | Framed UND tokens | BF16 K/V per example | -|---:|---:| -| 256 | 36 MiB | -| 512 | 72 MiB | -| 1,024 | 144 MiB | -| 1,790 | 251.7 MiB | -| 2,048 | 288 MiB | +| ----------------: | -------------------: | +| 256 | 36 MiB | +| 512 | 72 MiB | +| 1,024 | 144 MiB | +| 1,790 | 251.7 MiB | +| 2,048 | 288 MiB | The local full BridgeData manifest has 1,222 examples. With the recipe's real Qwen tokenization, it has 1,082 framed tokens on average (p50 1,004, p95 1,586, maximum @@ -159,12 +164,12 @@ remote or read-through backend becomes attractive. The Nano language model contains approximately: -| Component | Parameters | BF16 logical size | -|---|---:|---:| -| GEN layer pathway | 6.946B | 12.94 GiB | -| UND layer pathway | 6.946B | 12.94 GiB | -| UND embeddings + LM head + final norm | 1.245B | 2.32 GiB | -| Total removable Reasoner | 8.191B | 15.26 GiB | +| Component | Parameters | BF16 logical size | +| ------------------------------------- | ---------: | ----------------: | +| GEN layer pathway | 6.946B | 12.94 GiB | +| UND layer pathway | 6.946B | 12.94 GiB | +| UND embeddings + LM head + final norm | 1.245B | 2.32 GiB | +| Total removable Reasoner | 8.191B | 15.26 GiB | A real Nano meta-model audit measured 15,136,811,008 parameters in the full language model and 6,946,075,648 after pruning: 8,190,735,360 parameters @@ -202,16 +207,21 @@ class ReasonerFeatureBatch: fingerprints: tuple[str, ...] = () class ReasonerFeatureProvider(Protocol): + @property + def signature(self) -> ReasonerFeatureSignature: ... + def submit( self, requests: Sequence[ReasonerFeatureRequest] ) -> Future[ReasonerFeatureBatch]: ... ``` -`submit` is a future-shaped contract for every provider. The current offline provider -performs its local read synchronously and returns an already-completed future. The -training path submits after final sequence packing and before noising, then resolves -the future before entering the FSDP forward. A future asynchronous provider can use -the same seam to overlap provider work with noising or other preparation. +`submit` is a future-shaped contract for every provider. The offline provider performs +its local read synchronously and returns an already-completed future; the remote +provider snapshots the small CPU request tensors and immediately schedules a streaming +RPC on a private worker thread. The training path submits after final sequence packing +and before noising, then resolves the future before entering the FSDP forward. This +currently overlaps noising and packed-sequence H2D only; VAE encode has already +completed. Moving request construction before VAE encode is a later performance step. `extract_reasoner_feature_batch(causal_lm, requests)` is the current UND-only reference extractor. The immutable cache APIs are in @@ -347,13 +357,14 @@ cannot prove whether EMA was loaded or left randomly initialized. The shipped full Nano DCP was audited against the pruned vision-SFT target: all 405 target tensors (397 language-model GEN tensors plus 8 VFM tensors) exist with matching shapes, and strict DCP subset loading succeeds while ignoring source-only -UND tensors. This validates the model-only warm-start shape contract; a distributed -generator-only save/optimizer/resume integration test is still pending. +UND tensors. Phase 3 additionally completed a real seven-rank remote-conditioning +optimizer step, wrote a generator-only DCP, and loaded that DCP through a +resume-only reshard gate. Resume-and-continue parity is still pending. Still to validate before production use: -- full-checkpoint to generator-only warm start for both DCP and safetensors; -- generator-only save/resume while preserving immutable cache identity; +- full-checkpoint to generator-only warm start from safetensors; +- generator-only resume-and-continue while preserving immutable cache identity; - reconstruction of a full inference model from the generator checkpoint and pinned Reasoner; and - mismatched generator-checkpoint/cache provenance rejection on resume. @@ -479,8 +490,9 @@ Cache the exact text feature, not only a video UUID. For the current SFT dataset ## Remote backend -This section is a design target. No remote transport, client, service, or -read-through write path is implemented in the current branch. +The remote MVP is implemented. The read-through write path, TLS/authentication, +content-addressed server cache, dynamic batching, and production load-balancer +deployment remain future work. Run Reasoner workers as a separate service allocation, not as ranks in the generator's FSDP/DDP process group. Mixing service ranks into the training world size would make @@ -491,29 +503,94 @@ Recommended data flow: ```text generator rank -- submit(token IDs, positions, fingerprint) --> request queue | reasoner replica - +-- VAE encode + noise + pack (overlap) | - |<----------- per-layer K/V or cache URI -----------------+ + +-- noise + packed H2D (limited overlap) | + |<------ checksummed layer/K/V chunks over gRPC -----------+ +-- validate fingerprint --> GEN-only forward/backward ``` +The service is one independent process with one complete Reasoner replica and no +PyTorch process group. `GetInfo` performs a fail-closed protocol, identity, dtype, and +tensor-geometry handshake before the Generator model is materialized. `Generate` is a +server-streaming RPC: its header pins request order and offsets, each tensor chunk has a +SHA-256 digest, and the trailer covers raw tensors in canonical +`K0 || V0 || K1 || V1 ...` order. The client discards every partial response. + +The v1 wire request preserves the complete `causal_offsets` array so later per-view +captions do not require a protocol redesign. The current runtime advertises only the +`single_document` capability and explicitly rejects multi-document requests; 7/11-view +cached attention and per-view-caption parity remain Phase 4 work. + +Install the optional transport dependency and start one service replica only after its +Reasoner checkpoint has been made locally available: + +```bash +uv sync --extra train --extra reasoner-remote + +CUDA_VISIBLE_DEVICES=0 LD_LIBRARY_PATH='' \ +python -m cosmos_framework.scripts.serve_reasoner_features \ + --sft-toml examples/toml/sft_config/vision_sft_nano.toml \ + --checkpoint /shared/checkpoints/iter_000000100 \ + --checkpoint-source regular \ + --reasoner-fingerprint \ + --tokenizer-fingerprint \ + --framing-fingerprint \ + --host 127.0.0.1 --port 50051 +``` + +The service loader resolves only the Reasoner-relevant part of the SFT recipe. +It prunes the unused training dataloaders, video VAE, and trainer checkpoint +input before OmegaConf resolution, so `DATASET_PATH`, `WAN_VAE_PATH`, and +`BASE_CHECKPOINT_PATH` are not required for this command. `--checkpoint` is the +only checkpoint path used by the service runtime. + +The MVP transport uses insecure gRPC and has no client authentication. It binds +to loopback by default; use a non-loopback `--host` only inside an isolated, +trusted network. Checksums detect corrupted payloads but do not protect against +tampering or unauthorized requests. TLS and authentication are required before +exposing a service across hosts or behind a production load balancer. + +Point the Generator recipe at that service: + +```toml +[model.reasoner_conditioning] +backend = "remote" +endpoint = "reasoner-lb.example:50051" +reasoner_fingerprint = "" +tokenizer_fingerprint = "" +framing_fingerprint = "" +strict_fingerprint = true +connect_timeout_s = 30.0 +request_timeout_s = 300.0 +request_max_retries = 2 +retry_backoff_s = 0.25 +``` + +Retries are bounded by one absolute request deadline and apply only to gRPC +`UNAVAILABLE` and `RESOURCE_EXHAUSTED`. Identity, protocol, fingerprint, checksum, +shape, and dtype failures are terminal. There is no silent fallback to a joint +Reasoner. + Operational requirements: - replicate the 8B Reasoner one per service GPU; prefer replica/data parallelism over tensor parallelism because Nano fits on one modern accelerator; -- dynamic-batch requests by total UND tokens, not request count; +- admit and queue by total UND tokens, not request count; the MVP serializes GPU + execution, and a later dynamic batcher must preserve variable-length causal isolation; - use immutable request IDs, bounded queues, backpressure, deadlines, and idempotent retries; -- keep a memory and/or NVMe content-addressed cache in front of Reasoner execution; -- expose queue time, Reasoner compute time, serialization time, bytes sent, cache-hit - rate, and client wait time; and -- fail closed on version mismatches. An unavailable service may fall back to an exact - disk hit, but not silently to a different Reasoner or quantization. +- add a memory and/or NVMe content-addressed cache in front of Reasoner execution; +- expose queue time, Reasoner compute time, D2H time, CPU stream-preparation time, and + bytes through the response trailer; production percentile export and cache-hit metrics are + still pending; and +- fail closed on version mismatches. A future read-through provider may use an exact + disk hit, but must not silently switch to a different Reasoner or quantization. For a 1,024-token prompt the BF16 response is 144 MiB. At `R` examples/s, required payload bandwidth is approximately `144 * R MiB/s`, before transport framing. This is usually modest relative to long-video generator step time, but it must be measured at the intended number of training ranks. Start with a simple streaming RPC into pinned -CPU buffers plus asynchronous H2D copies. Add CUDA IPC/NVLink for same-node workers or +CPU buffers. Pinned staging and layerwise asynchronous H2D are not yet implemented. +Add CUDA IPC/NVLink for same-node workers or CUDA-aware UCX/RDMA for cross-node workers only if profiling shows transport on the critical path. @@ -527,25 +604,29 @@ generator nodes. Keep the feature contract in training/model code; a training job must not import Ray or other heavyweight inference-only dependencies. -| Area | Current status | -|---|---| -| `configs/base/defaults/model_config.py` | Implemented typed runtime configuration and compatibility validation. | -| `configs/toml_config/sft_config.py` | Implemented the flat `[model.reasoner_conditioning]` TOML schema. | -| `model/generator/mot/unified_mot.py` | Implemented optional UND construction, GEN-only execution, and `prune_und_pathway_`. | -| `model/generator/mot/cosmos3_vfm_network.py` | Implemented generator-only packed layout without text embedding. | -| `model/generator/omni_mot_model.py` | Implemented inline/offline lifecycle, pre-noise submission, synchronized resolution, structural pruning, and startup guards. Remote/read-through remain pending. | -| `model/generator/reasoner_features.py` | Implemented request/result types, UND-only extraction, inline/static memory, and exact two-way external-K/V attention. | -| `model/generator/reasoner_feature_cache.py` | Implemented immutable cache identity, fingerprints, bounded/resumable sharded writers, distributed finalization, manifest validation, and offline provider. | -| `model/generator/mot/multiview_attention.py` | Pending cache-aware per-view/maskless/Flex support. | -| `data/generator/local_datasets/sft_reasoner_documents.py` | Implemented finite deterministic standard-SFT document enumeration with training framing parity. | -| `data/generator/` | Centralized multi-batch prefetch/cache-handle plumbing remains pending. | -| `checkpoint/reasoner_only.py` | Implemented strict regular/EMA Reasoner-only local-DCP loading with independent reads and safe Generator/visual source-leaf omission; generator resume/composition validation remains pending. | -| `scripts/extract_reasoner_features.py` | Implemented resumable distributed shared-POSIX cache extraction and atomic publication without NCCL collectives. | -| `examples/toml/sft_config/` | Pending opt-in Nano cached-Reasoner recipe after full-scale parity passes. | - -An optional remote server can live under the inference/serving tree, but it should -implement the neutral wire schema without making the training client depend on that -server implementation. +| Area | Current status | +| --------------------------------------------------------- | ------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------ | +| `configs/base/defaults/model_config.py` | Implemented typed runtime configuration and compatibility validation. | +| `configs/toml_config/sft_config.py` | Implemented the flat `[model.reasoner_conditioning]` TOML schema. | +| `model/generator/mot/unified_mot.py` | Implemented optional UND construction, GEN-only execution, and `prune_und_pathway_`. | +| `model/generator/mot/cosmos3_vfm_network.py` | Implemented generator-only packed layout without text embedding. | +| `model/generator/omni_mot_model.py` | Implemented inline/offline/remote lifecycle, pre-noise submission, synchronized resolution, structural pruning, and provider-neutral startup signature guards. Read-through remains pending. | +| `model/generator/reasoner_features.py` | Implemented identity/fingerprint/signature/request/result types, UND-only extraction, inline/static memory, and exact two-way external-K/V attention. | +| `model/generator/reasoner_feature_cache.py` | Implemented immutable cache identity, fingerprints, bounded/resumable sharded writers, distributed finalization, manifest validation, and offline provider. | +| `model/generator/reasoner_runtime.py` | Implemented shared Reasoner-only construction, strict DCP restore, request admission/fingerprint validation, and serialized execution for offline extraction and serving. | +| `model/generator/reasoner_remote.py` | Implemented tensor codec, protocol validation, streaming response assembly, checksums, asynchronous client, deadlines, and bounded transient retries. | +| `model/generator/reasoner_remote_server.py` | Implemented one-replica token-bounded gRPC service, single-GPU execution, timing metadata, backpressure, and unhealthy-after-OOM/runtime-invariant behavior. | +| `protos/reasoner_features/v1/` | Implemented versioned `GetInfo` and server-streaming `Generate` protobuf schema and generated Python bindings. | +| `model/generator/mot/multiview_attention.py` | Pending cache-aware per-view/maskless/Flex support. | +| `data/generator/local_datasets/sft_reasoner_documents.py` | Implemented finite deterministic standard-SFT document enumeration with training framing parity. | +| `data/generator/` | Centralized multi-batch prefetch/cache-handle plumbing remains pending. | +| `checkpoint/reasoner_only.py` | Implemented strict regular/EMA Reasoner-only local-DCP loading with independent reads, safe Generator/visual source-leaf omission, and a seven-rank generator-only save/resume-load smoke; resume-and-continue/composition validation remains pending. | +| `scripts/extract_reasoner_features.py` | Implemented resumable distributed shared-POSIX cache extraction and atomic publication without NCCL collectives. | +| `scripts/serve_reasoner_features.py` | Implemented one-process/one-GPU remote service launch after load-before-listen readiness. | +| `examples/toml/sft_config/` | Pending opt-in Nano cached-Reasoner recipe after full-scale parity passes. | + +The training model imports only the provider/client module when `backend="remote"`; +it does not import the server implementation or Ray/vLLM inference infrastructure. ## Implementation sequence @@ -613,13 +694,60 @@ offline reduced mean step time by 17.05%, increased aggregate token throughput b the expected direction and magnitude, but the hot 1.18-GiB cache is not a proxy for full-corpus shared-filesystem behavior. -### Phase 3: remote/read-through provider — pending +### Phase 3: remote provider — MVP implemented; read-through pending + +Completed: + +1. Reused the provider-neutral request/result, signature, and fingerprint schema. +2. Added a versioned protobuf handshake and chunked server-streaming gRPC transport. +3. Added the asynchronous client/provider, absolute deadlines, bounded transient + retries, strict integrity validation, and synchronized distributed failure surface. +4. Added a shared Reasoner runtime and one-process/one-GPU service with token-bounded + backpressure and no Generator process-group membership. +5. Added CPU codec/corruption/retry tests and a localhost in-process gRPC end-to-end + test. +6. Loaded the real 36-layer Nano DCP into independent H20 replicas and verified all K/V + tensors bitwise across direct extraction and the localhost remote service. +7. Added all-rank fail-closed gates for provider startup, full feature-signature + validation, synchronous request construction/submission, and asynchronous + resolution before any rank enters the corresponding FSDP collective. +8. Completed one real optimizer step with one dedicated Reasoner H20 and seven + Generator FSDP H20s, including warm start, forward/backward, optimizer/EMA, + and a generator-only checkpoint save. + +Still pending: -1. Reuse the implemented request/result and fingerprint schema. -2. Add remote client/service transport and asynchronous prefetch. -3. Benchmark one service GPU against increasing generator-rank counts and scale - replicas based on measured queueing, not a fixed assumed ratio. -4. Add persistent read-through writes only after cache-key and atomicity tests pass. +1. Benchmark one service GPU against 1/2/4/8 Generator ranks and scale replicas from + measured queue, compute, D2H, network, client-wait, and Generator-step percentiles. +2. Repeat parity and throughput with representative/max-length prompts and batched + requests rather than the short correctness smoke. +3. Move request submission early enough to overlap VAE work, then re-run throughput + parity. +4. Implement token-budget dynamic batching only after variable-length output parity. +5. Add persistent read-through writes only after cache-key and atomicity tests pass. + +The real-model correctness smoke used one H20 for the service and a second H20 for the +direct reference, the shipped `Cosmos3-Nano` regular DCP, BF16, and a 10-token framed +request. Every K/V tensor across 36 layers was bitwise equal. Direct extraction took +0.293 s; localhost remote round-trip took 0.326 s. The server reported 0.288 s compute, +2.0 ms D2H, 3.0 ms CPU stream preparation, and negligible queue time. The service +allocated about 15.26 GiB after loading and peaked at about 15.40 GiB reserved. These numbers +validate the transport and runtime, not capacity: the prompt was intentionally short, +the network path was loopback, and there were no concurrent Generator ranks. + +The end-to-end training smoke reserved physical GPU 0 for the Reasoner and ran a +seven-rank Generator FSDP job on GPUs 1--7 with a 16,384-token packing cap. One +optimizer step completed at loss 0.2512 with no CUDA OOM or RPC error. Generator +CUDA allocator peaks were 20.41--20.46 GiB allocated and 21.47--21.51 GiB reserved; +physical peaks were 23.20--25.11 GiB. The Reasoner peaked at 18.49 GiB. The reported +71.79-second iteration included a 31.37-second, 104-GiB checkpoint save, so it is +not comparable to the earlier steady-state eight-rank benchmark. The saved model +contains only 405 live plus 405 EMA Generator leaves and no Reasoner embedding, +LM-head, or UND-expert leaves. Raw logs and telemetry are under +`outputs/nano_sft_reasoner_benchmark/remote_smoke_7gpu_1step_20260922_v1`. +After the final all-rank signature gate was added, the same seven-rank topology +also completed a resume-only load of that checkpoint in 38.70 seconds without +an error or additional checkpoint write. ### Phase 4: broaden support — pending @@ -635,17 +763,22 @@ Validated in the current implementation: - inline capture followed by static GEN-only replay; - tiny dense Qwen output and all-generator-gradient parity between the full cached model and `prune_und_pathway_` model on one H20; -- absence of UND parameters from the structurally pruned state dict; and +- absence of UND parameters from the structurally pruned state dict; - immutable-cache round trips, batching across shards, content-addressed prompt reuse, cache misses, stale fingerprints, malformed manifests, and corrupt shard - checksums. + checksums; and +- remote BF16 tensor codec and multi-chunk round trips, corrupt checksum/order + rejection, idempotent transient retry behavior, identity mismatch rejection, and a + localhost gRPC provider/service end-to-end test. Remaining correctness gates: - one optimizer step produces equivalent selected weights; - normal prompt, CFG-null prompt, maximum-length prompt, and per-view captions; - full checkpoint to generator-only warm start, generator-only resume, and full-model - inference composition; and + inference composition; +- representative/max-length remote K/V parity, service restart/OOM behavior, and + capacity scaling; and - end-to-end stale cache/config provenance is rejected during distributed startup and resume, before any FSDP forward. diff --git a/docs/nano_sft_remote_reasoner_test_report.md b/docs/nano_sft_remote_reasoner_test_report.md new file mode 100644 index 000000000..f00d04140 --- /dev/null +++ b/docs/nano_sft_remote_reasoner_test_report.md @@ -0,0 +1,218 @@ +# Nano SFT remote Reasoner: Phase 3 test report + +Date: 2026-09-22 + +Branch: `feat/nano-sft-remote-reasoner` + +Base commit: `bd6ec1c` (`feat(training): decouple Nano Reasoner with offline K/V conditioning`) + +## Scope + +This report covers the Phase 3 MVP that removes the frozen Reasoner from +Generator training ranks and obtains the same per-layer K/V tensors from a +separate one-process/one-GPU gRPC service. + +The implemented validation surface includes: + +- shared offline/remote Reasoner construction and strict DCP loading; +- versioned identity, feature-signature, capability, and limit handshake; +- exact BF16 tensor codec, canonical `K0, V0, K1, V1, ...` ordering, per-chunk + and whole-response SHA-256 validation; +- asynchronous provider submission with one absolute deadline and bounded + transient retries; +- request/token admission limits, bounded gRPC concurrency, serialized GPU + execution, and CPU staging before streaming; +- cancellation/deadline handling without queued ghost compute; +- OOM and permanent runtime-invariant unhealthy-state propagation to + already-queued requests; +- provider cleanup on constructor, dataloader, dry-run, training-success, and + training-failure paths; +- all-rank startup failure synchronization before any Generator rank enters + FSDP construction; and +- fail-closed rejection of multiple causal documents (including the current + per-view-caption representation) and specialized multiview attention metadata. + +## Latest CPU regression run + +Command: + +```bash +LD_LIBRARY_PATH='' .venv/bin/python -m pytest -q \ + cosmos_framework/model/generator/reasoner_remote_test.py \ + cosmos_framework/model/generator/reasoner_remote_lifecycle_test.py \ + cosmos_framework/model/generator/reasoner_remote_server_test.py \ + cosmos_framework/model/generator/reasoner_runtime_test.py \ + cosmos_framework/model/generator/reasoner_features_test.py \ + cosmos_framework/model/generator/reasoner_feature_cache_test.py \ + cosmos_framework/scripts/extract_reasoner_features_test.py \ + cosmos_framework/scripts/serve_reasoner_features_test.py \ + cosmos_framework/model/generator/omni_mot_reasoner_conditioning_test.py \ + cosmos_framework/configs/toml_config/sft_config_test.py \ + cosmos_framework/trainer/distillation_test.py \ + cosmos_framework/model/model_lifecycle_test.py +``` + +Result: **186 passed in 65.67 seconds**. + +The run emitted 31 existing `PytestUnknownMarkWarning` messages from +`trainer/distillation_test.py` for the legacy `L0`/`CPU` markers. There were no +test failures, skips caused by this feature, OOMs, or hangs. + +New failure-mode regressions specifically verify: + +- executor concurrency overflow returns `RESOURCE_EXHAUSTED`; +- explicit cancellation and deadline expiry do not execute a queued request; +- an OOM marks the replica unhealthy before the next queued request can run; +- an invalid runtime output signature or fingerprint order marks the replica + unhealthy before queued and subsequent requests can execute; +- malformed protocol/identity/dtype/order/shape/checksum inputs fail closed; +- partial streams are cancelled and token reservations are released; +- channel shutdown interrupts an in-flight RPC and retry backoff; +- constructor/handshake failures close both channel and executor; +- a rank-local startup/handshake failure is surfaced on every distributed rank, + while providers created successfully on peer ranks are closed; +- rank-local provider/identity/sample-key/request-build/submit validation errors + are deferred into failed futures and surfaced through the same all-rank + pre-forward failure gate; and +- stream timing excludes client/network backpressure. + +The service-config regression additionally unsets `DATASET_PATH`, +`WAN_VAE_PATH`, and `BASE_CHECKPOINT_PATH` before loading the shipped Nano SFT +recipe. The service now prunes its unused dataloader, video-VAE, and training +checkpoint subtrees before OmegaConf resolution, while preserving explicit +Reasoner model overrides. A Reasoner-only replica therefore no longer requires +those training-only environment variables. + +## Static validation + +- Ruff check and format check passed for every modified or added Python file. +- Pyrefly reported **0 errors** (15 existing suppressions); it also reported the + existing `ignore-missing-source` extra-key warning from `pyrefly.toml`. +- `git diff --check` passed. +- The standard `uv-lock` pre-commit hook and `uv 0.11.14 uv lock --check` + passed without skipping or rewriting the lock file. + +## Issue encountered and resolved + +The first real service launch failed during generic SFT-config resolution, +before the Reasoner was constructed, because the complete training recipe also +resolved dataset, VAE, and training-checkpoint environment variables. Supplying +those variables confirmed the GPU path, but they are unrelated to Reasoner +serving. The service-specific loader now replaces those unused subtrees before +resolution, and the no-environment regression above prevents the dependency +from returning. + +## Real H20 correctness smoke + +The shipped `examples/checkpoints/Cosmos3-Nano` regular DCP was loaded into two +independent H20 processes: one direct `ReasonerFeatureRuntime` reference and one +localhost gRPC service. The request was a 10-token BF16 framed prompt. + +| Measurement | Result | +| --------------------------------- | -------------------------------------: | +| Reasoner layers compared | 36 | +| K/V equality | Bitwise equal for every K and V tensor | +| Direct extraction | 0.293149 s | +| Remote localhost round trip | 0.326337 s | +| Server Reasoner compute | 0.288081 s | +| Server device-to-host copy | 0.002015 s | +| Server CPU stream preparation | 0.003000 s | +| Service allocated after load | 16,383,586,816 B (15.26 GiB) | +| Service peak reserved during load | 16,536,043,520 B (15.40 GiB) | + +The final smoke also verified the advertised four-request preflight limit. The +service process was stopped after the test and both GPUs returned to zero +reported allocation. This is a correctness smoke, not a capacity result: it +uses a short prompt, loopback networking, one request, and no concurrent +Generator ranks. + +## Real seven-rank Generator training smoke + +The complete training path was then exercised with GPU 0 reserved for one +Reasoner service replica and physical GPUs 1--7 running a seven-rank Generator +FSDP job. The job used the official eight-video BridgeData sample, a 16,384-token +packing cap, BF16, full activation checkpointing, FP32 EMA, gradient accumulation +one, eager execution, and one optimizer step. It warm-started from the shipped +Nano DCP and saved a new seven-way DCP. + +| Measurement | Result | +| ------------------------------------------- | ---------------------------: | +| Optimizer steps | 1 | +| Final loss | 0.2512 | +| Reported iteration time | 71.79 s | +| Checkpoint portion of iteration | 31.37 s | +| End-to-end `torchrun` wall time | 101.35 s | +| Generator peak allocated, min/mean/max | 20.414 / 20.441 / 20.459 GiB | +| Generator peak reserved, min/mean/max | 21.467 / 21.487 / 21.506 GiB | +| Generator physical peak, `nvidia-smi` | 23.196--25.110 GiB/GPU | +| Reasoner physical peak, `nvidia-smi` | 18.493 GiB | +| Saved checkpoint size | 104 GiB | +| CUDA OOM / RPC failure / training exception | 0 | + +The saved model metadata contains 810 leaves: 405 live Generator leaves and +405 EMA Generator leaves. It contains no `embed_tokens`, `lm_head`, or +`moe_und` parameter, confirming that the training checkpoint does not recreate +the frozen Reasoner. After graceful shutdown, all eight GPUs returned to zero +reported memory use. + +After adding the final all-rank feature-signature gate, the same seven-rank job +was started again against a fresh service process and resumed the saved DCP at +iteration 1. It completed the resume-only path in 38.70 seconds without entering +another training step or writing another checkpoint. This exercises the final +provider handshake, signature collective, seven-way DCP reshard/load, and +provider shutdown success path; all GPUs again returned to zero memory use. + +Artifacts: + +- run root: + `outputs/nano_sft_reasoner_benchmark/remote_smoke_7gpu_1step_20260922_v1`; +- Generator log: `logs/generator_train.log`; +- Reasoner startup log: `logs/reasoner_service.log`; +- 500 ms physical-GPU telemetry: `logs/nvidia_smi.csv`; +- final-code resume gate: `logs/generator_resume_gate.log` and + `logs/reasoner_service_final_gate.log`; and +- checkpoint: + `train/cosmos3/sft/nano_remote_smoke_7gpu/checkpoints/iter_000000001`. + +These raw artifacts are local-only: the run root is under the gitignored +`outputs/` directory and therefore does not travel with this branch. The +measurements and conclusions above are committed in this report; archive the +run root separately when raw logs or the 104 GiB checkpoint must be handed off. + +This is an end-to-end correctness and resource smoke, not a throughput comparison +with the earlier eight-rank runs: it uses seven Generator ranks, a shorter packing +cap, one step, and the reported iteration time includes checkpoint writing. + +## Existing offline A/B baseline + +The earlier 8-H20 short-run comparison remains the performance baseline: + +| Backend | Mean step time | Mean peak allocated per Generator GPU | +| ---------------------------- | -------------: | ------------------------------------: | +| Joint Reasoner + Generator | 57.887 s | 52.128 GiB | +| Offline K/V + Generator only | 48.015 s | 36.498 GiB | + +Offline conditioning reduced mean step time by 17.05% and peak allocated memory +by 15.63 GiB per Generator GPU in that short run. Remote conditioning has the +same Generator-side structural pruning, but its end-to-end training throughput +must still be measured under realistic service concurrency and networking. + +## Remaining gates + +- Run 1/2/4/8 concurrent Generator ranks against one service replica and record + queue, compute, D2H, stream preparation, network/client wait, and training-step + percentiles. +- Repeat remote parity with representative and maximum-length prompts. +- Compare the remote step's loss and selected gradients against an identically + configured `inline` or `offline` run. +- Derive and verify the Reasoner, tokenizer, and framing fingerprints from the + actual checkpoint and tokenizer artifacts. The MVP strictly compares the + operator-supplied labels, but does not yet prove that a label is the content + digest of the artifact it names. +- Add TLS/authentication and deployment health/metrics before cross-host + production use. The MVP binds to loopback by default and non-loopback serving + is restricted operationally to an isolated, trusted network. +- Add multiple causal-document and specialized cached-attention support before + using 7-view/11-view per-view captions; protocol v1 currently rejects the + multi-document input deliberately. Shared-caption multiview also needs its + separate attention-layout parity gate. diff --git a/pyproject.toml b/pyproject.toml index 6115f6c11..3b04e3a79 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -60,6 +60,10 @@ guardrail = [ ] interactive = [ ] +reasoner-remote = [ + "grpcio>=1.78", + "protobuf>=6.31.1", +] serve = [ "fastapi", "httpx", diff --git a/uv.lock b/uv.lock index 0a05554a4..74af8c644 100644 --- a/uv.lock +++ b/uv.lock @@ -113,15 +113,17 @@ conflicts = [[ ], [ { package = "cosmos-framework", group = "cu128" }, { package = "cosmos-framework", group = "cu128-train" }, - { package = "cosmos-framework", group = "cu130" }, + { package = "cosmos-framework", group = "cu130-torch213" }, { package = "cosmos-framework", group = "cu130-torch213-train" }, { package = "cosmos-framework", group = "cu130-train" }, ], [ + { package = "cosmos-framework", group = "cu128" }, + { package = "cosmos-framework", group = "cu128-train" }, + { package = "cosmos-framework", group = "cu130" }, { package = "cosmos-framework", group = "cu130-torch213-train" }, - { package = "cosmos-framework", group = "vllm" }, + { package = "cosmos-framework", group = "cu130-train" }, ], [ { package = "cosmos-framework", group = "cu128-train" }, - { package = "cosmos-framework", group = "cu130" }, { package = "cosmos-framework", group = "cu130-torch213" }, { package = "cosmos-framework", group = "cu130-torch213-train" }, { package = "cosmos-framework", group = "cu130-train" }, @@ -133,11 +135,11 @@ conflicts = [[ ], [ { package = "cosmos-framework", group = "cu128" }, { package = "cosmos-framework", group = "cu128-train" }, - { package = "cosmos-framework", group = "cu130-torch213" }, { package = "cosmos-framework", group = "cu130-torch213-train" }, { package = "cosmos-framework", group = "cu130-train" }, ], [ { package = "cosmos-framework", group = "cu128-train" }, + { package = "cosmos-framework", group = "cu130" }, { package = "cosmos-framework", group = "cu130-torch213" }, { package = "cosmos-framework", group = "cu130-torch213-train" }, { package = "cosmos-framework", group = "cu130-train" }, @@ -146,10 +148,8 @@ conflicts = [[ { package = "cosmos-framework", group = "cu130-torch213-train" }, { package = "cosmos-framework", group = "cu130-train" }, ], [ - { package = "cosmos-framework", group = "cu128" }, - { package = "cosmos-framework", group = "cu128-train" }, { package = "cosmos-framework", group = "cu130-torch213-train" }, - { package = "cosmos-framework", group = "cu130-train" }, + { package = "cosmos-framework", group = "vllm" }, ]] [manifest] @@ -1666,12 +1666,16 @@ guardrail = [ { name = "retinaface-py" }, { name = "sentencepiece" }, ] +reasoner-remote = [ + { name = "grpcio" }, + { name = "protobuf" }, +] serve = [ { name = "fastapi" }, { name = "gradio" }, { name = "httpx" }, - { name = "ray", version = "2.46.0", source = { registry = "https://pypi.org/simple" }, extra = ["serve"], marker = "python_full_version < '3.11' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra != 'group-16-cosmos-framework-cu130-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "ray", version = "2.55.1", source = { registry = "https://pypi.org/simple" }, extra = ["serve"], marker = "python_full_version >= '3.12' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra != 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "ray", version = "2.46.0", source = { registry = "https://pypi.org/simple" }, extra = ["serve"], marker = "python_full_version < '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "ray", version = "2.55.1", source = { registry = "https://pypi.org/simple" }, extra = ["serve"], marker = "python_full_version >= '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, ] train = [ { name = "aioboto3" }, @@ -1974,6 +1978,7 @@ requires-dist = [ { name = "fvcore", marker = "extra == 'train'" }, { name = "glfw", marker = "extra == 'train'" }, { name = "gradio", marker = "extra == 'serve'" }, + { name = "grpcio", marker = "extra == 'reasoner-remote'", specifier = ">=1.78" }, { name = "h5py", marker = "extra == 'train'" }, { name = "httpx", marker = "extra == 'serve'" }, { name = "hydra-core" }, @@ -2017,6 +2022,7 @@ requires-dist = [ { name = "polars", marker = "extra == 'train'" }, { name = "polyscope", marker = "extra == 'train'" }, { name = "protobuf", marker = "extra == 'guardrail'" }, + { name = "protobuf", marker = "extra == 'reasoner-remote'", specifier = ">=6.31.1" }, { name = "psycopg2-binary", marker = "extra == 'train'" }, { name = "py3nvml", marker = "extra == 'train'" }, { name = "pycocotools", marker = "extra == 'train'" }, @@ -2061,7 +2067,7 @@ requires-dist = [ { name = "xatlas", marker = "extra == 'train'" }, { name = "zarr", marker = "extra == 'train'" }, ] -provides-extras = ["guardrail", "interactive", "serve", "train"] +provides-extras = ["guardrail", "interactive", "reasoner-remote", "serve", "train"] [package.metadata.requires-dev] cu128 = [ @@ -8236,9 +8242,9 @@ name = "opentelemetry-exporter-prometheus" version = "0.61b0" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "opentelemetry-api", marker = "python_full_version >= '3.12' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra != 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "opentelemetry-sdk", marker = "python_full_version >= '3.12' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra != 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "prometheus-client", marker = "python_full_version >= '3.12' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra != 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "opentelemetry-api", marker = "python_full_version >= '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "opentelemetry-sdk", marker = "python_full_version >= '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "prometheus-client", marker = "python_full_version >= '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, ] sdist = { url = "https://files.pythonhosted.org/packages/4a/20/9e818fd364d12e8d0cfdce4a3b2d82e24d98c4ceebb315de6b6770b5f214/opentelemetry_exporter_prometheus-0.61b0.tar.gz", hash = "sha256:7c4919bd8e79abd62b610767e80f42c9c3a06c5183f4dd9141eedeb57aea284b", size = 15136, upload-time = "2026-03-04T14:17:26.275Z" } wheels = [ @@ -10142,22 +10148,19 @@ name = "ray" version = "2.46.0" source = { registry = "https://pypi.org/simple" } resolution-markers = [ - "python_full_version == '3.11.*' and sys_platform == 'linux'", - "python_full_version == '3.11.*' and sys_platform != 'darwin' and sys_platform != 'linux'", - "python_full_version == '3.11.*' and sys_platform == 'darwin'", "python_full_version < '3.11' and sys_platform == 'linux'", "python_full_version < '3.11' and sys_platform != 'darwin' and sys_platform != 'linux'", "python_full_version < '3.11' and sys_platform == 'darwin'", ] dependencies = [ - { name = "click", marker = "python_full_version < '3.11' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra != 'group-16-cosmos-framework-cu130-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "filelock", marker = "python_full_version < '3.11' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra != 'group-16-cosmos-framework-cu130-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "jsonschema", marker = "python_full_version < '3.11' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra != 'group-16-cosmos-framework-cu130-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "msgpack", marker = "python_full_version < '3.11' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra != 'group-16-cosmos-framework-cu130-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "packaging", marker = "python_full_version < '3.11' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra != 'group-16-cosmos-framework-cu130-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "protobuf", marker = "python_full_version < '3.11' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra != 'group-16-cosmos-framework-cu130-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "pyyaml", marker = "python_full_version < '3.11' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra != 'group-16-cosmos-framework-cu130-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "requests", marker = "python_full_version < '3.11' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra != 'group-16-cosmos-framework-cu130-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "click", marker = "python_full_version < '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "filelock", marker = "python_full_version < '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "jsonschema", marker = "python_full_version < '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "msgpack", marker = "python_full_version < '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "packaging", marker = "python_full_version < '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "protobuf", marker = "python_full_version < '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "pyyaml", marker = "python_full_version < '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "requests", marker = "python_full_version < '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, ] wheels = [ { url = "https://files.pythonhosted.org/packages/25/a2/0cc3dc138d149dbbf69a90ded986118ae9420a5ac0dbf99e0c5883152950/ray-2.46.0-cp310-cp310-macosx_10_15_x86_64.whl", hash = "sha256:719244b84df79502e5f09497f256618d94d78d66fbaf229422008a0568d3a0ff", size = 68482152, upload-time = "2025-05-07T21:04:19.538Z" }, @@ -10183,21 +10186,21 @@ wheels = [ [package.optional-dependencies] serve = [ - { name = "aiohttp", marker = "python_full_version < '3.11' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra != 'group-16-cosmos-framework-cu130-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "aiohttp-cors", marker = "python_full_version < '3.11' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra != 'group-16-cosmos-framework-cu130-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "colorful", marker = "python_full_version < '3.11' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra != 'group-16-cosmos-framework-cu130-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "fastapi", marker = "python_full_version < '3.11' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra != 'group-16-cosmos-framework-cu130-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "grpcio", marker = "python_full_version < '3.11' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra != 'group-16-cosmos-framework-cu130-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "opencensus", marker = "python_full_version < '3.11' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra != 'group-16-cosmos-framework-cu130-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "prometheus-client", marker = "python_full_version < '3.11' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra != 'group-16-cosmos-framework-cu130-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "py-spy", marker = "python_full_version < '3.11' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra != 'group-16-cosmos-framework-cu130-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "pydantic", marker = "python_full_version < '3.11' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra != 'group-16-cosmos-framework-cu130-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "requests", marker = "python_full_version < '3.11' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra != 'group-16-cosmos-framework-cu130-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "smart-open", marker = "python_full_version < '3.11' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra != 'group-16-cosmos-framework-cu130-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "starlette", marker = "python_full_version < '3.11' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra != 'group-16-cosmos-framework-cu130-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "uvicorn", extra = ["standard"], marker = "python_full_version < '3.11' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra != 'group-16-cosmos-framework-cu130-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "virtualenv", marker = "python_full_version < '3.11' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra != 'group-16-cosmos-framework-cu130-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "watchfiles", marker = "python_full_version < '3.11' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version >= '3.12' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra != 'group-16-cosmos-framework-cu130-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "aiohttp", marker = "python_full_version < '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "aiohttp-cors", marker = "python_full_version < '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "colorful", marker = "python_full_version < '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "fastapi", marker = "python_full_version < '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "grpcio", marker = "python_full_version < '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "opencensus", marker = "python_full_version < '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "prometheus-client", marker = "python_full_version < '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "py-spy", marker = "python_full_version < '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "pydantic", marker = "python_full_version < '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "requests", marker = "python_full_version < '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "smart-open", marker = "python_full_version < '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "starlette", marker = "python_full_version < '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "uvicorn", extra = ["standard"], marker = "python_full_version < '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "virtualenv", marker = "python_full_version < '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "watchfiles", marker = "python_full_version < '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, ] [[package]] @@ -10213,7 +10216,10 @@ resolution-markers = [ "python_full_version >= '3.13' and platform_machine != 'aarch64' and platform_machine != 'x86_64' and sys_platform != 'linux' and extra != 'group-16-cosmos-framework-cu128' and extra != 'group-16-cosmos-framework-cu128-train' and extra != 'group-16-cosmos-framework-cu130' and extra != 'group-16-cosmos-framework-cu130-torch213' and extra != 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train'", "python_full_version == '3.12.*' and sys_platform == 'linux' and extra != 'group-16-cosmos-framework-cu128' and extra != 'group-16-cosmos-framework-cu128-train' and extra != 'group-16-cosmos-framework-cu130' and extra != 'group-16-cosmos-framework-cu130-torch213' and extra != 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train'", "python_full_version == '3.12.*' and sys_platform != 'darwin' and sys_platform != 'linux' and extra != 'group-16-cosmos-framework-cu128' and extra != 'group-16-cosmos-framework-cu128-train' and extra != 'group-16-cosmos-framework-cu130' and extra != 'group-16-cosmos-framework-cu130-torch213' and extra != 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train'", + "python_full_version == '3.11.*' and sys_platform == 'linux' and extra != 'group-16-cosmos-framework-cu128' and extra != 'group-16-cosmos-framework-cu128-train' and extra != 'group-16-cosmos-framework-cu130' and extra != 'group-16-cosmos-framework-cu130-torch213' and extra != 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train'", + "python_full_version == '3.11.*' and sys_platform != 'darwin' and sys_platform != 'linux' and extra != 'group-16-cosmos-framework-cu128' and extra != 'group-16-cosmos-framework-cu128-train' and extra != 'group-16-cosmos-framework-cu130' and extra != 'group-16-cosmos-framework-cu130-torch213' and extra != 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train'", "python_full_version == '3.12.*' and sys_platform == 'darwin' and extra != 'group-16-cosmos-framework-cu128' and extra != 'group-16-cosmos-framework-cu128-train' and extra != 'group-16-cosmos-framework-cu130' and extra != 'group-16-cosmos-framework-cu130-torch213' and extra != 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train'", + "python_full_version == '3.11.*' and sys_platform == 'darwin' and extra != 'group-16-cosmos-framework-cu128' and extra != 'group-16-cosmos-framework-cu128-train' and extra != 'group-16-cosmos-framework-cu130' and extra != 'group-16-cosmos-framework-cu130-torch213' and extra != 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train'", "python_full_version >= '3.13' and platform_machine == 'x86_64' and sys_platform == 'linux' and extra != 'group-16-cosmos-framework-cu128' and extra != 'group-16-cosmos-framework-cu128-train' and extra != 'group-16-cosmos-framework-cu130' and extra != 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra != 'group-16-cosmos-framework-cu130-train' and extra != 'group-16-cosmos-framework-vllm'", "python_full_version >= '3.13' and platform_machine == 'x86_64' and sys_platform != 'linux' and extra != 'group-16-cosmos-framework-cu128' and extra != 'group-16-cosmos-framework-cu128-train' and extra != 'group-16-cosmos-framework-cu130' and extra != 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra != 'group-16-cosmos-framework-cu130-train' and extra != 'group-16-cosmos-framework-vllm'", "python_full_version >= '3.13' and platform_machine != 'x86_64' and sys_platform == 'linux' and extra != 'group-16-cosmos-framework-cu128' and extra != 'group-16-cosmos-framework-cu128-train' and extra != 'group-16-cosmos-framework-cu130' and extra != 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra != 'group-16-cosmos-framework-cu130-train' and extra != 'group-16-cosmos-framework-vllm'", @@ -10275,14 +10281,14 @@ resolution-markers = [ "python_full_version == '3.11.*' and sys_platform == 'darwin' and extra != 'group-16-cosmos-framework-cu128' and extra != 'group-16-cosmos-framework-cu128-train' and extra != 'group-16-cosmos-framework-cu130' and extra != 'group-16-cosmos-framework-cu130-torch213' and extra != 'group-16-cosmos-framework-cu130-torch213-train' and extra != 'group-16-cosmos-framework-cu130-train'", ] dependencies = [ - { name = "click", marker = "python_full_version >= '3.12' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra != 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "filelock", marker = "python_full_version >= '3.12' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra != 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "jsonschema", marker = "python_full_version >= '3.12' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra != 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "msgpack", marker = "python_full_version >= '3.12' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra != 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "packaging", marker = "python_full_version >= '3.12' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra != 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "protobuf", marker = "python_full_version >= '3.12' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra != 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "pyyaml", marker = "python_full_version >= '3.12' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra != 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "requests", marker = "python_full_version >= '3.12' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra != 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "click", marker = "python_full_version >= '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "filelock", marker = "python_full_version >= '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "jsonschema", marker = "python_full_version >= '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "msgpack", marker = "python_full_version >= '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "packaging", marker = "python_full_version >= '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "protobuf", marker = "python_full_version >= '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "pyyaml", marker = "python_full_version >= '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "requests", marker = "python_full_version >= '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, ] wheels = [ { url = "https://files.pythonhosted.org/packages/7e/d0/a85097dd53aaca1a44acc4dd0b3d2c0e9233179433e2ee326e4018ab3cf7/ray-2.55.1-cp310-cp310-macosx_12_0_arm64.whl", hash = "sha256:2d5786661e192148719accc959def6cdcabd7a24cd9008005bf3d0e3c8cfd529", size = 65829601, upload-time = "2026-04-22T20:09:10.013Z" }, @@ -10304,24 +10310,24 @@ wheels = [ [package.optional-dependencies] serve = [ - { name = "aiohttp", marker = "python_full_version >= '3.12' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra != 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "aiohttp-cors", marker = "python_full_version >= '3.12' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra != 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "colorful", marker = "python_full_version >= '3.12' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra != 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "fastapi", marker = "python_full_version >= '3.12' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra != 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "grpcio", marker = "python_full_version >= '3.12' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra != 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "opencensus", marker = "python_full_version >= '3.12' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra != 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "opentelemetry-exporter-prometheus", marker = "python_full_version >= '3.12' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra != 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "opentelemetry-proto", marker = "python_full_version >= '3.12' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra != 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "opentelemetry-sdk", marker = "python_full_version >= '3.12' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra != 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "prometheus-client", marker = "python_full_version >= '3.12' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra != 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "py-spy", marker = "python_full_version >= '3.12' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra != 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "pydantic", marker = "python_full_version >= '3.12' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra != 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "requests", marker = "python_full_version >= '3.12' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra != 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "smart-open", marker = "python_full_version >= '3.12' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra != 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "starlette", marker = "python_full_version >= '3.12' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra != 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "uvicorn", extra = ["standard"], marker = "python_full_version >= '3.12' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra != 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "virtualenv", marker = "python_full_version >= '3.12' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra != 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, - { name = "watchfiles", marker = "python_full_version >= '3.12' or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version == '3.11.*' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version == '3.11.*' and extra != 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (python_full_version < '3.11' and extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "aiohttp", marker = "python_full_version >= '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "aiohttp-cors", marker = "python_full_version >= '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "colorful", marker = "python_full_version >= '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "fastapi", marker = "python_full_version >= '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "grpcio", marker = "python_full_version >= '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "opencensus", marker = "python_full_version >= '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "opentelemetry-exporter-prometheus", marker = "python_full_version >= '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "opentelemetry-proto", marker = "python_full_version >= '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "opentelemetry-sdk", marker = "python_full_version >= '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "prometheus-client", marker = "python_full_version >= '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "py-spy", marker = "python_full_version >= '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "pydantic", marker = "python_full_version >= '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "requests", marker = "python_full_version >= '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "smart-open", marker = "python_full_version >= '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "starlette", marker = "python_full_version >= '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "uvicorn", extra = ["standard"], marker = "python_full_version >= '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "virtualenv", marker = "python_full_version >= '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, + { name = "watchfiles", marker = "python_full_version >= '3.11' or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu128-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu128-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-torch213-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213' and extra == 'group-16-cosmos-framework-vllm') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-cu130-train') or (extra == 'group-16-cosmos-framework-cu130-torch213-train' and extra == 'group-16-cosmos-framework-vllm')" }, ] [[package]] From 40373ca299652510b2da7dd9b8a5a9e13379754e Mon Sep 17 00:00:00 2001 From: Codex Date: Tue, 22 Sep 2026 09:14:41 +0800 Subject: [PATCH 3/4] docs(training): record nonzero remote Reasoner updates --- docs/nano_sft_decoupled_reasoner.md | 134 ++++++++++------ docs/nano_sft_remote_reasoner_test_report.md | 156 ++++++++++++++++--- 2 files changed, 220 insertions(+), 70 deletions(-) diff --git a/docs/nano_sft_decoupled_reasoner.md b/docs/nano_sft_decoupled_reasoner.md index 75ddfdc6e..9f945a5cf 100644 --- a/docs/nano_sft_decoupled_reasoner.md +++ b/docs/nano_sft_decoupled_reasoner.md @@ -357,14 +357,25 @@ cannot prove whether EMA was loaded or left randomly initialized. The shipped full Nano DCP was audited against the pruned vision-SFT target: all 405 target tensors (397 language-model GEN tensors plus 8 VFM tensors) exist with matching shapes, and strict DCP subset loading succeeds while ignoring source-only -UND tensors. Phase 3 additionally completed a real seven-rank remote-conditioning -optimizer step, wrote a generator-only DCP, and loaded that DCP through a -resume-only reshard gate. Resume-and-continue parity is still pending. +UND tensors. Phase 3 first completed a seven-rank remote execution smoke whose +scheduled optimizer call used effective LR 0. That run produced gradients and +Adam moment state, wrote a generator-only DCP, and passed a resume-only reshard +gate, but 0/405 live Generator tensors changed. It validated the execution and +checkpoint structure, not a learned parameter update. + +A subsequent two-iteration run performed its first nonzero-LR update at `2e-6`; +all 405 live Generator tensors changed relative to the warm-start checkpoint. The +same job then restored model, optimizer, scheduler, and trainer state from +iteration 2 and continued through iteration 3 at LR `4e-6`; all 405 live tensors +changed again. Operational resume-and-continue is therefore validated. The run +did not restore a dataloader cursor because no dataloader checkpoint was written, +so exact datastream-position resume and uninterrupted-versus-resumed numerical +parity remain pending. Still to validate before production use: - full-checkpoint to generator-only warm start from safetensors; -- generator-only resume-and-continue while preserving immutable cache identity; +- exact datastream-position resume and uninterrupted-versus-resumed numerical parity; - reconstruction of a full inference model from the generator checkpoint and pinned Reasoner; and - mismatched generator-checkpoint/cache provenance rejection on resume. @@ -604,26 +615,26 @@ generator nodes. Keep the feature contract in training/model code; a training job must not import Ray or other heavyweight inference-only dependencies. -| Area | Current status | -| --------------------------------------------------------- | ------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------ | -| `configs/base/defaults/model_config.py` | Implemented typed runtime configuration and compatibility validation. | -| `configs/toml_config/sft_config.py` | Implemented the flat `[model.reasoner_conditioning]` TOML schema. | -| `model/generator/mot/unified_mot.py` | Implemented optional UND construction, GEN-only execution, and `prune_und_pathway_`. | -| `model/generator/mot/cosmos3_vfm_network.py` | Implemented generator-only packed layout without text embedding. | -| `model/generator/omni_mot_model.py` | Implemented inline/offline/remote lifecycle, pre-noise submission, synchronized resolution, structural pruning, and provider-neutral startup signature guards. Read-through remains pending. | -| `model/generator/reasoner_features.py` | Implemented identity/fingerprint/signature/request/result types, UND-only extraction, inline/static memory, and exact two-way external-K/V attention. | -| `model/generator/reasoner_feature_cache.py` | Implemented immutable cache identity, fingerprints, bounded/resumable sharded writers, distributed finalization, manifest validation, and offline provider. | -| `model/generator/reasoner_runtime.py` | Implemented shared Reasoner-only construction, strict DCP restore, request admission/fingerprint validation, and serialized execution for offline extraction and serving. | -| `model/generator/reasoner_remote.py` | Implemented tensor codec, protocol validation, streaming response assembly, checksums, asynchronous client, deadlines, and bounded transient retries. | -| `model/generator/reasoner_remote_server.py` | Implemented one-replica token-bounded gRPC service, single-GPU execution, timing metadata, backpressure, and unhealthy-after-OOM/runtime-invariant behavior. | -| `protos/reasoner_features/v1/` | Implemented versioned `GetInfo` and server-streaming `Generate` protobuf schema and generated Python bindings. | -| `model/generator/mot/multiview_attention.py` | Pending cache-aware per-view/maskless/Flex support. | -| `data/generator/local_datasets/sft_reasoner_documents.py` | Implemented finite deterministic standard-SFT document enumeration with training framing parity. | -| `data/generator/` | Centralized multi-batch prefetch/cache-handle plumbing remains pending. | -| `checkpoint/reasoner_only.py` | Implemented strict regular/EMA Reasoner-only local-DCP loading with independent reads, safe Generator/visual source-leaf omission, and a seven-rank generator-only save/resume-load smoke; resume-and-continue/composition validation remains pending. | -| `scripts/extract_reasoner_features.py` | Implemented resumable distributed shared-POSIX cache extraction and atomic publication without NCCL collectives. | -| `scripts/serve_reasoner_features.py` | Implemented one-process/one-GPU remote service launch after load-before-listen readiness. | -| `examples/toml/sft_config/` | Pending opt-in Nano cached-Reasoner recipe after full-scale parity passes. | +| Area | Current status | +| --------------------------------------------------------- | ---------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | +| `configs/base/defaults/model_config.py` | Implemented typed runtime configuration and compatibility validation. | +| `configs/toml_config/sft_config.py` | Implemented the flat `[model.reasoner_conditioning]` TOML schema. | +| `model/generator/mot/unified_mot.py` | Implemented optional UND construction, GEN-only execution, and `prune_und_pathway_`. | +| `model/generator/mot/cosmos3_vfm_network.py` | Implemented generator-only packed layout without text embedding. | +| `model/generator/omni_mot_model.py` | Implemented inline/offline/remote lifecycle, pre-noise submission, synchronized resolution, structural pruning, and provider-neutral startup signature guards. Read-through remains pending. | +| `model/generator/reasoner_features.py` | Implemented identity/fingerprint/signature/request/result types, UND-only extraction, inline/static memory, and exact two-way external-K/V attention. | +| `model/generator/reasoner_feature_cache.py` | Implemented immutable cache identity, fingerprints, bounded/resumable sharded writers, distributed finalization, manifest validation, and offline provider. | +| `model/generator/reasoner_runtime.py` | Implemented shared Reasoner-only construction, strict DCP restore, request admission/fingerprint validation, and serialized execution for offline extraction and serving. | +| `model/generator/reasoner_remote.py` | Implemented tensor codec, protocol validation, streaming response assembly, checksums, asynchronous client, deadlines, and bounded transient retries. | +| `model/generator/reasoner_remote_server.py` | Implemented one-replica token-bounded gRPC service, single-GPU execution, timing metadata, backpressure, and unhealthy-after-OOM/runtime-invariant behavior. | +| `protos/reasoner_features/v1/` | Implemented versioned `GetInfo` and server-streaming `Generate` protobuf schema and generated Python bindings. | +| `model/generator/mot/multiview_attention.py` | Pending cache-aware per-view/maskless/Flex support. | +| `data/generator/local_datasets/sft_reasoner_documents.py` | Implemented finite deterministic standard-SFT document enumeration with training framing parity. | +| `data/generator/` | Centralized multi-batch prefetch/cache-handle plumbing remains pending. | +| `checkpoint/reasoner_only.py` | Implemented strict regular/EMA Reasoner-only local-DCP loading with independent reads, safe Generator/visual source-leaf omission, seven-rank generator-only save/resume-load, a nonzero-LR update smoke, and operational resume-and-continue; dataloader-cursor parity and full-model inference composition remain pending. | +| `scripts/extract_reasoner_features.py` | Implemented resumable distributed shared-POSIX cache extraction and atomic publication without NCCL collectives. | +| `scripts/serve_reasoner_features.py` | Implemented one-process/one-GPU remote service launch after load-before-listen readiness. | +| `examples/toml/sft_config/` | Pending opt-in Nano cached-Reasoner recipe after full-scale parity passes. | The training model imports only the provider/client module when `backend="remote"`; it does not import the server implementation or Ray/vLLM inference infrastructure. @@ -711,9 +722,13 @@ Completed: 7. Added all-rank fail-closed gates for provider startup, full feature-signature validation, synchronous request construction/submission, and asynchronous resolution before any rank enters the corresponding FSDP collective. -8. Completed one real optimizer step with one dedicated Reasoner H20 and seven - Generator FSDP H20s, including warm start, forward/backward, optimizer/EMA, - and a generator-only checkpoint save. +8. Completed a seven-rank remote forward/backward and checkpoint smoke. Its initial + optimizer call used effective LR 0, so it did not change Generator weights. +9. Completed a fresh two-iteration run whose second optimizer call used LR `2e-6`; + an exhaustive audit found all 405 live Generator tensors changed. +10. Restored model, optimizer, scheduler, and trainer state from iteration 2 and + continued the same job through iteration 3 at LR `4e-6`; all 405 live tensors + changed again. Dataloader-cursor and uninterrupted-run parity remain pending. Still pending: @@ -735,19 +750,39 @@ allocated about 15.26 GiB after loading and peaked at about 15.40 GiB reserved. validate the transport and runtime, not capacity: the prompt was intentionally short, the network path was loopback, and there were no concurrent Generator ranks. -The end-to-end training smoke reserved physical GPU 0 for the Reasoner and ran a -seven-rank Generator FSDP job on GPUs 1--7 with a 16,384-token packing cap. One -optimizer step completed at loss 0.2512 with no CUDA OOM or RPC error. Generator -CUDA allocator peaks were 20.41--20.46 GiB allocated and 21.47--21.51 GiB reserved; -physical peaks were 23.20--25.11 GiB. The Reasoner peaked at 18.49 GiB. The reported -71.79-second iteration included a 31.37-second, 104-GiB checkpoint save, so it is -not comparable to the earlier steady-state eight-rank benchmark. The saved model -contains only 405 live plus 405 EMA Generator leaves and no Reasoner embedding, -LM-head, or UND-expert leaves. Raw logs and telemetry are under -`outputs/nano_sft_reasoner_benchmark/remote_smoke_7gpu_1step_20260922_v1`. -After the final all-rank signature gate was added, the same seven-rank topology -also completed a resume-only load of that checkpoint in 38.70 seconds without -an error or additional checkpoint write. +The first end-to-end execution smoke reserved physical GPU 0 for the Reasoner and +ran a seven-rank Generator FSDP job on GPUs 1--7 with a 16,384-token packing cap. +Its only scheduled optimizer call used effective LR 0. The pass completed at loss +0.2512 with global gradient norm 0.49541 and no CUDA OOM or RPC error, but an +exhaustive audit found 0/405 live Generator tensors changed. Generator CUDA +allocator peaks were 20.41--20.46 GiB allocated and 21.47--21.51 GiB reserved; +physical peaks were 23.20--25.11 GiB. The Reasoner peaked at 18.49 GiB. The +reported 71.79-second iteration included a 31.37-second, 104-GiB checkpoint save, +so it is not comparable to the earlier steady-state eight-rank benchmark. Its +saved model contains only 405 live plus 405 EMA Generator leaves and no Reasoner +embedding, LM-head, or UND-expert leaves. Raw logs and telemetry are under +`outputs/nano_sft_reasoner_benchmark/remote_smoke_7gpu_1step_20260922_v1`. After +the final all-rank signature gate was added, the same topology completed a +resume-only load in 38.70 seconds without an additional training step. + +A fresh two-iteration run then exercised a real update: its optimizer calls used +LR 0 and `2e-6`. All 405 live Generator tensors changed relative to the shipped +warm-start DCP, covering 96.3922% of their elements with maximum absolute delta +`2.026557922e-6` and no NaN or Inf. Its two iterations reported 41.21 and 52.28 +seconds; the latter included a 32.02-second checkpoint save. Physical peaks were +18.577 GiB on the Reasoner GPU and 30.007--31.718 GiB on the Generator GPUs. This +run used the recipe-default CFG dropout and modality-conditioning distribution, +so its loss and timing are not a paired A/B comparison with the first smoke. + +The same job then restarted the service, loaded its iteration-2 model, optimizer, +scheduler, and trainer state in 10.49 seconds, and completed iteration 3 at the +restored LR of `4e-6`. The saved optimizer steps all advanced from 2 to 3, and an +iter-2-to-iter-3 audit found all 405 live tensors changed, with maximum absolute +delta `4.053115845e-6` and no NaN or Inf. The dataloader subtree was absent and +explicitly skipped, so this proves operational training-state continuation rather +than exact datastream-position or uninterrupted-versus-resumed parity. The fresh +and resumed artifacts are under +`outputs/nano_sft_reasoner_benchmark/remote_smoke_7gpu_2step_20260922_v2`. ### Phase 4: broaden support — pending @@ -764,6 +799,9 @@ Validated in the current implementation: - tiny dense Qwen output and all-generator-gradient parity between the full cached model and `prune_und_pathway_` model on one H20; - absence of UND parameters from the structurally pruned state dict; +- full Nano DCP subset warm start, a real nonzero-LR update of all 405 selected + Generator tensors, and seven-rank model/optimizer/scheduler/trainer continuation + from a generator-only checkpoint; - immutable-cache round trips, batching across shards, content-addressed prompt reuse, cache misses, stale fingerprints, malformed manifests, and corrupt shard checksums; and @@ -773,12 +811,15 @@ Validated in the current implementation: Remaining correctness gates: -- one optimizer step produces equivalent selected weights; +- loss, gradients, and post-update selected weights match an identically configured + `joint`, `inline`, or `offline` run; +- a resumed update matches an uninterrupted run with identical data, persisted + dataloader position, and RNG state; - normal prompt, CFG-null prompt, maximum-length prompt, and per-view captions; -- full checkpoint to generator-only warm start, generator-only resume, and full-model - inference composition; -- representative/max-length remote K/V parity, service restart/OOM behavior, and - capacity scaling; and +- safetensors full-checkpoint to generator-only warm start and full-model inference + composition; +- representative/max-length remote K/V parity, real-service OOM recovery, and capacity + scaling; and - end-to-end stale cache/config provenance is rejected during distributed startup and resume, before any FSDP forward. @@ -788,7 +829,8 @@ Distributed gates: mode; compile-enabled validation remains pending; - no Reasoner parameter appears in either regular or EMA model state on training ranks; - no collective divergence when a cache/service request fails; and -- deterministic sample-to-feature association across resume and dataloader workers. +- deterministic sample-to-feature association across dataloader workers, including a + persisted cursor on resume, remains pending. Performance report: diff --git a/docs/nano_sft_remote_reasoner_test_report.md b/docs/nano_sft_remote_reasoner_test_report.md index f00d04140..d96f4653f 100644 --- a/docs/nano_sft_remote_reasoner_test_report.md +++ b/docs/nano_sft_remote_reasoner_test_report.md @@ -126,19 +126,30 @@ reported allocation. This is a correctness smoke, not a capacity result: it uses a short prompt, loopback networking, one request, and no concurrent Generator ranks. -## Real seven-rank Generator training smoke - -The complete training path was then exercised with GPU 0 reserved for one -Reasoner service replica and physical GPUs 1--7 running a seven-rank Generator -FSDP job. The job used the official eight-video BridgeData sample, a 16,384-token -packing cap, BF16, full activation checkpointing, FP32 EMA, gradient accumulation -one, eager execution, and one optimizer step. It warm-started from the shipped -Nano DCP and saved a new seven-way DCP. +## Real seven-rank execution/resource smoke (zero-LR first iteration) + +The complete remote-conditioned execution path was exercised with GPU 0 +reserved for one Reasoner service replica and physical GPUs 1--7 running a +seven-rank Generator FSDP job. The job used the official eight-video BridgeData +sample, a 16,384-token packing cap, BF16, full activation checkpointing, FP32 +EMA, gradient accumulation one, and eager execution. It warm-started from the +shipped Nano DCP and saved a new seven-way DCP. + +The run completed one forward/backward pass and one scheduled optimizer call. +However, the recipe starts its 50-step warmup at an LR multiplier of zero, so +that optimizer call used an effective learning rate of zero. The run produced a +nonzero gradient norm and initialized Adam moment state, but an exhaustive +comparison against the warm-start DCP found that 0/405 live Generator tensors +changed (`max_abs_delta=0`). It therefore validates the remote execution, +optimizer-state, checkpoint, and resource paths, not a learned parameter update. | Measurement | Result | | ------------------------------------------- | ---------------------------: | -| Optimizer steps | 1 | +| Optimizer calls / nonzero-LR updates | 1 / 0 | +| Effective optimizer LR | 0 | | Final loss | 0.2512 | +| Global gradient norm | 0.49541 | +| Live Generator tensors changed | 0 / 405 (`max_abs_delta=0`) | | Reported iteration time | 71.79 s | | Checkpoint portion of iteration | 31.37 s | | End-to-end `torchrun` wall time | 101.35 s | @@ -151,16 +162,19 @@ Nano DCP and saved a new seven-way DCP. The saved model metadata contains 810 leaves: 405 live Generator leaves and 405 EMA Generator leaves. It contains no `embed_tokens`, `lm_head`, or -`moe_und` parameter, confirming that the training checkpoint does not recreate -the frozen Reasoner. After graceful shutdown, all eight GPUs returned to zero -reported memory use. +`moe_und` parameter, confirming that the checkpoint does not recreate the +frozen Reasoner. This is a structural checkpoint result only; because the +effective learning rate was zero, it does not demonstrate learned Generator +weights or remote-training numerical parity. After graceful shutdown, all +eight GPUs returned to zero reported memory use. After adding the final all-rank feature-signature gate, the same seven-rank job was started again against a fresh service process and resumed the saved DCP at iteration 1. It completed the resume-only path in 38.70 seconds without entering -another training step or writing another checkpoint. This exercises the final -provider handshake, signature collective, seven-way DCP reshard/load, and -provider shutdown success path; all GPUs again returned to zero memory use. +another training step or writing another checkpoint. This validates provider +startup, signature synchronization, seven-way DCP reshard/load, and cleanup. +Because the loaded iteration already equalled `max_iter`, it did not validate +resume-and-continue training. Artifacts: @@ -174,14 +188,106 @@ Artifacts: - checkpoint: `train/cosmos3/sft/nano_remote_smoke_7gpu/checkpoints/iter_000000001`. -These raw artifacts are local-only: the run root is under the gitignored -`outputs/` directory and therefore does not travel with this branch. The -measurements and conclusions above are committed in this report; archive the -run root separately when raw logs or the 104 GiB checkpoint must be handed off. +This is an end-to-end execution-path and resource smoke, not a training- +correctness or throughput comparison. It uses seven Generator ranks, a shorter +packing cap, one zero-LR step, and an iteration time that includes checkpoint +writing. + +## Real nonzero-LR Generator update and continuation + +A fresh two-iteration run used the same one-Reasoner-plus-seven-Generator +topology and saved only `iter_000000002`. Its first optimizer call used LR 0; +after the first scheduler advance, the second call used LR `2e-6`. + +| Measurement | Result | +| ------------------------------------------- | ---------------------------: | +| Optimizer calls / nonzero-LR updates | 2 / 1 | +| Effective optimizer LRs | `0`, `2e-6` | +| Iteration 1 loss / gradient norm / time | 0.2496 / 0.54616 / 41.21 s | +| Iteration 2 loss / gradient norm / time | 0.2131 / 0.35519 / 52.28 s | +| Iteration 2 checkpoint save | 32.02 s | +| End-to-end `torchrun` wall time | 120.62 s | +| Generator physical peak, min/mean/max | 30.007 / 30.922 / 31.718 GiB | +| Reasoner physical peak | 18.577 GiB | +| CUDA OOM / RPC failure / training exception | 0 | + +An exhaustive CPU-streaming DCP comparison against the shipped warm-start +checkpoint found that all 405 live Generator tensors changed. In total, +6,714,184,413 of 6,965,486,784 elements changed (96.3922%). The maximum +absolute delta was `2.026557922e-6`, mean absolute delta was +`2.529132027e-7`, L2 delta was `0.03197588908`, and relative L2 delta was +`1.500317368e-5`; neither checkpoint contained NaN or Inf values. All 405 +optimizer step states were 2. The saved scheduler state had `last_epoch=2` and +LR `4e-6` for the next optimizer update. + +This validates a real nonzero-LR Generator update through remote conditioning. +It does not establish learning quality or numerical parity against `joint`, +`inline`, or `offline`. Unlike the original one-iteration smoke, this run used +the recipe-default CFG-dropout rate of 0.1 and the default T2V/I2V/V2V +conditioning distribution of 0.7/0.2/0.1. The old and new loss/timing values +must therefore not be treated as a paired A/B comparison. + +### Same-job resume and continued update + +The same job was restarted against a fresh Reasoner service and restored the +iteration-2 model, optimizer, scheduler, and trainer state in 10.49 seconds. +Iteration 3 then used the restored LR of `4e-6`, reported loss 0.2138 and global +gradient norm 0.55963, and saved `iter_000000003`. Its reported 73.23-second +iteration included a 30.29-second checkpoint save; end-to-end process wall time +was 102.831 seconds. + +The iter-2-to-iter-3 DCP audit found that 405/405 live Generator tensors changed: +6,763,395,246 of 6,965,486,784 elements (97.0987%), with maximum absolute delta +`4.053115845e-6`, mean absolute delta `4.378727690e-7`, relative L2 delta +`2.647678690e-5`, and no NaN or Inf. All optimizer step states advanced from 2 +to 3; the scheduler advanced from `last_epoch=2` to 3 and saved the next LR as +`6e-6`. The resulting model DCP still has exactly 405 live plus 405 EMA +Generator leaves and no Reasoner leaves. + +Physical peaks during the resumed process were 18.386 GiB on the Reasoner GPU +and 30.019--32.972 GiB on the seven Generator GPUs (31.104 GiB mean). No OOM, +RPC, NCCL, or training error occurred, and all eight GPUs returned to zero +memory after cleanup. + +The checkpoint loader requested dataloader state, but the iteration-2 checkpoint +did not contain a `dataloader` subtree, so every rank explicitly skipped it. +This result validates operational resume-and-continue for model, optimizer, +scheduler, and trainer state. It does not validate exact datastream-position +resume or uninterrupted-versus-resumed numerical parity. + +One launch attempt failed before CUDA initialization because the training script +was passed to `torchrun` by file path. That placed `cosmos_framework/scripts` +first on `sys.path`, causing the sibling `hydra.py` to shadow the installed +Hydra package. The supported module entrypoint succeeded: + +```bash +torchrun --nproc_per_node=7 --module cosmos_framework.scripts.train ... +``` -This is an end-to-end correctness and resource smoke, not a throughput comparison -with the earlier eight-rank runs: it uses seven Generator ranks, a shorter packing -cap, one step, and the reported iteration time includes checkpoint writing. +Artifacts: + +- run root: + `outputs/nano_sft_reasoner_benchmark/remote_smoke_7gpu_2step_20260922_v2`; +- fresh-run logs: `logs/generator_train.log`, `logs/reasoner_service.log`, and + `logs/nvidia_smi.csv`; +- resume logs: `logs/generator_resume_iter2_to_iter3.log`, + `logs/reasoner_service_resume_iter3.log`, and + `logs/nvidia_smi_resume_iter3.csv`; +- failed-entrypoint log: + `logs/generator_resume_iter2_to_iter3.failed_direct_script.log`; +- state audits: `logs/training_state_audit_iter2.json` and + `logs/training_state_audit_iter3.json`; +- derived GPU summary: `logs/gpu_telemetry_summary.json`; +- exhaustive DCP audits: + `logs/generator_weight_delta_iter2_vs_base.json` and + `logs/generator_weight_delta_iter3_vs_iter2.json`; and +- checkpoints: + `train/cosmos3/sft/nano_remote_nonzero_7gpu/checkpoints/iter_000000002` + and `iter_000000003` (104 GiB each). + +These raw artifacts are local-only: the run roots are under the gitignored +`outputs/` directory and therefore do not travel with this branch. Archive them +separately when raw logs or the 104 GiB checkpoints must be handed off. ## Existing offline A/B baseline @@ -203,8 +309,10 @@ must still be measured under realistic service concurrency and networking. queue, compute, D2H, stream preparation, network/client wait, and training-step percentiles. - Repeat remote parity with representative and maximum-length prompts. -- Compare the remote step's loss and selected gradients against an identically - configured `inline` or `offline` run. +- Compare the remote step's loss, selected gradients, and post-update weights + against an identically configured `inline` or `offline` run. +- Persist and restore dataloader position, then compare a resumed update against + an uninterrupted run with identical samples and RNG state. - Derive and verify the Reasoner, tokenizer, and framing fingerprints from the actual checkpoint and tokenizer artifacts. The MVP strictly compares the operator-supplied labels, but does not yet prove that a label is the content From de1b3dc3c5fc295e98ac7cb4d005f058ada29400 Mon Sep 17 00:00:00 2001 From: Codex Date: Tue, 22 Sep 2026 11:03:42 +0800 Subject: [PATCH 4/4] docs(training): add Chinese remote Reasoner test report --- ...nano_sft_remote_reasoner_test_report_zh.md | 304 ++++++++++++++++++ 1 file changed, 304 insertions(+) create mode 100644 docs/nano_sft_remote_reasoner_test_report_zh.md diff --git a/docs/nano_sft_remote_reasoner_test_report_zh.md b/docs/nano_sft_remote_reasoner_test_report_zh.md new file mode 100644 index 000000000..45503635f --- /dev/null +++ b/docs/nano_sft_remote_reasoner_test_report_zh.md @@ -0,0 +1,304 @@ +# Nano SFT Reasoner 解耦与远程服务测试报告(中文) + +日期:2026-09-22 + +分支:`feat/nano-sft-remote-reasoner` + +对应英文报告:`docs/nano_sft_remote_reasoner_test_report.md` + +## 结论摘要 + +本阶段已经完成 Nano SFT 冻结 Reasoner 与可训练 Generator 的结构性解耦, +并实现了两种可实际使用的外部条件输入方式: + +- `offline`:预先抽取每层 Reasoner K/V,训练时从本地 safetensors cache 读取; +- `remote`:由独立的单 GPU gRPC 服务计算 Reasoner K/V,Generator 训练进程只接收结果。 + +`vision_sft_nano` 实际训练的是 405 个 Generator tensor,包括 397 个 +`moe_gen` tensor,以及 `time_embedder`、`vae2llm`、`llm2vae` 共 8 个 +tensor。外部 backend 会在模型实例化、FSDP 包装和 EMA 创建之前裁掉冻结的 +UND/Reasoner 路径,因此 Generator rank 不再承担以下开销: + +- Reasoner 参数实例化与 FSDP shard; +- Reasoner 参数 all-gather; +- Reasoner forward; +- FP32 EMA 中的第二份 Reasoner 权重。 + +目前的核心测试结论如下: + +1. Offline 8xH20 短测中,平均 step time 下降 17.05%,每张 Generator GPU + 的 peak allocated memory 下降 15.63 GiB。 +2. Remote localhost correctness smoke 中,36 层全部 K/V tensor 与直接抽取 + 结果逐 bit 相等。 +3. Remote 真实训练中,第二个 optimizer call 使用非零 LR `2e-6`,405/405 + 个 Generator tensor 全部发生变化。 +4. 同一 job 从 iteration 2 恢复后,以 LR `4e-6` 完成 iteration 3,405/405 + 个 Generator tensor 再次发生变化。 +5. 所有成功运行中均未出现 CUDA OOM、RPC、NCCL 或训练异常,任务结束后 + 8 张 GPU 均恢复到 0 MiB 占用。 + +当前结果证明外部 Reasoner 条件输入可以完成真实 Generator 更新和 +model/optimizer/scheduler/trainer 的续训,但还不代表已经完成生产部署验证或 +与 joint/inline/offline 的严格数值等价验证。 + +## 已实现的能力 + +### Provider-neutral 条件接口 + +实现了 `joint`、`inline`、`offline`、`remote` backend: + +- `joint` 保持现有训练行为,是默认值; +- `inline` 在本进程先运行 Reasoner,再用静态 K/V 执行 Generator,是数值 + 对照路径,但不会节省 Reasoner 显存; +- `offline` 从不可变 cache 读取 K/V; +- `remote` 从独立服务异步获取 K/V; +- `read_through` 仅保留配置契约,目前仍然 fail closed。 + +外部 backend 使用已经完成 RoPE 的 Reasoner K/V。Generator Q、Generator +K/V、attention、MLP 和 residual 路径仍然在线计算并参与反向传播。 + +### Offline backend + +- Reasoner-only DCP loader,不构建 Generator、VAE 或 EMA; +- 与训练一致的 caption/document framing; +- 可恢复、可分布式运行的抽取 CLI; +- layer-major safetensors shard; +- 原子发布、SHA-256 校验、fingerprint 和完整性检查; +- Generator-only warm start、checkpoint 和 EMA 初始化处理。 + +### Remote backend + +- versioned protobuf 和 gRPC server-streaming 协议; +- 启动 handshake、feature signature 和 fingerprint 校验; +- per-chunk 与完整 response checksum; +- 异步请求、绝对 deadline、有限重试和取消; +- request/token admission、bounded queue 和 backpressure; +- OOM 或 runtime invariant 错误后的 unhealthy replica 传播; +- Generator 各 rank 在进入对应 FSDP collective 前同步失败状态; +- success、dry-run 和 failure 路径上的 provider 清理。 + +## 自动化验证 + +最终相关测试集结果:**186 passed**。 + +同时通过: + +- Ruff check 与 format check; +- Pyrefly,0 errors; +- `git diff --check`; +- pre-commit; +- `uv lock --check`。 + +测试覆盖 cache 完整性、tensor codec、协议损坏、checksum/order/shape/dtype +校验、retry、deadline、cancellation、queue admission、OOM 传播、分布式失败 +同步、checkpoint load,以及资源清理。 + +## H20 实测结果 + +### 1. Remote K/V 正确性 + +使用发布的 `Cosmos3-Nano` regular DCP,在两个独立 H20 进程中分别运行直接 +`ReasonerFeatureRuntime` 和 localhost gRPC service。输入为 10-token BF16 +framed prompt。 + +| 指标 | 结果 | +| --------------------------- | ----------------------: | +| 比较的 Reasoner 层数 | 36 | +| K/V 一致性 | 每层 K、V 均逐 bit 相等 | +| 直接抽取 | 0.293149 s | +| localhost remote round trip | 0.326337 s | +| 服务端 Reasoner compute | 0.288081 s | +| 服务端 D2H | 0.002015 s | +| CPU stream preparation | 0.003000 s | +| 服务 load 后 allocated | 15.26 GiB | +| 服务 load peak reserved | 15.40 GiB | + +这是 transport/runtime correctness smoke,不是容量或并发吞吐测试。 + +### 2. 第一次 7-rank smoke 的 zero-LR 修正 + +第一次端到端运行使用 GPU 0 作为 Reasoner service,GPU 1--7 作为 +Generator FSDP ranks。该运行完成了 forward、backward、gradient clipping、 +optimizer/EMA 路径和 checkpoint 写入。 + +但 recipe 的 50-step warmup 从 multiplier 0 开始,因此第一个 optimizer +call 的有效 LR 为 0。完整 DCP 对比结果是: + +- gradient norm:0.49541; +- 0/405 个 live Generator tensor 发生变化; +- `max_abs_delta=0`。 + +因此,这次运行只证明 remote execution、梯度、Adam state、EMA/checkpoint +plumbing 和资源路径正常,不能作为“Generator 权重已经学习更新”的证据。 + +### 3. Fresh 2-step 非零 LR 更新 + +使用同样的 1 Reasoner + 7 Generator 拓扑重新运行两个 iteration: + +| 指标 | 结果 | +| ------------------------------------ | ---------------------------: | +| optimizer calls / 非零 LR updates | 2 / 1 | +| 两次 optimizer LR | `0`, `2e-6` | +| Iteration 1:loss / grad norm / time | 0.2496 / 0.54616 / 41.21 s | +| Iteration 2:loss / grad norm / time | 0.2131 / 0.35519 / 52.28 s | +| Iteration 2 checkpoint save | 32.02 s | +| `torchrun` wall time | 120.62 s | +| Reasoner physical peak | 18.577 GiB | +| Generator physical peak min/mean/max | 30.007 / 30.922 / 31.718 GiB | + +CPU-only 全量 DCP 对比结果: + +- 405/405 个 live Generator tensor 发生变化; +- 6,714,184,413 / 6,965,486,784 个元素变化,即 96.3922%; +- maximum absolute delta:`2.026557922e-6`; +- mean absolute delta:`2.529132027e-7`; +- relative L2 delta:`1.500317368e-5`; +- base 与 candidate 均无 NaN/Inf; +- 405 个 optimizer step state 全部为 2; +- scheduler `last_epoch=2`,下一步 LR 为 `4e-6`。 + +该运行使用 recipe 默认的 CFG dropout 0.1 和 T2V/I2V/V2V 分布 +0.7/0.2/0.1,与第一次固定 T2V、无 dropout 的 smoke 不是严格 paired A/B, +两次运行的 loss 和 time 不应直接用于性能对比。 + +### 4. Iteration 2 到 iteration 3 的续训 + +重启独立 Reasoner service 后,同一训练 job 从 `iter_000000002` 恢复: + +- model、optimizer、scheduler 和 trainer state 恢复成功; +- checkpoint load:10.49 s; +- iteration 3 使用恢复后的 LR `4e-6`; +- loss:0.2138; +- gradient norm:0.55963; +- iteration time:73.23 s,其中 checkpoint save 为 30.29 s; +- process wall time:102.831 s; +- 保存 `iter_000000003`。 + +Iter-2 与 iter-3 的 CPU-only 全量对比结果: + +- 405/405 个 live Generator tensor 再次变化; +- 6,763,395,246 / 6,965,486,784 个元素变化,即 97.0987%; +- maximum absolute delta:`4.053115845e-6`; +- relative L2 delta:`2.647678690e-5`; +- 两个 checkpoint 均无 NaN/Inf; +- 405 个 optimizer step state 从 2 全部推进到 3; +- scheduler `last_epoch=3`,下一步 LR 为 `6e-6`; +- iter-3 model 仍然只有 405 个 live + 405 个 EMA Generator leaves, + Reasoner leaves 为 0。 + +Resume 运行的 physical memory peak: + +| 角色 | 峰值显存 | +| -------------------------- | ---------------------------: | +| Reasoner GPU | 18.386 GiB | +| Generator GPU min/mean/max | 30.019 / 31.104 / 32.972 GiB | + +重要限制:checkpoint loader 请求了 dataloader state,但 iteration-2 checkpoint +中不存在 `dataloader` subtree,因此所有 rank 都明确跳过了 dataloader cursor +恢复。本测试证明 model/optimizer/scheduler/trainer 的 operational +resume-and-continue,不证明 datastream position 恢复,也不证明与 uninterrupted +3-step run 数值等价。 + +### 5. Offline 与原始 Joint 的短测对比 + +该 8xH20 A/B 使用完整 Nano、FSDP、full activation checkpointing、FP32 EMA、 +eager mode 和同一个 8-video BridgeData sample。统计 steady steps 5--10: + +| Backend | Mean step time | Mean peak allocated / Generator GPU | +| ---------------------------- | -------------: | ----------------------------------: | +| Joint Reasoner + Generator | 57.887 s | 52.128 GiB | +| Offline K/V + Generator only | 48.015 s | 36.498 GiB | + +Offline 的变化: + +- mean step time:-17.05%; +- aggregate token throughput:+20.56%; +- mean peak allocated memory:-15.630 GiB/GPU。 + +这是短时 hot-cache benchmark,不代表完整数据集抽取、共享文件系统带宽或长期 +训练吞吐。 + +## 运行中遇到的问题 + +### Reasoner service 不应解析训练专用环境变量 + +最初 service 启动会在通用 SFT config resolution 阶段解析 dataset、VAE 和 +training checkpoint 环境变量。实现已经在 service loader 中裁掉这些无关 +subtree,并增加无相关环境变量的回归测试。 + +### `torchrun` 必须使用 module entrypoint + +一次 resume 尝试将 `cosmos_framework/scripts/train.py` 文件路径直接传给 +`torchrun`,导致同目录的 `hydra.py` 遮蔽安装的 Hydra package,出现: + +```text +No module named 'hydra.core'; 'hydra' is not a package +``` + +正确启动方式是: + +```bash +torchrun --nproc_per_node= --module cosmos_framework.scripts.train ... +``` + +该问题发生在 CUDA 初始化之前,不是模型、checkpoint、OOM 或 RPC 故障。 + +## 当前限制与后续 gate + +合入生产训练前仍需完成: + +1. 使用完全相同的数据、dropout、RNG 和配置,对 remote 与 + joint/inline/offline 的 loss、gradient 和 post-update weights 做严格 A/B。 +2. 保存并恢复 dataloader cursor,对比 uninterrupted 与 resumed update。 +3. 使用 representative 和 maximum-length prompt 做 remote parity 与吞吐测试。 +4. 测试一个 Reasoner replica 服务 1/2/4/8 个 Generator ranks 时的 queue、 + compute、D2H、network/client wait 和 step latency percentile。 +5. 为 7-view/11-view per-view caption 与 specialized multiview attention 增加 + multi-document 支持;protocol v1 当前对此明确 fail closed。 +6. 实现或验证 read-through cache、dynamic batching、layerwise H2D 和持久化 + service cache。 +7. 从实际 checkpoint/tokenizer 内容自动派生 digest,而不是依赖人工标签。 +8. 跨节点生产部署前增加 TLS、authentication、health check 和 metrics。 +9. 完成 full-corpus extraction、共享存储压力和 compile-enabled 长跑测试。 + +## 本地测试产物 + +原始日志和 checkpoint 位于 gitignored `outputs/`,不会随 PR 提交: + +```text +outputs/nano_sft_reasoner_benchmark/remote_smoke_7gpu_1step_20260922_v1 +outputs/nano_sft_reasoner_benchmark/remote_smoke_7gpu_2step_20260922_v2 +``` + +第二次运行中的主要证据: + +```text +logs/generator_train.log +logs/generator_resume_iter2_to_iter3.log +logs/reasoner_service.log +logs/reasoner_service_resume_iter3.log +logs/nvidia_smi.csv +logs/nvidia_smi_resume_iter3.csv +logs/gpu_telemetry_summary.json +logs/training_state_audit_iter2.json +logs/training_state_audit_iter3.json +logs/generator_weight_delta_iter2_vs_base.json +logs/generator_weight_delta_iter3_vs_iter2.json +``` + +两个 Generator-only checkpoint 分别约 104 GiB: + +```text +train/cosmos3/sft/nano_remote_nonzero_7gpu/checkpoints/iter_000000002 +train/cosmos3/sft/nano_remote_nonzero_7gpu/checkpoints/iter_000000003 +``` + +## 最终判断 + +当前实现已经证明:冻结 Reasoner 可以从 Nano Generator SFT rank 中结构性移除, +offline 和 remote backend 均能驱动真实的 Generator 参数更新,并显著降低训练 +rank 显存。Offline 方案已经显示明确的显存与 step-time 收益;Remote 方案已经 +通过 K/V 正确性、非零 LR 更新和 operational resume gate。 + +建议当前阶段定位为可 review、可继续扩展的 MVP。正式生产化仍应以严格数值 +A/B、dataloader resume、代表性并发吞吐、multiview 支持和安全部署能力为准。