diff --git a/.dockerignore b/.dockerignore index 0dfe444b6..2c1ae741d 100644 --- a/.dockerignore +++ b/.dockerignore @@ -1,5 +1,7 @@ .venv .git +**/__pycache__ +**/*.pyc /checkpoints /datasets /output diff --git a/.github/workflows/pre-commit.yml b/.github/workflows/pre-commit.yml index 8d202c8f5..37d6eaab9 100644 --- a/.github/workflows/pre-commit.yml +++ b/.github/workflows/pre-commit.yml @@ -19,4 +19,22 @@ jobs: - uses: actions/setup-python@v6 - uses: astral-sh/setup-uv@v7 - run: uvx pre-commit@4.5.1 run -a -c ci/.pre-commit-config-base.yaml + # The pinned hook requires rumdl 0.1.62, which has no PyPI distribution. + # Build the genuine release without changing the hook or formatter version. + - name: Build pinned rumdl wheel + env: + RUSTUP_HOME: ${{ runner.temp }}/rumdl-rustup + CARGO_HOME: ${{ runner.temp }}/rumdl-cargo + RUSTUP_TOOLCHAIN: '1.94.0' + run: | + git init "$RUNNER_TEMP/rumdl-source" + git -C "$RUNNER_TEMP/rumdl-source" fetch --depth=1 https://github.com/rvben/rumdl.git 8e22c9b16f49c7209355106f731500cc1aacef20 + git -C "$RUNNER_TEMP/rumdl-source" checkout --detach FETCH_HEAD + rustup toolchain install "$RUSTUP_TOOLCHAIN" --profile minimal --no-self-update + uvx maturin@1.15.0 build --release --locked \ + --manifest-path "$RUNNER_TEMP/rumdl-source/Cargo.toml" \ + --out "$RUNNER_TEMP/rumdl-wheels" + uvx --no-index --find-links "$RUNNER_TEMP/rumdl-wheels" rumdl@0.1.62 --version - run: uvx pre-commit@4.5.1 run -a + env: + PIP_FIND_LINKS: ${{ runner.temp }}/rumdl-wheels diff --git a/Dockerfile b/Dockerfile index 74d1e5bcf..e8e74c060 100644 --- a/Dockerfile +++ b/Dockerfile @@ -7,6 +7,25 @@ ARG CUDA_VERSION=13.0.2 ARG BASE_IMAGE=nvidia/cuda:${CUDA_VERSION}-cudnn-devel-ubuntu24.04 FROM ${BASE_IMAGE} +ARG SOURCE_COMMIT +ARG SOURCE_TREE +ARG SOURCE_DIRTY=1 +ARG BUILD_TIMESTAMP +ARG REQUIRE_SOURCE_PROVENANCE=0 +ARG BASE_IMAGE +ARG CUDA_VERSION +LABEL org.opencontainers.image.revision="${SOURCE_COMMIT}" \ + org.opencontainers.image.created="${BUILD_TIMESTAMP}" \ + com.nvidia.cosmos.source-tree="${SOURCE_TREE}" \ + com.nvidia.cosmos.backend="cosmos-framework" +ENV SOURCE_COMMIT="${SOURCE_COMMIT}" \ + SOURCE_TREE="${SOURCE_TREE}" \ + SOURCE_DIRTY="${SOURCE_DIRTY}" \ + BUILD_TIMESTAMP="${BUILD_TIMESTAMP}" \ + REQUIRE_SOURCE_PROVENANCE="${REQUIRE_SOURCE_PROVENANCE}" \ + PROVENANCE_BASE_IMAGE="${BASE_IMAGE}" \ + CUDA_VERSION="${CUDA_VERSION}" + # Set the DEBIAN_FRONTEND environment variable to avoid interactive prompts during apt operations. ENV DEBIAN_FRONTEND=noninteractive @@ -28,7 +47,8 @@ COPY --from=ghcr.io/astral-sh/uv:0.12.2 /uv /uvx /usr/local/bin/ # Copy from the cache instead of linking since it's a mounted volume ENV UV_LINK_MODE=copy # Cache python downloads -ENV UV_PYTHON_CACHE_DIR=/root/.cache/uv/python +ENV UV_PYTHON_CACHE_DIR=/opt/uv-python-cache \ + UV_PYTHON_INSTALL_DIR=/opt/uv-python # Install just: https://just.systems/man/en/pre-built-binaries.html RUN curl --proto '=https' --tlsv1.2 -sSf https://just.systems/install.sh | bash -s -- --to /usr/local/bin --tag 1.46.0 @@ -40,7 +60,8 @@ WORKDIR /workspace # Install python RUN --mount=type=cache,target=/root/.cache/uv \ --mount=type=bind,source=.python-version,target=.python-version \ - uv python install + uv python install && \ + chmod -R a+rX /opt/uv-python /opt/uv-python-cache # Install into virtual environment RUN echo "$CUDA_VERSION" | sed -E 's/^([0-9]+)\.([0-9]+).*/cu\1\2/' > /root/.cuda-name @@ -49,7 +70,7 @@ RUN --mount=type=cache,target=/root/.cache/uv \ --mount=type=bind,source=pyproject.toml,target=pyproject.toml \ --mount=type=bind,source=.python-version,target=.python-version \ --mount=type=bind,source=packages,target=packages \ - uv sync --locked --no-install-project --no-editable --all-extras --group=$(cat /root/.cuda-name) --group=vllm + uv sync --locked --no-install-project --no-editable --all-extras --group=$(cat /root/.cuda-name)-train ENV PATH="/workspace/.venv/bin:$PATH" # Set to 0 to skip the apex build, which is by far the slowest layer. apex is optional: @@ -66,6 +87,18 @@ RUN --mount=type=cache,target=/root/.cache/uv \ echo "INSTALL_APEX=$INSTALL_APEX, skipping apex"; \ fi +# Package the exact source state after the expensive dependency layers so source +# edits do not rebuild apex. Managed platforms do not require a host bind mount. +COPY . /workspace +RUN --mount=type=cache,target=/root/.cache/uv \ + uv pip install --no-deps . + +RUN /workspace/.venv/bin/python /workspace/docker/write_image_provenance.py && \ + chmod a+rx /workspace /workspace/docker /workspace/docker/entrypoint.sh && \ + chmod -R a+rX /opt/cosmos /workspace/.venv /workspace/cosmos_framework /workspace && \ + test -x /workspace/docker/entrypoint.sh && \ + test -x /workspace/.venv/bin/python + # Triton bundled ptxas doesn't support latest GPU architectures ENV TRITON_PTXAS_PATH="/usr/local/cuda/bin/ptxas" diff --git a/cosmos_framework/callbacks/iter_speed.py b/cosmos_framework/callbacks/iter_speed.py index 2ef103c87..aebf44578 100644 --- a/cosmos_framework/callbacks/iter_speed.py +++ b/cosmos_framework/callbacks/iter_speed.py @@ -114,7 +114,7 @@ def on_training_step_end( ) -> None: if self.hit_counter < self.hit_thres: log.info( - f"Iteration {iteration}: " + f"[RANK {log.RANK}] Iteration {iteration}: " f"Hit counter: {self.hit_counter + 1}/{self.hit_thres} | " f"Loss: {loss.detach().item():.4f} | " f"Time: {time.time() - self.last_hit_time:.2f}s", diff --git a/cosmos_framework/callbacks/loss_spike_rollback.py b/cosmos_framework/callbacks/loss_spike_rollback.py new file mode 100644 index 000000000..576c1838c --- /dev/null +++ b/cosmos_framework/callbacks/loss_spike_rollback.py @@ -0,0 +1,332 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 +"""Detect gradient-norm spikes and rewind the model past them. + +A constant learning rate with no decay leaves LoRA SFT marginally stable: most steps +sit at a gradient norm well under the clip threshold, then a single unlucky batch +produces a norm one or two orders of magnitude larger. Clipping bounds the *magnitude* +of that step but not its *direction*, and AdamW folds the anomalous direction into its +first and second moments. Because AdamW's update is scale-invariant, the poisoned +moments keep steering the model for many steps after the offending batch is gone, and +the run either degrades or diverges outright. + +Detection uses the gradient norm rather than the loss because the loss is a lagging +indicator. Measured on Cosmos3-Nano LoRA SFT, at the spike the gradient norm separates +from its running baseline by 10.4x while the loss still reads a healthy 0.2742 and +separates by only 2.8x -- the gradient norm fires roughly 100 steps before the loss +makes the problem visible. + +Recovery rewinds rather than skips. Skipping the update leaves the poisoned moments in +place, and once the loss is elevated every subsequent norm looks like a spike, so a +skip-based guard freezes a diverged run instead of rescuing it (observed: 170 of 180 +steps skipped in a single epoch). This callback keeps a short ring of parameter and +optimizer-moment snapshots and restores the oldest one, which discards the moments as +well as the weights, then backs the learning rate off so the retry is less likely to +trip the same edge. + +MEMORY: the ring holds ``rollback_depth`` copies of every trainable parameter plus its +two AdamW moments. That is roughly 1.2GB for a rank-64 LoRA adapter, but it scales with +the *trainable* parameter count, so full fine-tuning of an 8B model would need hundreds +of gigabytes. Hence ``enabled`` defaults to False; turn it on for PEFT. +""" + +from collections import deque +from typing import Any + +import torch + +from cosmos_framework.utils import log +from cosmos_framework.utils.callback import Callback + +try: # ImaginaireModel is only needed for type hints. + from cosmos_framework.model._base import ImaginaireModel +except Exception: # pragma: no cover - typing convenience only + ImaginaireModel = Any # type: ignore[assignment,misc] + + +def _real_optimizers(optimizer: Any) -> list[torch.optim.Optimizer]: + """Normalise the trainer's optimizer argument to a list of real optimizers. + + The reasoner path hands over an ``OptimizersContainer``, which owns one optimizer per + model part and exposes neither ``param_groups`` nor ``state``. It is iterable, so + iterate; a plain ``torch.optim.Optimizer`` is wrapped in a single-element list. + """ + if hasattr(optimizer, "param_groups"): + return [optimizer] + try: + return list(iter(optimizer)) + except TypeError: + return [] + + +def _real_schedulers(scheduler: Any) -> list[Any]: + """Normalise the trainer's scheduler argument to a list of real LR schedulers. + + ``SchedulersContainer`` holds one ``LambdaLR`` per optimizer under ``.schedulers``. + """ + if hasattr(scheduler, "base_lrs"): + return [scheduler] + inner = getattr(scheduler, "schedulers", None) + if inner: + return list(inner) + return [] + + +class LossSpikeRollback(Callback): + """Rewind past gradient-norm spikes instead of training through them. + + Hooks used, and why each: + + ``on_after_backward`` + Reads the gradient norm BEFORE :class:`~cosmos_framework.callbacks.grad_clip.GradClip` + rescales it. Clipping compresses every spike down to ``clip_norm``, so a norm read + after clipping carries no signal at all. This hook is the last point at which the + true magnitude is still visible, regardless of the order callbacks are registered in. + + ``on_before_zero_grad`` + Runs immediately AFTER the optimizer step, so restoring here undoes the damaging + update itself rather than trying to pre-empt it. Cancelling the step from + ``on_before_optimizer_step`` is not possible -- the trainer calls ``_optimizer_step`` + unconditionally -- and merely zeroing the gradients would not help, since AdamW still + moves the weights from momentum alone. + """ + + def __init__( + self, + enabled: bool = False, + grad_norm_factor: float = 10.0, + window: int = 50, + min_observations: int = 12, + rollback_depth: int = 4, + max_consecutive: int = 8, + lr_backoff: float = 0.5, + lr_recovery: float = 1.02, + lr_min_scale: float = 0.1, + lr_ceiling_decay: float = 0.8, + baseline_inflation_cap: float = 4.0, + backoff_cooldown: int = 50, + recovery_health_factor: float = 2.0, + ) -> None: + super().__init__() + self.enabled = enabled + self.grad_norm_factor = grad_norm_factor + self.window = window + self.min_observations = min_observations + self.rollback_depth = max(1, rollback_depth) + self.max_consecutive = max_consecutive + self.lr_backoff = lr_backoff + self.lr_recovery = lr_recovery + self.lr_min_scale = lr_min_scale + self.lr_ceiling_decay = lr_ceiling_decay + self.baseline_inflation_cap = baseline_inflation_cap + self.backoff_cooldown = backoff_cooldown + self.recovery_health_factor = recovery_health_factor + + self._norms: deque[float] = deque(maxlen=window) + self._ring: deque[dict[str, Any]] = deque(maxlen=self.rollback_depth) + self._pending_reason: str | None = None + self._consecutive = 0 + self._lr_scale = 1.0 + # Ceiling that recovery walks back toward. Ratchets down on every rollback so a run + # that keeps tripping settles at a lower rate instead of climbing back to a rate it + # has already demonstrated it cannot hold. + self._lr_ceiling = 1.0 + self._base_lrs: list[list[float]] | None = None + # Lowest median seen since arming: the run's demonstrated healthy scale. + self._healthy_baseline: float | None = None + # Iteration of the last learning-rate backoff, so a burst of spikes counts as one + # episode rather than as one escalation per rollback. + self._last_backoff: int | None = None + self.rollbacks = 0 + + # ---------------------------------------------------------------- detection + + def _baseline(self) -> float | None: + """Median of the recent window, or None until enough steps have been seen. + + The median rather than the mean because a spike that is already inside the window + would drag a mean upward and mask the next one. + """ + if len(self._norms) < self.min_observations: + return None + ordered = sorted(self._norms) + mid = len(ordered) // 2 + median = ordered[mid] if len(ordered) % 2 else 0.5 * (ordered[mid - 1] + ordered[mid]) + + # Anchor the baseline to the healthiest scale this run has demonstrated. Without + # this the baseline chases a deteriorating run upward: once the window fills with + # elevated norms, the relative test silently demands an ever larger spike to trip. + # Observed in practice -- a baseline that drifted from ~0.4 to 9.03 needed a norm + # above 90 to fire, so the guard went quiet exactly when it was needed most. + if self._healthy_baseline is None or median < self._healthy_baseline: + self._healthy_baseline = median + return min(median, self.baseline_inflation_cap * self._healthy_baseline) + + @torch.no_grad() + def on_after_backward(self, model: "ImaginaireModel", iteration: int = 0) -> None: + if not self.enabled: + return + total_sq = 0.0 + for parameter in model.parameters(): + if parameter.requires_grad and parameter.grad is not None: + grad = parameter.grad + # DTensor under FSDP: reduce to the local shard's contribution. The + # baseline is a ratio against this same quantity, so a consistently + # partial norm still detects a spike. + grad = grad.to_local() if hasattr(grad, "to_local") else grad + total_sq += float(grad.detach().float().pow(2).sum()) + grad_norm = total_sq**0.5 + + if not (grad_norm == grad_norm) or grad_norm in (float("inf"), float("-inf")): + self._pending_reason = f"non-finite gradient norm ({grad_norm})" + return + + baseline = self._baseline() + if baseline is not None and baseline > 0.0 and grad_norm > self.grad_norm_factor * baseline: + self._pending_reason = ( + f"gradient norm {grad_norm:.4f} exceeds {self.grad_norm_factor:.1f}x " + f"the median of the last {len(self._norms)} steps ({baseline:.4f})" + ) + return + + self._pending_reason = None + self._norms.append(grad_norm) + + def _is_healthy(self) -> bool: + """Is the current gradient-norm scale back near the run's demonstrated healthy one?""" + if self._healthy_baseline is None or len(self._norms) < self.min_observations: + return True # not enough evidence to withhold recovery + ordered = sorted(self._norms) + mid = len(ordered) // 2 + median = ordered[mid] if len(ordered) % 2 else 0.5 * (ordered[mid - 1] + ordered[mid]) + return median <= self.recovery_health_factor * self._healthy_baseline + + # ------------------------------------------------------------ snapshot/undo + + @torch.no_grad() + def _capture(self, model: "ImaginaireModel", optimizer: torch.optim.Optimizer) -> None: + params = {name: p.detach().clone() for name, p in model.named_parameters() if p.requires_grad} + moments: dict[tuple[int, int], dict[str, torch.Tensor]] = {} + for opt_index, opt in enumerate(_real_optimizers(optimizer)): + for index, parameter in enumerate(p for group in opt.param_groups for p in group["params"]): + state = opt.state.get(parameter) + if not state: + continue + moments[(opt_index, index)] = { + key: value.detach().clone() for key, value in state.items() if isinstance(value, torch.Tensor) + } + self._ring.append({"params": params, "moments": moments}) + + @torch.no_grad() + def _restore(self, model: "ImaginaireModel", optimizer: torch.optim.Optimizer) -> bool: + if not self._ring: + return False + snapshot = self._ring[0] + named = dict(model.named_parameters()) + for name, saved in snapshot["params"].items(): + target = named.get(name) + if target is not None: + target.detach().copy_(saved) + opts = _real_optimizers(optimizer) + flats = [[p for group in opt.param_groups for p in group["params"]] for opt in opts] + for (opt_index, index), saved_state in snapshot["moments"].items(): + if opt_index >= len(opts) or index >= len(flats[opt_index]): + continue + state = opts[opt_index].state.get(flats[opt_index][index]) + if not state: + continue + for key, saved in saved_state.items(): + if key in state and isinstance(state[key], torch.Tensor): + state[key].copy_(saved) + # Everything newer than the restored point came after the poisoned step, so it is + # not a safe place to rewind to a second time. The gradient-norm history is NOT + # cleared alongside it: those samples were all taken before the spike (the spiking + # value is never appended), they describe the state we just rewound to, and + # discarding them forces the baseline to be rebuilt from post-rollback steps that + # may themselves be unhealthy. + self._ring.clear() + return True + + def _scale_lr(self, optimizer: torch.optim.Optimizer, scheduler: Any, factor: float) -> None: + """Back the learning rate off by rescaling the scheduler's ``base_lrs``. + + Writing ``param_group["lr"]`` directly would not survive: ``LambdaLR`` recomputes + ``lr = base_lrs[i] * lr_lambda(step)`` on every ``scheduler.step()``, so a direct + write is discarded on the very next iteration. ``base_lrs`` is the input to that + product and therefore the only durable place to apply the backoff. + """ + schedulers = _real_schedulers(scheduler) + if self._base_lrs is None: + if schedulers: + self._base_lrs = [list(sched.base_lrs) for sched in schedulers] + else: + self._base_lrs = [[group["lr"] for group in opt.param_groups] for opt in _real_optimizers(optimizer)] + if factor < 1.0: + # Back off RELATIVE TO THE CEILING, not to the current scale. Compounding from + # the current scale makes a cluster of spikes multiply out: four rollbacks + # inside thirteen steps drove 0.5**4 straight into the floor, and the run then + # spent its remaining 274 steps at a tenth of the intended rate -- rescued from + # divergence but undertrained, which cost roughly ten points of accuracy. + # The ceiling carries the persistent penalty; this is the temporary dip. + self._lr_ceiling = max(self.lr_min_scale, self._lr_ceiling * self.lr_ceiling_decay) + self._lr_scale = max(self.lr_min_scale, self._lr_ceiling * self.lr_backoff) + else: + self._lr_scale = max(self.lr_min_scale, min(self._lr_ceiling, self._lr_scale * factor)) + + if schedulers: + for sched, bases in zip(schedulers, self._base_lrs): + sched.base_lrs = [base * self._lr_scale for base in bases] + return + for opt, bases in zip(_real_optimizers(optimizer), self._base_lrs): + for group, base in zip(opt.param_groups, bases): + group["lr"] = base * self._lr_scale + + def on_before_zero_grad( + self, + model: "ImaginaireModel", + optimizer: torch.optim.Optimizer, + scheduler: torch.optim.lr_scheduler.LRScheduler, + iteration: int = 0, + ) -> None: + if not self.enabled: + return + + if self._pending_reason is None: + self._consecutive = 0 + # Recover only once the run is demonstrably healthy again, not merely because + # this particular step did not spike. A degraded run looks clean between its + # spikes, so a recovery gated on "no spike this step" walks the rate back up + # while the model is still sick -- observed climbing to 0.627 of base while the + # median loss sat at 0.69, five times its healthy value, for 280 steps. + if self._lr_scale < self._lr_ceiling and self._is_healthy(): + self._scale_lr(optimizer, scheduler, self.lr_recovery) + self._capture(model, optimizer) + return + + reason = self._pending_reason + self._pending_reason = None + self._consecutive += 1 + if self._consecutive > self.max_consecutive: + log.warning( + f"loss-spike guard: {self._consecutive} consecutive spikes at iteration {iteration}; " + "standing down so a genuinely diverged run is not frozen mid-epoch" + ) + return + + if self._restore(model, optimizer): + self.rollbacks += 1 + # Rewinding repeatedly is cheap and safe; cutting the rate repeatedly is not. + # Only the first rollback of an episode moves the rate. + fresh_episode = self._last_backoff is None or (iteration - self._last_backoff) >= self.backoff_cooldown + if fresh_episode: + self._last_backoff = iteration + self._scale_lr(optimizer, scheduler, self.lr_backoff) + log.warning( + f"loss-spike guard: rolled back at iteration {iteration} ({reason}); " + f"learning rate scaled to {self._lr_scale:.3f} of base [rollback #{self.rollbacks}]" + ) + else: + log.warning( + f"loss-spike guard: spike at iteration {iteration} ({reason}) but no snapshot is " + "available yet; letting the step stand" + ) diff --git a/cosmos_framework/callbacks/loss_spike_rollback_test.py b/cosmos_framework/callbacks/loss_spike_rollback_test.py new file mode 100644 index 000000000..2a94569cf --- /dev/null +++ b/cosmos_framework/callbacks/loss_spike_rollback_test.py @@ -0,0 +1,248 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 +"""Tests for the gradient-norm spike guard. + +Natural spikes are not reproducible run to run, so the guard is exercised with an +injected gradient rather than by waiting for training to misbehave. +""" + +import torch +from torch import nn +from torch.optim.lr_scheduler import LambdaLR + +from cosmos_framework.callbacks.loss_spike_rollback import LossSpikeRollback + + +def _harness(**kwargs): + torch.manual_seed(0) + model = nn.Linear(8, 8) + optimizer = torch.optim.AdamW(model.parameters(), lr=0.01) + scheduler = LambdaLR(optimizer, lr_lambda=lambda step: 1.0) + defaults = dict(enabled=True, window=10, min_observations=5, rollback_depth=2, grad_norm_factor=10.0) + defaults.update(kwargs) + return model, optimizer, scheduler, LossSpikeRollback(**defaults) + + +def _step(model, optimizer, scheduler, callback, grad_scale): + for parameter in model.parameters(): + parameter.grad = torch.full_like(parameter, grad_scale) + callback.on_after_backward(model) + optimizer.step() + callback.on_before_zero_grad(model, optimizer, scheduler) + scheduler.step() + optimizer.zero_grad(set_to_none=True) + + +def _step_at(model, optimizer, scheduler, callback, grad_scale, iteration): + for parameter in model.parameters(): + parameter.grad = torch.full_like(parameter, grad_scale) + callback.on_after_backward(model, iteration=iteration) + optimizer.step() + callback.on_before_zero_grad(model, optimizer, scheduler, iteration=iteration) + scheduler.step() + optimizer.zero_grad(set_to_none=True) + + +def _spike_at(model, optimizer, scheduler, callback, iteration): + _step_at(model, optimizer, scheduler, callback, 10.0, iteration) + + +def test_no_rollback_on_steady_gradients(): + model, optimizer, scheduler, callback = _harness() + for _ in range(20): + _step(model, optimizer, scheduler, callback, 0.01) + assert callback.rollbacks == 0 + assert scheduler.base_lrs == [0.01] + + +def test_rollback_restores_weights_and_backs_off_lr(): + model, optimizer, scheduler, callback = _harness() + for _ in range(20): + _step(model, optimizer, scheduler, callback, 0.01) + + # The ring holds the last `rollback_depth` snapshots; the guard restores the OLDEST. + expected = {name: tensor.clone() for name, tensor in callback._ring[0]["params"].items()} + + _step(model, optimizer, scheduler, callback, 10.0) # ~1000x the steady norm + + assert callback.rollbacks == 1, "guard did not fire on an injected spike" + for name, parameter in model.named_parameters(): + torch.testing.assert_close(parameter.detach(), expected[name]) + # Backoff is ceiling-relative: ceiling ratchets 1.0 -> 0.8, then dips by 0.5. + assert scheduler.base_lrs == [0.01 * 0.8 * 0.5], f"unexpected rate {scheduler.base_lrs}" + + +def test_lr_backoff_survives_scheduler_step(): + """The whole point of rescaling base_lrs rather than param_group['lr']. + + LambdaLR recomputes lr = base_lrs * lambda(step) on every step, so a direct write to + param_group['lr'] would be discarded on the next iteration and the backoff would be + silently ineffective. + """ + model, optimizer, scheduler, callback = _harness() + for _ in range(20): + _step(model, optimizer, scheduler, callback, 0.01) + _step(model, optimizer, scheduler, callback, 10.0) + lr_right_after = optimizer.param_groups[0]["lr"] + for _ in range(3): + _step(model, optimizer, scheduler, callback, 0.01) + assert optimizer.param_groups[0]["lr"] <= lr_right_after * 1.1 + assert optimizer.param_groups[0]["lr"] < 0.01, "backoff was undone by scheduler.step()" + + +def test_stands_down_after_sustained_divergence(): + """A run that has truly diverged must not be frozen by an unyielding guard.""" + model, optimizer, scheduler, callback = _harness(max_consecutive=3) + for _ in range(20): + _step(model, optimizer, scheduler, callback, 0.01) + for _ in range(10): + _step(model, optimizer, scheduler, callback, 10.0) + assert callback.rollbacks <= 3, f"guard never stood down ({callback.rollbacks} rollbacks)" + + +def test_disabled_is_inert(): + model, optimizer, scheduler, callback = _harness(enabled=False) + for _ in range(20): + _step(model, optimizer, scheduler, callback, 0.01) + _step(model, optimizer, scheduler, callback, 10.0) + assert callback.rollbacks == 0 + assert not callback._ring, "disabled guard should not pay the snapshot memory cost" + + +def test_baseline_does_not_chase_a_deteriorating_run(): + """The failure that let a real run escape the guard. + + Once the window fills with elevated norms the median rises, and a purely relative test + then demands an ever larger spike to trip. Observed in training: the baseline drifted + from ~0.4 to 9.03, so firing required a norm above 90 and the guard went quiet exactly + when it was needed. The baseline is anchored to the healthiest scale the run has shown. + """ + model, optimizer, scheduler, callback = _harness(baseline_inflation_cap=4.0) + for _ in range(20): + _step(model, optimizer, scheduler, callback, 0.01) + healthy = callback._baseline() + assert healthy is not None + + # Feed a sustained elevation an order of magnitude above healthy, short of tripping. + for _ in range(60): + for parameter in model.parameters(): + parameter.grad = torch.full_like(parameter, 0.05) + callback.on_after_backward(model) + callback.on_before_zero_grad(model, optimizer, scheduler) + + inflated = callback._baseline() + assert inflated <= 4.0 * healthy * 1.001, f"baseline inflated to {inflated} from {healthy}" + + +def test_repeated_rollbacks_ratchet_the_learning_rate_down(): + """A run that keeps tripping must not climb back to a rate it cannot hold.""" + model, optimizer, scheduler, callback = _harness(max_consecutive=100) + for _ in range(20): + _step(model, optimizer, scheduler, callback, 0.01) + + ceilings = [] + for _ in range(3): + _step(model, optimizer, scheduler, callback, 10.0) # spike + for _ in range(40): # long clean stretch: recovery walks up to the ceiling + _step(model, optimizer, scheduler, callback, 0.01) + ceilings.append(callback._lr_ceiling) + + assert ceilings == sorted(ceilings, reverse=True), f"ceiling did not ratchet down: {ceilings}" + assert ceilings[-1] < 1.0 + assert callback._lr_scale <= callback._lr_ceiling + 1e-9 + + +def test_rollback_keeps_the_healthy_norm_window(): + """Clearing the window on rollback discards the only healthy reference available.""" + model, optimizer, scheduler, callback = _harness() + for _ in range(20): + _step(model, optimizer, scheduler, callback, 0.01) + before = len(callback._norms) + _step(model, optimizer, scheduler, callback, 10.0) + assert callback.rollbacks == 1 + assert len(callback._norms) == before, "healthy gradient-norm history was discarded" + + +def test_clustered_rollbacks_do_not_compound_into_the_floor(): + """The failure that rescued a run from divergence but left it undertrained. + + Four rollbacks inside thirteen steps compounded 0.5**4 straight into the learning-rate + floor; the run then spent its remaining 274 steps at a tenth of the intended rate and + lost about ten points of accuracy. A burst of spikes is one episode, not four + escalations, and the dip is measured from the ceiling rather than from wherever the + scale happens to have landed. + """ + model, optimizer, scheduler, callback = _harness(max_consecutive=100, backoff_cooldown=50) + for _ in range(20): + _step(model, optimizer, scheduler, callback, 0.01) + + iteration = 0 + for burst in range(4): # four spikes a few steps apart: one episode + _spike_at(model, optimizer, scheduler, callback, iteration) + iteration += 1 + for _ in range(3): # clean steps repopulate the ring a rollback cleared + _step_at(model, optimizer, scheduler, callback, 0.01, iteration) + iteration += 1 + + assert callback.rollbacks == 4, "every spike should still be rewound" + assert callback._lr_scale > callback.lr_min_scale, ( + f"clustered rollbacks collapsed the rate to the floor ({callback._lr_scale})" + ) + assert callback._lr_ceiling > 0.5, f"ceiling over-ratcheted within one episode ({callback._lr_ceiling})" + + +def test_separated_episodes_still_ratchet(): + """Spikes far apart are genuinely separate episodes and should each cost rate.""" + model, optimizer, scheduler, callback = _harness(max_consecutive=100, backoff_cooldown=10) + for _ in range(20): + _step(model, optimizer, scheduler, callback, 0.01) + + ceilings = [] + for episode in range(3): + it = episode * 100 + _spike_at(model, optimizer, scheduler, callback, it) + ceilings.append(callback._lr_ceiling) + for offset in range(1, 6): # clean steps rebuild the ring for the next episode + _step_at(model, optimizer, scheduler, callback, 0.01, it + offset) + + assert ceilings == sorted(ceilings, reverse=True) and ceilings[-1] < ceilings[0] + + +def test_recovery_waits_for_the_run_to_be_healthy_again(): + """A degraded run looks clean between spikes, and must not be handed its rate back. + + Observed in training: the rate climbed back to 0.627 of base while the median loss sat + at 0.69, roughly five times its healthy value, and stayed there for 280 steps. + """ + model, optimizer, scheduler, callback = _harness(recovery_health_factor=2.0) + for _ in range(20): + _step(model, optimizer, scheduler, callback, 0.01) + _step(model, optimizer, scheduler, callback, 10.0) # spike -> backoff + suppressed = callback._lr_scale + assert suppressed < callback._lr_ceiling + + # Sustained elevation well above healthy but below the spike threshold: no step trips + # the guard, yet the run is plainly not well. + for _ in range(80): + _step(model, optimizer, scheduler, callback, 0.08) + assert callback.rollbacks == 1, "elevation should not itself trip the guard" + # A few recovery steps land before the window accumulates enough elevated samples for + # the median to register the degradation; what matters is that recovery then stops + # well short of the ceiling, where an ungated 1.02-per-step walk would have arrived. + assert callback._lr_scale < suppressed * 1.15, ( + f"rate recovered to {callback._lr_scale} while the run was still degraded" + ) + assert callback._lr_scale < 0.7 * callback._lr_ceiling, ( + f"rate reached {callback._lr_scale}, close to the ceiling {callback._lr_ceiling}" + ) + + +def test_recovery_resumes_once_gradients_return_to_normal(): + model, optimizer, scheduler, callback = _harness(recovery_health_factor=2.0) + for _ in range(20): + _step(model, optimizer, scheduler, callback, 0.01) + _step(model, optimizer, scheduler, callback, 10.0) + suppressed = callback._lr_scale + for _ in range(80): # genuinely healthy steps + _step(model, optimizer, scheduler, callback, 0.01) + assert callback._lr_scale > suppressed, "rate never recovered despite a healthy run" diff --git a/cosmos_framework/callbacks/workflow_status.py b/cosmos_framework/callbacks/workflow_status.py new file mode 100644 index 000000000..1d1679f9b --- /dev/null +++ b/cosmos_framework/callbacks/workflow_status.py @@ -0,0 +1,530 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +"""Cosmos-compatible lifecycle and metric logging for Cosmos Framework training.""" + +from __future__ import annotations + +import json +import os +import time +from datetime import datetime, timedelta +from pathlib import Path +from typing import Any + +import torch +import torch.distributed as dist + +from cosmos_framework.utils import distributed, log +from cosmos_framework.utils.callback import Callback + + +def _to_json_value(value: Any) -> Any: + """Recursively convert tensors and array-like values to JSON-safe values.""" + if isinstance(value, dict): + return {str(key): _to_json_value(item) for key, item in value.items()} + if isinstance(value, (list, tuple)): + return [_to_json_value(item) for item in value] + if isinstance(value, torch.Tensor): + value = value.detach() + if value.numel() == 1: + return value.item() + return value.cpu().tolist() + if hasattr(value, "item"): + try: + return value.item() + except (TypeError, ValueError): + pass + if hasattr(value, "tolist"): + try: + return value.tolist() + except (TypeError, ValueError): + pass + return value + + +class _WorkflowStatusWriter: + """Write lifecycle records directly; no external logging package is required.""" + + def __init__(self, filename: str) -> None: + self.filename = filename + Path(filename).parent.mkdir(parents=True, exist_ok=True) + + def write( + self, + *, + status: str, + message: str, + data: dict[str, Any] | None = None, + kpi: dict[str, Any] | None = None, + verbosity: str = "INFO", + ) -> None: + data = _to_json_value(data or {}) + kpi = _to_json_value(kpi or {}) + + now = datetime.now() + payload: dict[str, Any] = { + **data, + "date": f"{now.month}/{now.day}/{now.year}", + "time": f"{now.hour}:{now.minute}:{now.second}", + "status": status, + "verbosity": verbosity, + "message": message, + } + if kpi: + payload["kpi"] = kpi + with open(self.filename, "a", encoding="utf-8") as status_file: + status_file.write(json.dumps(payload, default=str) + "\n") + + +def write_early_failure(error: BaseException) -> bool: + """Write a terminal Cosmos record before the callback/config exists. + + The orchestration layer supplies ``COSMOS_STATUS_FILE`` for direct launches, + or the normal Cosmos job/result variables. No implicit host path is used. + """ + status_path = os.environ.get("COSMOS_STATUS_FILE") + if not status_path: + job_id = os.environ.get("COSMOS_JOB_ID") + results_root = os.environ.get("COSMOS_RESULTS_ROOT") + if job_id and results_root: + status_path = os.path.join(results_root, job_id, "status.json") + if not status_path: + api_job_id = os.environ.get("COSMOS_API_JOB_ID") + results_root = os.environ.get("COSMOS_API_RESULTS_DIR") + if api_job_id and results_root: + status_path = os.path.join(results_root, api_job_id, "status.json") + if not status_path: + return False + _WorkflowStatusWriter(status_path).write( + status="FAILURE", + verbosity="ERROR", + message=f"Cosmos Framework training failed before callback initialization: {error}", + data={"phase": "preflight_or_initialization", "error_type": type(error).__name__}, + ) + return True + + +class WorkflowStatusCallback(Callback): + """Write Cosmos lifecycle, training, and validation records from rank zero. + + The output path is resolved in this order: + + 1. ``status_file_path`` when explicitly configured. + 2. ``$COSMOS_RESULTS_ROOT/$COSMOS_JOB_ID/status.json`` (Cosmos SDK). + 3. ``$COSMOS_API_RESULTS_DIR/$COSMOS_API_JOB_ID/status.json`` (Cosmos API). + 4. ``/status.json`` for direct launches. + """ + + def __init__( + self, + enabled: bool = False, + status_file_path: str | None = None, + experiment_name: str = "", + logging_interval: int = 1, + validation_heartbeat_interval: int = 1, + ) -> None: + if logging_interval < 1: + raise ValueError("logging_interval must be >= 1") + if validation_heartbeat_interval < 1: + raise ValueError("validation_heartbeat_interval must be >= 1") + self.enabled = enabled + self.status_file_path = status_file_path + self.experiment_name = experiment_name + self.logging_interval = logging_interval + self.validation_heartbeat_interval = validation_heartbeat_interval + self._writer: _WorkflowStatusWriter | None = None + self._train_start_time = 0.0 + self._step_start_time = 0.0 + self._validation_batches = 0 + self._validation_loss_numerator = 0.0 + self._validation_loss_denominator = 0 + self._validation_local_numerators: list[torch.Tensor] = [] + self._validation_local_denominators: list[torch.Tensor] = [] + self._train_loss_numerator = 0.0 + self._train_loss_denominator = 0 + self.last_training_loss: float | None = None + self.last_validation_loss: float | None = None + + def _is_rank_zero(self) -> bool: + if dist.is_available() and dist.is_initialized(): + return distributed.is_rank0() + return int(os.environ.get("RANK", os.environ.get("LOCAL_RANK", "0"))) == 0 + + def _component_name(self) -> str: + if self.experiment_name: + return self.experiment_name + return getattr(self.config.job, "name", "Cosmos Framework SFT") + + def _resolve_status_file(self) -> str: + if self.status_file_path: + return self.status_file_path + + job_id = os.environ.get("COSMOS_JOB_ID") + if job_id: + results_root = os.environ.get("COSMOS_RESULTS_ROOT") + if not results_root: + raise RuntimeError("COSMOS_RESULTS_ROOT is required when COSMOS_JOB_ID is set") + return os.path.join(results_root, job_id, "status.json") + + api_job_id = os.environ.get("COSMOS_API_JOB_ID") + if api_job_id: + results_root = os.environ.get("COSMOS_API_RESULTS_DIR") + if not results_root: + raise RuntimeError("COSMOS_API_RESULTS_DIR is required when COSMOS_API_JOB_ID is set") + return os.path.join(results_root, api_job_id, "status.json") + + return os.path.join(self.config.job.path_local, "status.json") + + def _get_writer(self) -> _WorkflowStatusWriter | None: + if not self.enabled or not self._is_rank_zero(): + return None + if self._writer is None: + self._writer = _WorkflowStatusWriter(self._resolve_status_file()) + return self._writer + + def _progress_data(self, iteration: int, seconds_per_step: float | None = None) -> dict[str, Any]: + trainer = getattr(self, "trainer", None) + max_step = int(getattr(trainer, "max_iterations", self.config.trainer.max_iter)) + if seconds_per_step is None: + elapsed = max(time.monotonic() - self._train_start_time, 0.0) + seconds_per_step = elapsed / max(iteration, 1) + eta_seconds = max(max_step - iteration, 0) * seconds_per_step + data = { + "component": self._component_name(), + "step": iteration, + "max_step": max_step, + "time_per_step": str(timedelta(seconds=seconds_per_step)), + "eta": str(timedelta(seconds=eta_seconds)), + } + steps_per_epoch = getattr(trainer, "steps_per_epoch", None) + num_epochs = getattr(trainer, "num_epochs", None) + if steps_per_epoch and num_epochs: + completed_epochs = iteration // steps_per_epoch + if iteration == 0: + epoch = 1 + step_in_epoch = 0 + elif iteration % steps_per_epoch == 0: + epoch = min(completed_epochs, num_epochs) + step_in_epoch = steps_per_epoch + else: + epoch = min(completed_epochs + 1, num_epochs) + step_in_epoch = iteration % steps_per_epoch + data.update( + { + "epoch": epoch, + "max_epoch": num_epochs, + "completed_epochs": min(completed_epochs, num_epochs), + "step_in_epoch": step_in_epoch, + "steps_per_epoch": steps_per_epoch, + "time_per_epoch": str(timedelta(seconds=seconds_per_step * steps_per_epoch)), + } + ) + return data + + def _epoch_label(self, iteration: int) -> str: + progress = self._progress_data(iteration, seconds_per_step=0.0) + if "epoch" not in progress: + return f"training step {iteration}/{progress['max_step']}" + return f"epoch {progress['epoch']}/{progress['max_epoch']}" + + @staticmethod + def _local_token_stats( + loss: torch.Tensor, + data_batch: dict[str, Any], + output_batch: dict[str, Any], + ) -> tuple[torch.Tensor, torch.Tensor]: + numerator = output_batch.get("loss_numerator") + denominator = output_batch.get("loss_denominator") + if numerator is not None and denominator is not None: + return numerator.detach(), denominator.detach().to(dtype=torch.long) + + # Backward-compatible fallback for non-VLM models. It is deliberately + # sample-weighted and is never used by the Cosmos3 VLM path, which + # always emits exact token statistics. + sample_count = next( + ( + int(value.shape[0]) + for value in data_batch.values() + if isinstance(value, torch.Tensor) and value.ndim > 0 + ), + 1, + ) + return ( + loss.detach() * sample_count, + torch.tensor(sample_count, device=loss.device, dtype=torch.long), + ) + + @classmethod + def _global_token_average( + cls, + loss: torch.Tensor, + data_batch: dict[str, Any], + output_batch: dict[str, Any], + ) -> tuple[float, float, int]: + numerator, denominator = cls._local_token_stats(loss, data_batch, output_batch) + numerator = numerator.clone() + denominator = denominator.clone() + if dist.is_available() and dist.is_initialized(): + dist.all_reduce(numerator, op=dist.ReduceOp.SUM) + dist.all_reduce(denominator, op=dist.ReduceOp.SUM) + global_denominator = int(denominator.item()) + global_numerator = float(numerator.item()) + return global_numerator / max(global_denominator, 1), global_numerator, global_denominator + + @staticmethod + def _reduce_accumulator(numerator: float, denominator: int) -> tuple[float, int]: + device = "cuda" if dist.is_available() and dist.is_initialized() else "cpu" + values = torch.tensor([numerator, float(denominator)], dtype=torch.float64, device=device) + if dist.is_available() and dist.is_initialized(): + dist.all_reduce(values, op=dist.ReduceOp.SUM) + return float(values[0].item()), int(values[1].item()) + + def on_train_start(self, model: Any, iteration: int = 0) -> None: + self._train_start_time = time.monotonic() + self._train_loss_numerator = 0.0 + self._train_loss_denominator = 0 + writer = self._get_writer() + if writer is not None: + writer.write( + status="STARTED", + message=f"Starting {self._component_name()} training", + data={ + **self._progress_data(iteration, seconds_per_step=0.0), + "parameter_summary": getattr(model, "parameter_summary", None), + }, + ) + log.info(f"Cosmos status will be logged to {writer.filename}") + + def on_training_step_start(self, model: Any, data: dict[str, Any], iteration: int = 0) -> None: + self._step_start_time = time.monotonic() + + def on_training_step_batch_end( + self, + model: Any, + data_batch: dict[str, Any], + output_batch: dict[str, Any], + loss: torch.Tensor, + iteration: int = 0, + ) -> None: + numerator, denominator = self._local_token_stats(loss, data_batch, output_batch) + self._train_loss_numerator += float(numerator.item()) + self._train_loss_denominator += int(denominator.item()) + + def on_training_step_end( + self, + model: Any, + data_batch: dict[str, Any], + output_batch: dict[str, Any], + loss: torch.Tensor, + iteration: int = 0, + ) -> None: + interval = int(self.config.trainer.logging_iter) * self.logging_interval + if iteration % interval != 0: + return + + average_loss, numerator, denominator = self._global_token_average(loss, data_batch, output_batch) + writer = self._get_writer() + if writer is None: + return + seconds_per_step = max(time.monotonic() - self._step_start_time, 0.0) + kpi = { + "train/step_loss": average_loss, + "train/step_loss_numerator": numerator, + "train/step_loss_denominator": denominator, + } + progress = self._progress_data(iteration, seconds_per_step=seconds_per_step) + if "epoch" in progress: + message = ( + f"Training epoch {progress['epoch']}/{progress['max_epoch']}, " + f"step {progress['step_in_epoch']}/{progress['steps_per_epoch']} " + f"(global step {iteration}/{progress['max_step']}) - Loss: {average_loss:.6f}" + ) + else: + message = f"Training step {iteration}/{progress['max_step']} - Loss: {average_loss:.6f}" + writer.write( + status="RUNNING", + message=message, + data=progress, + kpi=kpi, + ) + + def on_validation_start(self, model: Any, dataloader_val: Any, iteration: int = 0) -> None: + self._validation_batches = 0 + self._validation_loss_numerator = 0.0 + self._validation_loss_denominator = 0 + self._validation_local_numerators = [] + self._validation_local_denominators = [] + writer = self._get_writer() + if writer is not None: + writer.write( + status="RUNNING", + message=f"Starting validation for {self._epoch_label(iteration)}", + data={**self._progress_data(iteration), "phase": "validation_starting"}, + ) + log.info(f"Starting validation for {self._epoch_label(iteration)}") + + def on_validation_step_end( + self, + model: Any, + data_batch: dict[str, Any], + output_batch: dict[str, Any], + loss: torch.Tensor, + iteration: int = 0, + ) -> None: + local_numerator, local_denominator = self._local_token_stats(loss, data_batch, output_batch) + self._validation_local_numerators.append(local_numerator) + self._validation_local_denominators.append(local_denominator) + self._validation_batches += 1 + + if self._validation_batches % self.validation_heartbeat_interval != 0: + return + average_loss, _, _ = self._global_token_average(loss, data_batch, output_batch) + writer = self._get_writer() + max_validation_batches = getattr(self.config.trainer, "max_val_iter", None) + batch_progress = ( + f"{self._validation_batches}/{max_validation_batches}" + if max_validation_batches is not None + else str(self._validation_batches) + ) + if writer is not None: + writer.write( + status="RUNNING", + message=( + f"Validation {self._epoch_label(iteration)}, batch {batch_progress} - Loss: {average_loss:.6f}" + ), + data={ + **self._progress_data(iteration), + "phase": "validation_batch_complete", + "validation_batch": self._validation_batches, + "max_validation_batches": max_validation_batches, + }, + kpi={"val/batch_loss": average_loss}, + ) + log.info(f"Validation {self._epoch_label(iteration)}, batch {batch_progress} - Loss: {average_loss:.6f}") + + def on_validation_end(self, model: Any, iteration: int = 0) -> None: + if self._validation_local_numerators: + # Preserve the historical Python-float accumulation order exactly, + # while replacing one CUDA synchronization per scalar per batch + # with one bounded transfer at validation end. + local_numerators = torch.stack(self._validation_local_numerators).detach().cpu().tolist() + local_denominators = torch.stack(self._validation_local_denominators).detach().cpu().tolist() + self._validation_loss_numerator = sum(float(value) for value in local_numerators) + self._validation_loss_denominator = sum(int(value) for value in local_denominators) + print( + "COSMOS_FRAMEWORK_VALIDATION_STATUS_REDUCTION_ATTESTATION " + f"mode=deferred_scalar_transfer batches={self._validation_batches} " + f"rank={os.environ.get('RANK', os.environ.get('LOCAL_RANK', '0'))}", + flush=True, + ) + self._validation_local_numerators = [] + self._validation_local_denominators = [] + numerator, denominator = self._reduce_accumulator( + self._validation_loss_numerator, self._validation_loss_denominator + ) + if denominator == 0: + log.warning("Cosmos validation logging saw zero samples; no val/loss record was written") + return + + self.last_validation_loss = numerator / denominator + writer = self._get_writer() + if writer is not None: + writer.write( + status="RUNNING", + message=( + f"Validation complete for {self._epoch_label(iteration)} - Loss: {self.last_validation_loss:.6f}" + ), + data={ + **self._progress_data(iteration), + "phase": "validation_complete", + "validation_batches": self._validation_batches, + "validation_loss_numerator": numerator, + "validation_valid_label_count": denominator, + }, + kpi={ + "val/loss": self.last_validation_loss, + "val/avg_loss": self.last_validation_loss, + "val/loss_numerator": numerator, + "val/valid_label_count": denominator, + }, + ) + log.info(f"Validation loss ({self._epoch_label(iteration)}): {self.last_validation_loss:.6f}") + + def on_save_checkpoint_success(self, iteration: int = 0, elapsed_time: float = 0) -> None: + checkpoint_path = None + checkpoint_root = Path(self.config.job.path_local) / "checkpoints" + latest = checkpoint_root / "latest_checkpoint.txt" + if latest.is_file(): + # DCP updates this fixed-width marker in place. A callback can + # observe trailing NUL padding while the payload is being + # replaced, so decode only the completed checkpoint-name prefix. + marker = latest.read_bytes() + checkpoint_name = marker.split(b"\x00", 1)[0].decode("utf-8").strip() + if checkpoint_name: + checkpoint_path = str((checkpoint_root / checkpoint_name).resolve()) + writer = self._get_writer() + if writer is not None: + writer.write( + status="RUNNING", + message=f"Checkpoint saved successfully at step {iteration}", + data={ + **self._progress_data(iteration), + "phase": "checkpoint_complete", + "checkpoint_iteration": iteration, + "checkpoint_elapsed_seconds": elapsed_time, + "checkpoint_path": checkpoint_path, + }, + ) + + def on_train_end(self, model: Any, iteration: int = 0) -> None: + numerator, denominator = self._reduce_accumulator(self._train_loss_numerator, self._train_loss_denominator) + if denominator == 0: + raise RuntimeError("Cosmos metric collection observed zero valid training labels") + self.last_training_loss = numerator / denominator + writer = self._get_writer() + if writer is not None: + writer.write( + status="RUNNING", + message=f"Training complete - token-weighted loss: {self.last_training_loss:.6f}", + data={ + **self._progress_data(iteration), + "phase": "training_complete", + "train_loss_numerator": numerator, + "train_valid_label_count": denominator, + }, + kpi={ + "train/avg_loss": self.last_training_loss, + "train/loss_numerator": numerator, + "train/valid_label_count": denominator, + }, + ) + + def on_app_end(self) -> None: + writer = self._get_writer() + if writer is not None: + writer.write( + status="SUCCESS", + message=f"{self._component_name()} training completed successfully", + data=self._progress_data( + int(getattr(getattr(self, "trainer", None), "max_iterations", self.config.trainer.max_iter)) + ), + kpi={ + **({"train/avg_loss": self.last_training_loss} if self.last_training_loss is not None else {}), + **( + {"val/loss": self.last_validation_loss, "val/avg_loss": self.last_validation_loss} + if self.last_validation_loss is not None + else {} + ), + }, + ) + + def on_exception(self, error: BaseException) -> None: + writer = self._get_writer() + if writer is not None: + writer.write( + status="FAILURE", + verbosity="ERROR", + message=f"{self._component_name()} training failed: {error}", + data={"component": self._component_name(), "error_type": type(error).__name__}, + ) diff --git a/cosmos_framework/callbacks/workflow_status_test.py b/cosmos_framework/callbacks/workflow_status_test.py new file mode 100644 index 000000000..f894ac10c --- /dev/null +++ b/cosmos_framework/callbacks/workflow_status_test.py @@ -0,0 +1,159 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +from __future__ import annotations + +import json +from types import SimpleNamespace + +import torch + +from cosmos_framework.callbacks.workflow_status import WorkflowStatusCallback + + +def _callback(tmp_path) -> WorkflowStatusCallback: + callback = WorkflowStatusCallback( + enabled=True, + status_file_path=str(tmp_path / "status.json"), + experiment_name="test", + ) + callback.config = SimpleNamespace( + job=SimpleNamespace(name="job", path_local=str(tmp_path)), + trainer=SimpleNamespace(max_iter=6, max_val_iter=2, logging_iter=1), + ) + callback.trainer = SimpleNamespace(max_iterations=6, num_epochs=2, steps_per_epoch=3) + return callback + + +def _records(tmp_path) -> list[dict]: + return [json.loads(line) for line in (tmp_path / "status.json").read_text(encoding="utf-8").splitlines()] + + +def test_workflow_status_callback_writes_training_validation_and_success(tmp_path) -> None: + callback = _callback(tmp_path) + callback.on_train_start(model=None, iteration=0) + callback.on_training_step_start(model=None, data={}, iteration=0) + callback.on_training_step_batch_end( + model=None, + data_batch={"input_ids": torch.zeros(2, 3)}, + output_batch={"loss_numerator": torch.tensor(6.0), "loss_denominator": torch.tensor(12)}, + loss=torch.tensor(0.5), + iteration=3, + ) + callback.on_training_step_end( + model=None, + data_batch={"input_ids": torch.zeros(2, 3)}, + output_batch={"loss_numerator": torch.tensor(6.0), "loss_denominator": torch.tensor(12)}, + loss=torch.tensor(0.5), + iteration=3, + ) + callback.on_validation_start(model=None, dataloader_val=None, iteration=3) + callback.on_validation_step_end( + model=None, + data_batch={"input_ids": torch.zeros(2, 3)}, + output_batch={"loss_numerator": torch.tensor(2.0), "loss_denominator": torch.tensor(8)}, + loss=torch.tensor(0.25), + iteration=3, + ) + callback.on_validation_end(model=None, iteration=3) + callback.on_train_end(model=None, iteration=6) + callback.on_app_end() + + records = _records(tmp_path) + assert [record["status"] for record in records] == [ + "STARTED", + "RUNNING", + "RUNNING", + "RUNNING", + "RUNNING", + "RUNNING", + "SUCCESS", + ] + assert records[1]["kpi"]["train/step_loss"] == 0.5 + assert records[1]["epoch"] == 1 + assert records[1]["max_epoch"] == 2 + assert records[1]["step_in_epoch"] == 3 + assert records[1]["steps_per_epoch"] == 3 + assert records[1]["max_step"] == 6 + assert records[2]["phase"] == "validation_starting" + assert records[3]["max_validation_batches"] == 2 + assert "epoch 1/2" in records[3]["message"] + assert records[4]["kpi"]["val/avg_loss"] == 0.25 + assert records[4]["kpi"]["val/loss_numerator"] == 2.0 + assert records[4]["kpi"]["val/valid_label_count"] == 8 + assert "epoch 1/2" in records[4]["message"] + assert records[5]["phase"] == "training_complete" + assert records[5]["kpi"]["train/avg_loss"] == 0.5 + assert records[5]["train_loss_numerator"] == 6.0 + assert records[5]["train_valid_label_count"] == 12 + assert records[5]["kpi"]["train/valid_label_count"] == 12 + assert records[-1]["kpi"]["val/loss"] == 0.25 + assert records[-1]["epoch"] == 2 + assert records[-1]["completed_epochs"] == 2 + + +def test_workflow_status_callback_reports_checkpoint_event(tmp_path) -> None: + callback = _callback(tmp_path) + checkpoint = tmp_path / "checkpoints" / "epoch_1" + checkpoint.mkdir(parents=True) + (tmp_path / "checkpoints" / "latest_checkpoint.txt").write_text("epoch_1\n") + callback.on_save_checkpoint_success(iteration=3, elapsed_time=1.25) + record = _records(tmp_path)[-1] + assert record["phase"] == "checkpoint_complete" + assert record["checkpoint_iteration"] == 3 + assert record["checkpoint_elapsed_seconds"] == 1.25 + assert record["checkpoint_path"] == str(checkpoint.resolve()) + + +def test_validation_stats_defer_scalar_transfer_until_end(tmp_path) -> None: + callback = _callback(tmp_path) + callback.validation_heartbeat_interval = 50 + callback.on_validation_start(model=None, dataloader_val=None, iteration=3) + + numerators = [torch.tensor(1.25), torch.tensor(2.5)] + denominators = [torch.tensor(4), torch.tensor(6)] + for numerator, denominator in zip(numerators, denominators): + callback.on_validation_step_end( + model=None, + data_batch={"input_ids": torch.zeros(1, 3)}, + output_batch={ + "loss_numerator": numerator, + "loss_denominator": denominator, + }, + loss=numerator / denominator, + iteration=3, + ) + + assert callback._validation_loss_numerator == 0.0 + assert callback._validation_loss_denominator == 0 + assert len(callback._validation_local_numerators) == 2 + callback.on_validation_end(model=None, iteration=3) + + assert callback._validation_loss_numerator == 3.75 + assert callback._validation_loss_denominator == 10 + assert callback.last_validation_loss == 0.375 + assert callback._validation_local_numerators == [] + assert callback._validation_local_denominators == [] + + +def test_workflow_status_callback_ignores_dcp_marker_nul_padding(tmp_path) -> None: + callback = _callback(tmp_path) + checkpoint = tmp_path / "checkpoints" / "epoch_1" + checkpoint.mkdir(parents=True) + (tmp_path / "checkpoints" / "latest_checkpoint.txt").write_bytes(b"epoch_1\x00\x00\x00") + + callback.on_save_checkpoint_success(iteration=3, elapsed_time=1.25) + + record = _records(tmp_path)[-1] + assert record["phase"] == "checkpoint_complete" + assert record["checkpoint_path"] == str(checkpoint.resolve()) + + +def test_workflow_status_callback_writes_failure(tmp_path) -> None: + callback = _callback(tmp_path) + callback.on_train_start(model=None, iteration=0) + callback.on_exception(RuntimeError("boom")) + + records = _records(tmp_path) + assert records[-1]["status"] == "FAILURE" + assert records[-1]["error_type"] == "RuntimeError" diff --git a/cosmos_framework/checkpoint/base.py b/cosmos_framework/checkpoint/base.py index 59539207f..1d607c8df 100644 --- a/cosmos_framework/checkpoint/base.py +++ b/cosmos_framework/checkpoint/base.py @@ -8,11 +8,11 @@ import torch -from cosmos_framework.utils.config import CheckpointConfig, JobConfig, ObjectStoreConfig -from cosmos_framework.utils.flags import INTERNAL from cosmos_framework.model._base import ImaginaireModel from cosmos_framework.utils import callback +from cosmos_framework.utils.config import CheckpointConfig, JobConfig, ObjectStoreConfig from cosmos_framework.utils.easy_io import easy_io +from cosmos_framework.utils.flags import INTERNAL @dataclass(frozen=True, slots=True) @@ -123,6 +123,7 @@ def save( scheduler: torch.optim.lr_scheduler.LRScheduler, grad_scaler: torch.amp.GradScaler, iteration: int, + epoch: int | None = None, ) -> None: pass diff --git a/cosmos_framework/checkpoint/dcp.py b/cosmos_framework/checkpoint/dcp.py index 30cbea646..0be655f11 100644 --- a/cosmos_framework/checkpoint/dcp.py +++ b/cosmos_framework/checkpoint/dcp.py @@ -78,9 +78,9 @@ ) from cosmos_framework.checkpoint.base import AbstractCheckpointer, CheckpointLoadSource from cosmos_framework.checkpoint.s3_filesystem import S3StorageReader, S3StorageWriter -from cosmos_framework.utils.config import CheckpointConfig, JobConfig from cosmos_framework.model._base import ImaginaireModel from cosmos_framework.utils import callback, distributed, log, misc +from cosmos_framework.utils.config import CheckpointConfig, JobConfig from cosmos_framework.utils.easy_io import easy_io from cosmos_framework.utils.generator.rand_state import get_rand_state_dict, set_rand_state_dict @@ -790,7 +790,7 @@ def keys_to_resume_during_load(self) -> tuple[set[str], CheckpointLoadSource | N # If the path doesn't end with specific checkpoint, read the latest # checkpoint file to determine the most recent checkpoint iteration. - if not re.search(r"/checkpoints/iter_\d{9}/?$", checkpoint_path): + if not re.search(r"/checkpoints/(?:iter_\d{9}|epoch_\d+)/?$", checkpoint_path): old_ckpt_path = checkpoint_path latest_ckpt_path = os.path.join(checkpoint_path, "checkpoints/latest_checkpoint.txt") @@ -1207,6 +1207,7 @@ def save( scheduler: torch.optim.lr_scheduler.LRScheduler, grad_scaler: torch.amp.GradScaler, iteration: int, + epoch: int | None = None, ) -> None: """Save network weights, optimizer parameters, scheduler parameters to a checkpoint. @@ -1223,7 +1224,7 @@ def save( if self.callbacks is not None: self.callbacks.on_save_checkpoint_start(model, iteration) - checkpoint_file = f"iter_{iteration:09}" + checkpoint_file = f"epoch_{epoch}" if epoch is not None else f"iter_{iteration:09}" # Use rank-specific key for RNG state to ensure each rank saves its own state rng_key = f"rng_state_{dist.get_rank()}" @@ -1246,7 +1247,7 @@ def save( self.callbacks.on_save_checkpoint(model, state_dict=to_save_dict) for k in to_save_dict.keys(): - output_dirname = os.path.join(self.save_dirname, f"iter_{iteration:09}/{k}") + output_dirname = os.path.join(self.save_dirname, checkpoint_file, k) to_save_dict[k] = (to_save_dict[k], output_dirname) if self.async_mode == AsyncMode.ASYNC_WITH_PINNED_MEM: diff --git a/cosmos_framework/checkpoint/dcp_distill.py b/cosmos_framework/checkpoint/dcp_distill.py index a1b5279c3..a6921e3a3 100644 --- a/cosmos_framework/checkpoint/dcp_distill.py +++ b/cosmos_framework/checkpoint/dcp_distill.py @@ -26,9 +26,6 @@ from torch.distributed.checkpoint.stateful import Stateful from torch.nn.modules.module import _IncompatibleKeys -from cosmos_framework.model._base import ImaginaireModel -from cosmos_framework.utils import log, misc -from cosmos_framework.utils.easy_io import easy_io from cosmos_framework.checkpoint.dcp import ( AsyncMode, CustomLoadPlanner, @@ -38,8 +35,11 @@ DistributedCheckpointer as _DistributedCheckpointer, ) from cosmos_framework.checkpoint.dcp import ModelWrapper as VFMModelWrapper -from cosmos_framework.utils.generator.rand_state import get_rand_state_dict, set_rand_state_dict +from cosmos_framework.model._base import ImaginaireModel from cosmos_framework.model.generator.distillation.optimizer import OptimizerContainerLike, is_optimizer_container +from cosmos_framework.utils import log, misc +from cosmos_framework.utils.easy_io import easy_io +from cosmos_framework.utils.generator.rand_state import get_rand_state_dict, set_rand_state_dict __all__: tuple[str, ...] = ( "DistributedCheckpointer", @@ -324,6 +324,7 @@ def save( scheduler: Any = None, grad_scaler: torch.amp.GradScaler | None = None, iteration: int = 0, + epoch: int | None = None, ) -> None: if self.async_mode == AsyncMode.ASYNC_WITH_PINNED_MEM: self._wait_for_previous_async_checkpoint() @@ -332,7 +333,7 @@ def save( self.callbacks.on_save_checkpoint_start(model, iteration) model_dict = model.model_dict() - checkpoint_file = f"iter_{iteration:09}" + checkpoint_file = f"epoch_{epoch}" if epoch is not None else f"iter_{iteration:09}" rng_key = f"rng_state_{dist.get_rank()}" to_save_dict: dict[str, Any] = { @@ -360,7 +361,7 @@ def save( to_save_dict["dataloader"] = dataloader_wrapper.state_dict() for key in list(to_save_dict.keys()): - output_dirname = os.path.join(self.save_dirname, f"iter_{iteration:09}/{key}") + output_dirname = os.path.join(self.save_dirname, checkpoint_file, key) to_save_dict[key] = (to_save_dict[key], output_dirname) if self.callbacks is not None: diff --git a/cosmos_framework/checkpoint/dummy.py b/cosmos_framework/checkpoint/dummy.py index 3cfa1babe..443241d96 100644 --- a/cosmos_framework/checkpoint/dummy.py +++ b/cosmos_framework/checkpoint/dummy.py @@ -22,6 +22,7 @@ def save( scheduler: torch.optim.lr_scheduler.LRScheduler, grad_scaler: torch.amp.GradScaler, iteration: int, + epoch: int | None = None, ) -> None: pass diff --git a/cosmos_framework/configs/base/experiment/sft/models/edge_model_config.py b/cosmos_framework/configs/base/experiment/sft/models/edge_model_config.py index 1b2d0c3cb..523dd34eb 100644 --- a/cosmos_framework/configs/base/experiment/sft/models/edge_model_config.py +++ b/cosmos_framework/configs/base/experiment/sft/models/edge_model_config.py @@ -169,7 +169,10 @@ ), tokenizer=L(build_processor_lazy)( repository="nvidia/Cosmos3-Edge", - revision="main", + # Pin the public release used to build the matching DCP checkpoint. + # A symbolic ``main`` forces Hugging Face's CLI to contact the Hub + # even when the exact snapshot is already staged for an offline job. + revision="2a00e87e9976dc3ed5533dd18caf4cdbc3a1bcb2", ), ), ) diff --git a/cosmos_framework/configs/base/reasoner/defaults/callbacks.py b/cosmos_framework/configs/base/reasoner/defaults/callbacks.py index 89e7006d5..9cf2dd951 100644 --- a/cosmos_framework/configs/base/reasoner/defaults/callbacks.py +++ b/cosmos_framework/configs/base/reasoner/defaults/callbacks.py @@ -7,23 +7,24 @@ from hydra.core.config_store import ConfigStore -from cosmos_framework.callbacks.manual_gc import ManualGarbageCollection -from cosmos_framework.utils.lazy_config import PLACEHOLDER -from cosmos_framework.utils.lazy_config import LazyCall as L -from cosmos_framework.utils.callback import LowPrecisionCallback, WandBCallback from cosmos_framework.callbacks.dataloader_state import DataLoaderStateCallback - from cosmos_framework.callbacks.grad_clip import GradClip from cosmos_framework.callbacks.hf_export import HFExportCallback from cosmos_framework.callbacks.iter_speed import IterSpeed from cosmos_framework.callbacks.learning_rate_logger import LearningRateLogger from cosmos_framework.callbacks.log_tensor_shape import LogTensorShapeCallback +from cosmos_framework.callbacks.loss_spike_rollback import LossSpikeRollback +from cosmos_framework.callbacks.manual_gc import ManualGarbageCollection from cosmos_framework.callbacks.param_count import ParamCount from cosmos_framework.callbacks.sampled_media_recorder import SampledMediaRecorder from cosmos_framework.callbacks.tokens_per_sec import VLMTokensPerSec from cosmos_framework.callbacks.wandb_log import WandbCallback as WandBCallbackMultiplier from cosmos_framework.callbacks.wandb_vis import VisualizationLoggingCallback +from cosmos_framework.callbacks.workflow_status import WorkflowStatusCallback from cosmos_framework.configs.base.defaults.job_monitor import JOB_MONITOR_CALLBACKS +from cosmos_framework.utils.callback import LowPrecisionCallback, WandBCallback +from cosmos_framework.utils.lazy_config import PLACEHOLDER +from cosmos_framework.utils.lazy_config import LazyCall as L # from cosmos_framework.utils.callback import NVTXCallback @@ -47,12 +48,22 @@ def register_callbacks(): save_s3="${upload_reproducible_setup}", ), grad_clip=L(GradClip)(clip_norm=1.0, force_finite=False), # use model + # Registered after grad_clip only for readability; the guard reads the gradient + # norm in on_after_backward, which runs before any clipping regardless of order. + loss_spike_rollback=L(LossSpikeRollback)(enabled=False), # use model + optimizer learning_rate_logger=L(LearningRateLogger)(every_n=10), low_precision=L(LowPrecisionCallback)( update_iter=1, config=PLACEHOLDER, trainer=PLACEHOLDER, ), # reads model.precision; no extra kwarg needed + workflow_status=L(WorkflowStatusCallback)( + enabled=False, + status_file_path=None, + experiment_name="", + logging_interval=1, + validation_heartbeat_interval=1, + ), sampled_media=L(SampledMediaRecorder)( enabled=False, output_uri=( diff --git a/cosmos_framework/configs/base/reasoner/defaults/policy_config.py b/cosmos_framework/configs/base/reasoner/defaults/policy_config.py index a48588095..83191332a 100644 --- a/cosmos_framework/configs/base/reasoner/defaults/policy_config.py +++ b/cosmos_framework/configs/base/reasoner/defaults/policy_config.py @@ -33,7 +33,21 @@ class PolicyConfig: # instead of averaging independently normalized microbatch ratios. normalize_weighted_ce_over_accumulation_window: bool = False - # Extra model config + # Parameter-efficient fine-tuning. These fields intentionally live on + # the VLM policy instead of reusing the VFM/MoT model config: the two + # backends have different module names and checkpoint layouts. + lora_enabled: bool = False + lora_rank: int = 16 + lora_alpha: float = 32.0 + lora_dropout: float = 0.0 + lora_target_modules: str = "q_proj,k_proj,v_proj,o_proj" + lora_bias: str = "none" + lora_use_rslora: bool = False + lora_modules_to_save: str = "" + lora_precision: str | None = None + + # Legacy free-form field retained for config compatibility. New recipes + # must use the explicit fields above so PEFT equivalence can be validated. lora: Union[str, None] = None enable_liger_kernel: bool = False # Dense Qwen3.5 caption opt-ins; other recipes retain the standard logits loss. @@ -47,6 +61,27 @@ class PolicyConfig: # "sdpa", or "eager" for fallback. attn_implementation: str = "cosmos" + # Qwen3-VL's patch projection is a non-overlapping Conv3d and is therefore + # algebraically equivalent to a linear projection. ``auto`` selects the + # linear implementation on A100 (SM80), where large FSDP runs have exposed + # a cuDNN Conv3d backward failure. No environment-time monkey patch is + # required. + qwen3_vl_patch_embed: str = "auto" + + def __attrs_post_init__(self) -> None: + if self.lora_rank <= 0: + raise ValueError("lora_rank must be positive") + if self.lora_alpha <= 0: + raise ValueError("lora_alpha must be positive") + if not 0.0 <= self.lora_dropout < 1.0: + raise ValueError("lora_dropout must be in [0, 1)") + if self.lora_bias not in {"none", "all", "lora_only"}: + raise ValueError("lora_bias must be one of: none, all, lora_only") + if self.lora_precision not in {None, "float32", "float16", "bfloat16"}: + raise ValueError("lora_precision must be float32, float16, bfloat16, or unset") + if self.qwen3_vl_patch_embed not in {"auto", "linear", "conv3d"}: + raise ValueError("qwen3_vl_patch_embed must be auto, linear, or conv3d") + @attrs.define(slots=False) class LBLConfig: diff --git a/cosmos_framework/configs/base/reasoner/defaults/vlm_policy.py b/cosmos_framework/configs/base/reasoner/defaults/vlm_policy.py index b5f8d8291..138510d00 100644 --- a/cosmos_framework/configs/base/reasoner/defaults/vlm_policy.py +++ b/cosmos_framework/configs/base/reasoner/defaults/vlm_policy.py @@ -1,6 +1,8 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: OpenMDW-1.1 +import importlib.util + from hydra.core.config_store import ConfigStore from cosmos_framework.configs.base.defaults.reasoner import VLMConfig @@ -65,7 +67,18 @@ # model_type "cosmos3_edge" (native HF metadata, no remote code; classes are # registered in-framework). Reasoner weights load directly from the snapshot — # the training loader follows its root safetensors index; no converter step. -cosmos3_edge_reasoner = PolicyConfig(backbone=VLMConfig(model_name="nvidia/Cosmos3-Edge")) +# The default "cosmos" adapter is Qwen3-VL-specific and rejects Edge's +# explicit attention mask, so Edge must not inherit it. Prefer flash-attn-2 +# where its wheels exist (x86); fall back to SDPA elsewhere (e.g. aarch64). +# An explicit attn_implementation in the experiment TOML still wins. +_EDGE_ATTN_IMPLEMENTATION = ( + "flash_attention_2" if importlib.util.find_spec("flash_attn") is not None else "sdpa" +) + +cosmos3_edge_reasoner = PolicyConfig( + backbone=VLMConfig(model_name="nvidia/Cosmos3-Edge"), + attn_implementation=_EDGE_ATTN_IMPLEMENTATION, +) def register_vlm_policy(): diff --git a/cosmos_framework/configs/base/reasoner/defaults/vlm_policy_test.py b/cosmos_framework/configs/base/reasoner/defaults/vlm_policy_test.py new file mode 100644 index 000000000..2b97b927d --- /dev/null +++ b/cosmos_framework/configs/base/reasoner/defaults/vlm_policy_test.py @@ -0,0 +1,19 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +"""Cosmos3-Edge attention default must match the platform's kernels.""" + +import importlib.util + +from cosmos_framework.configs.base.reasoner.defaults.vlm_policy import cosmos3_edge_reasoner + + +def test_edge_policy_never_defaults_to_cosmos_adapter(): + # The "cosmos" NATTEN adapter is Qwen3-VL-specific and rejects Edge's + # explicit attention mask; Edge must default to a mask-capable impl. + assert cosmos3_edge_reasoner.attn_implementation in {"flash_attention_2", "sdpa"} + + +def test_edge_policy_matches_flash_attn_availability(): + expected = "flash_attention_2" if importlib.util.find_spec("flash_attn") else "sdpa" + assert cosmos3_edge_reasoner.attn_implementation == expected diff --git a/cosmos_framework/configs/base/reasoner/experiment/dataflow_roles.py b/cosmos_framework/configs/base/reasoner/experiment/dataflow_roles.py index cffceb809..07d049c99 100644 --- a/cosmos_framework/configs/base/reasoner/experiment/dataflow_roles.py +++ b/cosmos_framework/configs/base/reasoner/experiment/dataflow_roles.py @@ -6,12 +6,21 @@ from __future__ import annotations +import json +import os +import re +import threading +from collections import OrderedDict +from copy import deepcopy from typing import Any import torch +from PIL import Image from torch.utils.data._utils.collate import default_collate from cosmos_framework.data.generator.dataflow.base import BatchCollator, RawItemProcessor +from cosmos_framework.data.generator.local_datasets.reasoning_qa import apply_reasoning_chat_template +from cosmos_framework.utils.generator.torchcodec_video import TorchCodecVideoReader from cosmos_framework.utils.reasoner.constant import IGNORE_INDEX, PROCESSOR_KEYS_TO_ADD @@ -215,3 +224,369 @@ def _pad_stack(key: str, fill, dtype) -> torch.Tensor: batch["sample_epoch"] = torch.tensor([0] * batch_size) batch["sample_index"] = torch.tensor([0] * batch_size) return batch + + +class _ProcessedVideoCacheProxy: + """On-demand worker-local cache around the HF video preprocessor.""" + + def __init__(self, processor: Any, capacity: int) -> None: + self._processor = processor + self.capacity = int(capacity) + self._entries: OrderedDict[tuple, Any] = OrderedDict() + self._lock = threading.Lock() + self._inflight: dict[tuple, threading.Event] = {} + self._hit_attested = False + + def __getattr__(self, name: str) -> Any: + processor = self.__dict__.get("_processor") + if processor is None: + raise AttributeError(name) + return getattr(processor, name) + + def __getstate__(self) -> dict[str, Any]: + state = self.__dict__.copy() + state.pop("_lock", None) + state["_entries"] = OrderedDict() + state["_inflight"] = {} + state["_hit_attested"] = False + return state + + def __setstate__(self, state: dict[str, Any]) -> None: + self.__dict__.update(state) + self._lock = threading.Lock() + self._inflight = {} + + @staticmethod + def _identity(videos: Any) -> tuple | None: + frames: list[tuple[int, tuple[int, int], str]] = [] + + def collect(value: Any) -> None: + if isinstance(value, Image.Image): + frames.append((id(value), value.size, value.mode)) + elif isinstance(value, (list, tuple)): + for item in value: + collect(item) + + collect(videos) + return tuple(frames) if frames else None + + def __call__(self, *args: Any, **kwargs: Any) -> Any: + videos = kwargs.get("videos", args[0] if args else None) + key = self._identity(videos) + if key is None or self.capacity <= 0: + return self._processor(*args, **kwargs) + + while True: + with self._lock: + cached = self._entries.get(key) + if cached is not None: + self._entries.move_to_end(key) + if not self._hit_attested: + print( + "COSMOS_FRAMEWORK_VALIDATION_PROCESSED_VIDEO_CACHE_HIT_ATTESTATION " + f"rank={os.environ.get('RANK', os.environ.get('LOCAL_RANK', '0'))} " + f"capacity={self.capacity}", + flush=True, + ) + self._hit_attested = True + return deepcopy(cached) + inflight = self._inflight.get(key) + if inflight is None: + inflight = threading.Event() + self._inflight[key] = inflight + owner = True + else: + owner = False + if owner: + break + inflight.wait() + + try: + output = self._processor(*args, **kwargs) + canonical = deepcopy(output) + with self._lock: + self._entries[key] = canonical + self._entries.move_to_end(key) + while len(self._entries) > self.capacity: + self._entries.popitem(last=False) + return output + finally: + with self._lock: + completed = self._inflight.pop(key, None) + if completed is not None: + completed.set() + + +class VideoSFTProcessor(VLMProcessor): + """Convert video-supervision records and uniformly sample media to PIL frames.""" + + @staticmethod + def _resolve_video_device(video_device: str) -> str: + """Bind a generic CUDA request to this torchrun process's local rank.""" + requested = str(video_device) + if requested != "cuda": + return requested + local_rank = os.environ.get("LOCAL_RANK") + if local_rank is None: + return requested + try: + rank = int(local_rank) + except ValueError as exc: + raise ValueError(f"LOCAL_RANK must be an integer, found {local_rank!r}") from exc + if rank < 0: + raise ValueError(f"LOCAL_RANK must be non-negative, found {rank}") + return f"cuda:{rank}" + + def __init__( + self, + processor: Any, + ignore_index: int = IGNORE_INDEX, + num_video_frames: int = 8, + video_cache_size: int = 8, + video_device: str = "cuda", + video_num_threads: int = 1, + processed_video_cache_size: int = 0, + video_max_pixels: int | str | None = 81920, + video_override_map: str | None = None, + system_prompt: str = "", + use_reasoning_chat_template: bool = False, + ) -> None: + super().__init__(processor=processor, ignore_index=ignore_index) + num_video_frames = int(num_video_frames) + video_cache_size = int(video_cache_size) + video_num_threads = int(video_num_threads) + processed_video_cache_size = int(processed_video_cache_size) + if num_video_frames < 1: + raise ValueError("num_video_frames must be >= 1") + if video_cache_size < 0: + raise ValueError("video_cache_size must be >= 0") + if processed_video_cache_size < 0: + raise ValueError("processed_video_cache_size must be >= 0") + self.num_video_frames = num_video_frames + self.video_cache_size = video_cache_size + self.requested_video_device = str(video_device) + self.video_device = self._resolve_video_device(self.requested_video_device) + self.video_num_threads = video_num_threads + self.processed_video_cache_size = processed_video_cache_size + hf_processor = getattr(processor, "processor", None) + video_processor = getattr(hf_processor, "video_processor", None) + if processed_video_cache_size: + if video_processor is None: + raise RuntimeError("processed video caching requires processor.video_processor") + hf_processor.video_processor = _ProcessedVideoCacheProxy( + video_processor, + processed_video_cache_size, + ) + print( + "COSMOS_FRAMEWORK_VALIDATION_PROCESSED_VIDEO_CACHE_ENABLED_ATTESTATION " + f"rank={os.environ.get('RANK', os.environ.get('LOCAL_RANK', '0'))} " + f"capacity={processed_video_cache_size} population=on_demand", + flush=True, + ) + self.video_overrides: dict[str, str] = {} + if video_override_map not in (None, ""): + override_path = os.path.abspath(os.path.expanduser(str(video_override_map))) + with open(override_path, encoding="utf-8") as override_file: + overrides = json.load(override_file) + if not isinstance(overrides, dict) or not all( + isinstance(source, str) and isinstance(target, str) for source, target in overrides.items() + ): + raise ValueError("video_override_map must be a JSON object of string paths") + self.video_overrides = overrides + self.video_max_pixels: int | None = None + if video_max_pixels not in (None, "", 0, "0"): + parsed_video_max_pixels = int(video_max_pixels) + if parsed_video_max_pixels < 1: + raise ValueError("video_max_pixels must be >= 1") + hf_processor = getattr(processor, "processor", processor) + video_processor = getattr(hf_processor, "video_processor", None) + size = getattr(video_processor, "size", None) + if not isinstance(size, dict): + raise ValueError("video_max_pixels requires a processor.video_processor.size mapping") + shortest_edge = size.get("shortest_edge") + if shortest_edge is not None and parsed_video_max_pixels < int(shortest_edge): + raise ValueError( + f"video_max_pixels ({parsed_video_max_pixels}) must be >= shortest_edge ({shortest_edge})" + ) + size["longest_edge"] = parsed_video_max_pixels + self.video_max_pixels = parsed_video_max_pixels + self.system_prompt = system_prompt + self.use_reasoning_chat_template = use_reasoning_chat_template + if self.use_reasoning_chat_template: + apply_reasoning_chat_template(processor) + self._video_cache: OrderedDict[str, tuple[list[Image.Image], float]] = OrderedDict() + self._video_cache_lock = threading.Lock() + self._video_inflight: dict[str, threading.Event] = {} + self._video_runtime_attested = False + + def __getstate__(self) -> dict[str, Any]: + """Drop process-local synchronization and cache state before spawn.""" + state = self.__dict__.copy() + state.pop("_video_cache_lock", None) + state["_video_cache"] = OrderedDict() + state["_video_inflight"] = {} + state["_video_runtime_attested"] = False + return state + + def __setstate__(self, state: dict[str, Any]) -> None: + """Recreate rank-local cache synchronization in a spawned worker.""" + self.__dict__.update(state) + self._video_cache_lock = threading.Lock() + self._video_inflight = {} + + def _decode_video(self, video_path: str) -> tuple[list[Image.Image], float]: + video_path = self.video_overrides.get(video_path, video_path) + video_path = os.path.abspath(os.path.expanduser(video_path)) + if self.video_cache_size > 0: + # Concurrent processing can request the same source video in one + # logical pool. Elect one decoder and let peers consume its cached + # result, avoiding duplicate GPU decoder sessions without prewarm. + while True: + with self._video_cache_lock: + cached = self._video_cache.get(video_path) + if cached is not None: + self._video_cache.move_to_end(video_path) + return cached + inflight = self._video_inflight.get(video_path) + if inflight is None: + inflight = threading.Event() + self._video_inflight[video_path] = inflight + decode_owner = True + else: + decode_owner = False + if decode_owner: + break + inflight.wait() + + try: + reader = TorchCodecVideoReader( + video_path, + num_threads=self.video_num_threads, + device=self.video_device, + ) + total_frames = len(reader) + if total_frames < 1: + raise ValueError(f"video-supervision media has zero frames: {video_path}") + sample_count = min(self.num_video_frames, total_frames) + if sample_count == 1: + indices = [0] + else: + indices = torch.linspace(0, total_frames - 1, steps=sample_count).round().to(dtype=torch.long).tolist() + frames_np = reader.get_frames_nhwc_uint8(indices) + decoded_device = str(reader.last_output_device) + if self.video_device.startswith("cuda"): + requested_device = torch.device(self.video_device) + actual_device = torch.device(decoded_device) + if actual_device.type != "cuda" or ( + requested_device.index is not None and actual_device.index != requested_device.index + ): + raise RuntimeError( + "TorchCodec did not decode on the requested CUDA device: " + f"requested={self.video_device} actual={decoded_device}" + ) + frames = [Image.fromarray(frame) for frame in frames_np] + + with self._video_cache_lock: + if not self._video_runtime_attested: + print( + "COSMOS_FRAMEWORK_VIDEO_RUNTIME " + f"rank={os.environ.get('RANK', os.environ.get('LOCAL_RANK', '0'))} " + "backend=torchcodec " + f"requested_device={self.requested_video_device} " + f"resolved_device={self.video_device} actual_device={decoded_device} " + f"video_cache_size={self.video_cache_size} " + f"decoder_threads={self.video_num_threads}", + flush=True, + ) + self._video_runtime_attested = True + + source_fps = reader.get_avg_fps() + average_stride = (indices[-1] - indices[0]) / max(len(indices) - 1, 1) if len(indices) > 1 else 1.0 + effective_fps = source_fps / max(average_stride, 1.0) + decoded = (frames, float(effective_fps)) + if self.video_cache_size > 0: + with self._video_cache_lock: + self._video_cache[video_path] = decoded + self._video_cache.move_to_end(video_path) + while len(self._video_cache) > self.video_cache_size: + self._video_cache.popitem(last=False) + return decoded + finally: + if self.video_cache_size > 0: + with self._video_cache_lock: + completed = self._video_inflight.pop(video_path, None) + if completed is not None: + completed.set() + + def _sharegpt_to_openai(self, item: dict) -> list[dict]: + if "messages" in item: + messages = deepcopy(item["messages"]) + video_inserted = False + for message in messages: + content = message.get("content") + if not isinstance(content, list): + continue + for part in content: + if part.get("type") != "video": + continue + video_path = part.get("video") + if not isinstance(video_path, str): + raise TypeError("task-aware video content must contain a string path") + frames, fps = self._decode_video(video_path) + part["video"] = frames + part["fps"] = fps + video_inserted = True + if not video_inserted and isinstance(item.get("video"), str): + frames, fps = self._decode_video(item["video"]) + for message in messages: + if message.get("role") != "user": + continue + content = message.get("content", "") + message["content"] = [ + {"type": "video", "video": frames, "fps": fps}, + {"type": "text", "text": content if isinstance(content, str) else ""}, + ] + break + return messages + + conversations = item.get("conversations", []) + video_path = item.get("video") + frames, fps = self._decode_video(video_path) + messages: list[dict] = [] + video_inserted = False + if self.system_prompt: + messages.append({"role": "system", "content": self.system_prompt}) + + for turn in conversations: + role = "user" if turn["from"] == "human" else "assistant" + text = re.sub(r"(\n)?(\n)?", "", turn["value"]).strip() + if role == "user" and not video_inserted: + content: Any = [ + {"type": "video", "video": frames, "fps": fps}, + {"type": "text", "text": text}, + ] + video_inserted = True + else: + content = text + messages.append({"role": role, "content": content}) + return messages + + def process(self, item: dict) -> dict: + sample = super().process(item) + video_path = item.get("video") + if isinstance(video_path, str): + video_path = self.video_overrides.get(video_path, video_path) + sample["cosmos_video_cache_key"] = os.path.realpath(os.path.abspath(os.path.expanduser(video_path))) + return sample + + +class VideoVLMCollator(VLMCollator): + """Preserve one stable video identity per sample for validation caching.""" + + def collate(self, samples: list[dict]) -> dict: + cache_keys = [sample.get("cosmos_video_cache_key") for sample in samples] + batch = super().collate(samples) + batch.pop("cosmos_video_cache_key", None) + if all(isinstance(key, str) for key in cache_keys): + batch["cosmos_video_cache_keys"] = cache_keys + return batch diff --git a/cosmos_framework/configs/base/reasoner/experiment/video_sft.py b/cosmos_framework/configs/base/reasoner/experiment/video_sft.py new file mode 100644 index 000000000..48c5dbcad --- /dev/null +++ b/cosmos_framework/configs/base/reasoner/experiment/video_sft.py @@ -0,0 +1,261 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +"""Dataset-neutral Nano and Edge video SFT experiment registrations.""" + +import os +from copy import deepcopy + +from hydra.core.config_store import ConfigStore + +from cosmos_framework.callbacks.cosmos_dataloader_state import CosmosDataLoaderStateCallback +from cosmos_framework.configs.base.reasoner.experiment.dataflow_roles import ( + VideoSFTProcessor, + VideoVLMCollator, + VLMCollator, +) +from cosmos_framework.data.generator.dataflow import ContiguousBatcher, CosmosDataLoader, MapDistributor +from cosmos_framework.data.generator.dataflow.distributors import MediaGroupedMapDistributor +from cosmos_framework.data.generator.local_datasets.reasoning_qa import ReasoningQADataset, VideoConversationDataset +from cosmos_framework.data.generator.processors import build_processor +from cosmos_framework.utils.lazy_config import LazyCall as L +from cosmos_framework.utils.lazy_config import LazyDict +from cosmos_framework.utils.reasoner.constant import IGNORE_INDEX + + +def _video_conversation_dataloader( + *, + annotation_env: str, + media_env: str, + limit_env: str, + shuffle: bool, + frame_env: str = "COSMOS_VIDEO_NUM_FRAMES", + cache_env: str = "COSMOS_VIDEO_CACHE_SIZE", + max_pixels_env: str = "COSMOS_VIDEO_MAX_PIXELS", + system_prompt_env: str = "COSMOS_VIDEO_SYSTEM_PROMPT", +) -> LazyDict: + validation_grouped = ( + not shuffle and os.environ.get("COSMOS_FRAMEWORK_VALIDATION_SHARD_STRATEGY", "stride") == "media_grouped" + ) + distributor_cls = MediaGroupedMapDistributor if validation_grouped else MapDistributor + max_batch_size = int(os.environ.get("COSMOS_FRAMEWORK_VALIDATION_BATCH_SIZE", "1")) if not shuffle else 1 + return L(CosmosDataLoader)( + distributor=L(distributor_cls)( + dataset=L(VideoConversationDataset)( + annotation_path=f"${{oc.env:{annotation_env}}}", + media_path=f"${{oc.env:{media_env}}}", + limit=f"${{oc.env:{limit_env},''}}", + ), + shuffle=shuffle, + seed="${oc.env:COSMOS_DATALOADER_SEED,42}", + name="train" if shuffle else "val", + ), + processor=L(VideoSFTProcessor)( + processor=L(build_processor)( + tokenizer_type="${model.config.policy.backbone.model_name}", + config_variant="hf", + ), + ignore_index=IGNORE_INDEX, + num_video_frames=f"${{oc.env:{frame_env},8}}", + video_cache_size=f"${{oc.env:{cache_env},8}}", + video_device="${oc.env:COSMOS_VIDEO_DECODER_DEVICE,cuda}", + video_num_threads="${oc.env:COSMOS_VIDEO_DECODER_THREADS,1}", + # Training was pinned to 0, so every epoch re-decoded every video. That is + # the dominant cost of a training step here: measured on a GB300, the wall + # step splits 2.00s waiting on the dataloader against 0.62s of compute, so + # decode is roughly three quarters of the run. A cache large enough to hold + # the split turns epochs after the first into cache hits. + # + # Left at 0 by default because the cache is per dataloader worker and holds + # decoded frames, so capacity has to be chosen against the dataset size and + # available host memory rather than assumed. + processed_video_cache_size=( + "${oc.env:COSMOS_FRAMEWORK_VALIDATION_PROCESSED_VIDEO_CACHE_SIZE,0}" + if not shuffle + else "${oc.env:COSMOS_FRAMEWORK_TRAIN_PROCESSED_VIDEO_CACHE_SIZE,0}" + ), + video_max_pixels=f"${{oc.env:{max_pixels_env},81920}}", + video_override_map="${oc.env:COSMOS_VIDEO_OVERRIDE_MAP,''}", + system_prompt=f"${{oc.env:{system_prompt_env},''}}", + ), + batcher=L(ContiguousBatcher)( + max_batch_size=max_batch_size, + max_tokens=81920, + drop_last=False, + ), + collator=L(VideoVLMCollator)(), + num_workers="${oc.env:COSMOS_FRAMEWORK_DATALOADER_NUM_WORKERS,1}", + prefetch_factor="${oc.env:COSMOS_FRAMEWORK_DATALOADER_PREFETCH_FACTOR,4}", + persistent_workers=True, + pin_memory=True, + multiprocessing_context="spawn", + processing_threads="${oc.env:COSMOS_FRAMEWORK_SFT_PROCESS_THREADS,8}", + ) + + +def _task_aware_video_dataloader( + *, + split: str, + shuffle: bool, + annotation_env: str | None = None, + media_env: str | None = None, + limit_env: str | None = None, + frame_env: str = "COSMOS_VIDEO_NUM_FRAMES", + cache_env: str = "COSMOS_VIDEO_CACHE_SIZE", + max_pixels_env: str = "COSMOS_VIDEO_MAX_PIXELS", + system_prompt_env: str = "COSMOS_VIDEO_SYSTEM_PROMPT", +) -> LazyDict: + annotation_env = annotation_env or f"COSMOS_VIDEO_{split.upper()}_ANNOTATIONS" + media_env = media_env or f"COSMOS_VIDEO_{split.upper()}_MEDIA_ROOTS" + limit_env = limit_env or f"COSMOS_VIDEO_{split.upper()}_LIMIT" + return L(CosmosDataLoader)( + distributor=L(MapDistributor)( + dataset=L(ReasoningQADataset)( + annotation_paths=f"${{oc.env:{annotation_env}}}", + media_root=f"${{oc.env:{media_env}}}", + response_mode="hybrid" if split == "train" else "answer", + system_prompt=f"${{oc.env:{system_prompt_env},''}}", + vision_kwargs={}, + max_samples=f"${{oc.env:{limit_env},''}}", + ), + shuffle=shuffle, + seed="${oc.env:COSMOS_DATALOADER_SEED,42}", + name=split, + ), + processor=L(VideoSFTProcessor)( + processor=L(build_processor)( + tokenizer_type="${model.config.policy.backbone.model_name}", + config_variant="hf", + ), + ignore_index=IGNORE_INDEX, + num_video_frames=f"${{oc.env:{frame_env},8}}", + video_cache_size=f"${{oc.env:{cache_env},8}}", + video_device="${oc.env:COSMOS_VIDEO_DECODER_DEVICE,cuda}", + video_num_threads="${oc.env:COSMOS_VIDEO_DECODER_THREADS,1}", + video_max_pixels=f"${{oc.env:{max_pixels_env},81920}}", + video_override_map="${oc.env:COSMOS_VIDEO_OVERRIDE_MAP,''}", + system_prompt="", + use_reasoning_chat_template=True, + ), + batcher=L(ContiguousBatcher)( + max_batch_size=1, + max_tokens=81920, + drop_last=False, + ), + collator=L(VLMCollator)(), + num_workers="${oc.env:COSMOS_FRAMEWORK_DATALOADER_NUM_WORKERS,1}", + prefetch_factor="${oc.env:COSMOS_FRAMEWORK_DATALOADER_PREFETCH_FACTOR,2}", + persistent_workers=True, + pin_memory=False, + multiprocessing_context="spawn", + processing_threads="${oc.env:COSMOS_FRAMEWORK_SFT_PROCESS_THREADS,8}", + ) + + +cosmos_video_conversation = LazyDict( + dict( + defaults=[ + {"override /checkpoint": "local"}, + {"override /data_train": None}, + {"override /data_val": None}, + {"override /model": "vlm_fsdp"}, + {"override /vlm_policy": "qwen3_vl_8b_instruct"}, + {"override /callbacks": ["basic_vlm", "basic_log"]}, + "_self_", + ], + job=dict( + project="cosmos3_reasoner", + group="cosmos_video_conversation_sft", + wandb_mode="disabled", + ), + trainer=dict( + callbacks=dict( + dataloader_state=L(CosmosDataLoaderStateCallback)(), + workflow_status=dict( + enabled=True, + logging_interval=1, + validation_heartbeat_interval=1, + ), + ), + max_iter=10, + logging_iter=1, + run_validation=True, + validation_iter=10, + max_val_iter=10, + run_validation_on_start=False, + grad_accum_iter=1, + ), + optimizer=dict( + lr=1.0e-4, + fused=True, + weight_decay=0.01, + betas=[0.9, 0.999], + lr_multipliers={"model.visual": 1.0}, + ), + model=dict( + config=dict( + policy=dict( + model_max_length=81920, + qwen_max_video_token_length=8192, + ), + freeze=dict(trainable_params=[".*"]), + parallelism=dict( + data_parallel_shard_degree=4, + data_parallel_replicate_degree=1, + ), + ), + ), + data_setting=dict( + max_tokens=81920, + qwen_max_video_token_length=8192, + ), + checkpoint=dict( + save_iter=100, + load_from_object_store=dict(enabled=False, credentials="", bucket=""), + save_to_object_store=dict(enabled=False, credentials="", bucket=""), + ), + dataloader_train=_video_conversation_dataloader( + annotation_env="COSMOS_VIDEO_TRAIN_ANNOTATION", + media_env="COSMOS_VIDEO_TRAIN_MEDIA", + limit_env="COSMOS_VIDEO_TRAIN_LIMIT", + shuffle=True, + ), + dataloader_val=_video_conversation_dataloader( + annotation_env="COSMOS_VIDEO_VAL_ANNOTATION", + media_env="COSMOS_VIDEO_VAL_MEDIA", + limit_env="COSMOS_VIDEO_VAL_LIMIT", + shuffle=False, + ), + upload_reproducible_setup=False, + ), + flags={"allow_objects": True}, +) + + +cosmos_task_aware_video_reasoning = deepcopy(cosmos_video_conversation) +cosmos_task_aware_video_reasoning["job"]["group"] = "cosmos_task_aware_video_reasoning_sft" +cosmos_task_aware_video_reasoning["dataloader_train"] = _task_aware_video_dataloader(split="train", shuffle=True) +cosmos_task_aware_video_reasoning["dataloader_val"] = _task_aware_video_dataloader(split="val", shuffle=False) + + +def _edge_recipe(recipe: LazyDict, group: str) -> LazyDict: + edge = deepcopy(recipe) + edge["defaults"][4] = {"override /vlm_policy": "cosmos3_edge_reasoner"} + edge["job"]["group"] = group + edge["optimizer"].pop("lr_multipliers", None) + edge["model"]["config"]["policy"]["model_max_length"] = 16000 + return edge + + +cosmos_video_conversation_edge = _edge_recipe(cosmos_video_conversation, "cosmos_video_conversation_edge_sft") +cosmos_task_aware_video_reasoning_edge = _edge_recipe( + cosmos_task_aware_video_reasoning, "cosmos_task_aware_video_reasoning_edge_sft" +) + +for name, node in ( + ("cosmos_video_conversation", cosmos_video_conversation), + ("cosmos_task_aware_video_reasoning", cosmos_task_aware_video_reasoning), + ("cosmos_video_conversation_edge", cosmos_video_conversation_edge), + ("cosmos_task_aware_video_reasoning_edge", cosmos_task_aware_video_reasoning_edge), +): + ConfigStore.instance().store(group="experiment", package="_global_", name=name, node=node) diff --git a/cosmos_framework/configs/base/reasoner/experiment/video_sft_test.py b/cosmos_framework/configs/base/reasoner/experiment/video_sft_test.py new file mode 100644 index 000000000..bbe09cf53 --- /dev/null +++ b/cosmos_framework/configs/base/reasoner/experiment/video_sft_test.py @@ -0,0 +1,551 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: OpenMDW-1.1 + +from __future__ import annotations + +import json +import pickle +import threading +import time +import types +from collections import OrderedDict +from concurrent.futures import ThreadPoolExecutor + +import pytest +from PIL import Image + +from cosmos_framework.configs.base.reasoner.experiment import video_sft +from cosmos_framework.configs.base.reasoner.experiment.video_sft import ( + VideoConversationDataset, + VideoSFTProcessor, + cosmos_task_aware_video_reasoning, + cosmos_task_aware_video_reasoning_edge, + cosmos_video_conversation_edge, +) +from cosmos_framework.data.generator.dataflow import ContiguousBatcher +from cosmos_framework.data.generator.local_datasets.reasoning_qa import ( + ReasoningQADataset, + apply_reasoning_chat_template, + parse_path_list, +) + + +def _install_fake_reasoning(monkeypatch) -> list[object]: + calls: list[object] = [] + + class FakeDataset: + def __init__(self, **kwargs) -> None: + calls.append(kwargs) + self._raw_length = 2 + + def __getitem__(self, index: int) -> list[dict]: + return [{"role": "assistant", "content": f"item-{index}"}] + + def fake_template(processor) -> None: + calls.append(processor) + + from cosmos_framework.data.reasoner import qa_dataset + + monkeypatch.setattr(qa_dataset, "ReasoningConversationDataset", FakeDataset) + monkeypatch.setattr(qa_dataset, "apply_chat_template_override", fake_template) + return calls + + +def test_video_conversation_dataset_resolves_media_paths_and_limit(tmp_path) -> None: + media = tmp_path / "videos" + media.mkdir() + (media / "clip.mp4").write_bytes(b"video") + annotations = tmp_path / "annotations.json" + annotations.write_text( + json.dumps( + [ + { + "video": "clip.mp4", + "conversations": [ + {"from": "human", "value": "