From 518b8d0824d0eb2b61c5201bfd0f826183708841 Mon Sep 17 00:00:00 2001 From: Daoyuan Li <94409450+DaoyuanLi2816@users.noreply.github.com> Date: Wed, 12 Aug 2026 00:09:08 -0700 Subject: [PATCH 1/3] Conform forward top-k loss with verl v0.8 --- .github/workflows/verl-bridge.yml | 4 +- PROJECT_STATE.md | 18 ++ pyproject.toml | 1 + src/miniverl/cache/store.py | 69 ++++++++ src/miniverl/config/__init__.py | 2 + src/miniverl/config/models.py | 31 ++++ src/miniverl/losses/chunked.py | 37 +++++ src/miniverl/losses/verl_topk.py | 99 +++++++++++ src/miniverl/schemas/cache.py | 6 + src/miniverl/teachers/local.py | 73 +++++--- src/miniverl/training/trainer.py | 156 +++++++++++++----- tests/conformance/test_verl_v08_loss.py | 141 ++++++++++++++++ tests/integration/test_prompt_opd_pipeline.py | 22 ++- tests/unit/test_cache.py | 48 ++++++ tests/unit/test_prompt_source_config.py | 23 +++ tests/unit/test_verl_forward_kl_topk.py | 70 ++++++++ 16 files changed, 738 insertions(+), 62 deletions(-) create mode 100644 src/miniverl/losses/verl_topk.py create mode 100644 tests/conformance/test_verl_v08_loss.py create mode 100644 tests/unit/test_verl_forward_kl_topk.py diff --git a/.github/workflows/verl-bridge.yml b/.github/workflows/verl-bridge.yml index 0b9af8e..5e0ab97 100644 --- a/.github/workflows/verl-bridge.yml +++ b/.github/workflows/verl-bridge.yml @@ -47,7 +47,7 @@ jobs: changed=$(git diff --name-only "origin/$BASE_REF"...HEAD) echo "changed files:" echo "$changed" - if echo "$changed" | grep -Eq '^(\.github/workflows/verl-bridge\.yml|pyproject\.toml|scripts/[^/]*verl_bridge[^/]*|src/miniverl/bridge/.*|tests/.*verl_bridge.*)$'; then + if echo "$changed" | grep -Eq '^(\.github/workflows/verl-bridge\.yml|pyproject\.toml|scripts/[^/]*verl_bridge[^/]*|src/miniverl/bridge/.*|src/miniverl/losses/verl_topk\.py|tests/.*verl_bridge.*|tests/conformance/test_verl_v08_loss\.py)$'; then echo "bridge=true" >> "$GITHUB_OUTPUT" else echo "bridge=false" >> "$GITHUB_OUTPUT" @@ -73,6 +73,8 @@ jobs: run: >- python -m pip install --no-deps "git+https://github.com/verl-project/verl.git@7aed6b230776f963fa09509c10d9c3a767d1102c" + - name: Compare forward_kl_topk values, diagnostics and gradients + run: pytest -q tests/conformance/test_verl_v08_loss.py -m verl_conformance - name: Generate and export the standards smoke bundle run: | python scripts/prepare_verl_bridge_smoke.py --out _verl-smoke-source diff --git a/PROJECT_STATE.md b/PROJECT_STATE.md index a7e4d2d..b6de58b 100644 --- a/PROJECT_STATE.md +++ b/PROJECT_STATE.md @@ -46,6 +46,24 @@ The pure-OPD integration executes two Parquet prompts through rollout, teacher scoring, a padded actor update, checkpoint and manifest without an environment or reward. +## v0.8.0 single-GPU verl OPD pivot — PR C + +The compatibility loss is now a separate `forward_kl_topk` implementation; +native `bucketed_topk_tail` is unchanged. The supported profile requires +forward KL, upstream-compatible `token-mean` aggregation, temperature 1.0 and +no sampled-token NLL mixture. Deterministic tensor tests import the official +verl v0.8.0 loss implementation at pinned commit +`7aed6b230776f963fa09509c10d9c3a767d1102c` and compare per-token loss, +teacher/student top-k mass, overlap diagnostics, the reduced scalar and student +gradients at `rtol=1e-6`, `atol=1e-7`. + +Teacher cache entries now bind the prompt-row digest, exact actor response +token IDs and policy version; the cache identity additionally binds teacher, +tokenizer, top-k, temperature and score-implementation version. A mismatched +response or implementation fails closed. The prompt-data integration exercises +rollout, exact-state scoring, cache reload, one token-mean optimizer update and +the emitted verl diagnostics. + ## v0.7.1 Product correction — RELEASE CANDIDATE Branch `v0.7.1-product-correction` starts from synchronized main diff --git a/pyproject.toml b/pyproject.toml index 6beff07..97aab6b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -195,6 +195,7 @@ markers = [ "gpu: requires a CUDA GPU (deselected in CPU CI)", "torch: requires the [train] extra (torch/transformers/peft)", "network: requires network access (deselected in CI)", + "verl_conformance: imports the exact pinned official verl source", "slow: takes more than a few seconds", ] diff --git a/src/miniverl/cache/store.py b/src/miniverl/cache/store.py index 720a3d6..ad0b68d 100644 --- a/src/miniverl/cache/store.py +++ b/src/miniverl/cache/store.py @@ -63,6 +63,23 @@ logger = get_logger("cache") +def _binding_checksum( + *, + prompt_row_digest: str | None, + actor_response_token_ids: list[int] | None, + policy_version: int, + score_implementation_version: str | None, +) -> str: + payload = { + "actor_response_token_ids": actor_response_token_ids, + "policy_version": policy_version, + "prompt_row_digest": prompt_row_digest, + "score_implementation_version": score_implementation_version, + } + encoded = json.dumps(payload, sort_keys=True, separators=(",", ":")).encode("utf-8") + return hashlib.sha256(encoded).hexdigest() + + def _replace_shard_file(source: Path, target: Path) -> None: source.replace(target) @@ -161,6 +178,7 @@ def create( top_k: int, temperature: float, loss_mode: str, + score_implementation_version: str | None = None, dtype: str = "float32", entries_per_shard: int = 32, overwrite: bool = False, @@ -200,6 +218,7 @@ def create( top_k=top_k, temperature=temperature, loss_mode=loss_mode, + score_implementation_version=score_implementation_version, dtype=dtype, entries_per_shard=entries_per_shard, ) @@ -251,6 +270,7 @@ def assert_compatible( top_k: int, temperature: float, loss_mode: str, + score_implementation_version: str | None = None, dtype: str, ) -> None: """Reject reuse when any objective or teacher identity component changed.""" @@ -276,6 +296,8 @@ def assert_compatible( "loss_mode": loss_mode, "dtype": dtype, } + if score_implementation_version is not None: + expected["score_implementation_version"] = score_implementation_version if self.index.schema_version >= 2: expected["tokenizer_identity"] = dict(tokenizer_identity or {}) expected["teacher_adapter_provenance"] = ( @@ -344,6 +366,12 @@ def write( span_counts: dict[str, int] = {} for name in batch.span_types: span_counts[name] = span_counts.get(name, 0) + 1 + binding_checksum = _binding_checksum( + prompt_row_digest=batch.prompt_row_digest, + actor_response_token_ids=batch.actor_response_token_ids, + policy_version=batch.policy_version, + score_implementation_version=self.index.score_implementation_version, + ) self._pending[batch.trajectory_id] = { "tensors": tensors, @@ -355,6 +383,9 @@ def write( "tail_is_exact_zero": tail_is_exact_zero, "selected_span_types": span_counts, "ordered_span_types": list(batch.span_types), + "prompt_row_digest": batch.prompt_row_digest, + "actor_response_token_ids": batch.actor_response_token_ids, + "binding_checksum": binding_checksum, }, } self._pending_order.append(batch.trajectory_id) @@ -375,6 +406,9 @@ def write( checksum=digest.hexdigest(), selected_span_types=span_counts, ordered_span_types=list(batch.span_types), + prompt_row_digest=batch.prompt_row_digest, + actor_response_token_ids=batch.actor_response_token_ids, + binding_checksum=binding_checksum, ) def _next_shard_name(self) -> str: @@ -430,6 +464,9 @@ def flush(self) -> None: checksum=meta["checksum"], selected_span_types=meta["selected_span_types"], ordered_span_types=meta["ordered_span_types"], + prompt_row_digest=meta["prompt_row_digest"], + actor_response_token_ids=meta["actor_response_token_ids"], + binding_checksum=meta["binding_checksum"], ) self._write_index(next_index) self.index = next_index @@ -452,6 +489,8 @@ def read( trajectory_id: str, *, expect_policy_version: int | None = None, + expect_prompt_row_digest: str | None = None, + expect_actor_response_token_ids: list[int] | None = None, device: str = "cpu", ) -> TeacherTargetBatch: """Load one trajectory's targets, enforcing the policy-version contract.""" @@ -471,6 +510,34 @@ def read( hint="that would make the update off-policy. Re-score the trajectory, " "or switch to run.mode=offline_kd if fixed targets are intended.", ) + if ( + expect_prompt_row_digest is not None + and entry.prompt_row_digest != expect_prompt_row_digest + ): + raise StaleCacheError( + f"teacher targets for {trajectory_id!r} have prompt-row digest " + f"{entry.prompt_row_digest!r}, expected {expect_prompt_row_digest!r}" + ) + if ( + expect_actor_response_token_ids is not None + and entry.actor_response_token_ids != expect_actor_response_token_ids + ): + raise StaleCacheError( + f"teacher targets for {trajectory_id!r} do not match the exact actor " + "response token IDs" + ) + if entry.binding_checksum is not None: + actual_binding = _binding_checksum( + prompt_row_digest=entry.prompt_row_digest, + actor_response_token_ids=entry.actor_response_token_ids, + policy_version=entry.policy_version, + score_implementation_version=self.index.score_implementation_version, + ) + if actual_binding != entry.binding_checksum: + raise CacheCorruptionError( + f"binding checksum mismatch for {trajectory_id!r}: expected " + f"{entry.binding_checksum[:16]}..., got {actual_binding[:16]}..." + ) shard_path = self.path / entry.shard if not shard_path.is_file(): raise CacheCorruptionError(f"shard {entry.shard} referenced by the index is missing") @@ -516,6 +583,8 @@ def read( temperature=entry.temperature, top_k=entry.top_k, span_types=span_types, + prompt_row_digest=entry.prompt_row_digest, + actor_response_token_ids=entry.actor_response_token_ids, ) def __contains__(self, trajectory_id: object) -> bool: diff --git a/src/miniverl/config/__init__.py b/src/miniverl/config/__init__.py index 72a5937..55871f6 100644 --- a/src/miniverl/config/__init__.py +++ b/src/miniverl/config/__init__.py @@ -14,6 +14,7 @@ GateConfig, GateSignal, LoRAConfig, + LossAggregation, LossConfig, LossMode, MemoryConfig, @@ -60,6 +61,7 @@ "GateSignal", "LoRAConfig", "LossConfig", + "LossAggregation", "LossMode", "MemoryConfig", "MemoryStrategy", diff --git a/src/miniverl/config/models.py b/src/miniverl/config/models.py index 03c7740..dd77c31 100644 --- a/src/miniverl/config/models.py +++ b/src/miniverl/config/models.py @@ -45,6 +45,7 @@ "Quantization", "TeacherContextMode", "LossMode", + "LossAggregation", "Divergence", "SelectorName", "MemoryStrategy", @@ -152,6 +153,14 @@ class LossMode(str, Enum): EXACT_FULL_VOCAB = "exact_full_vocab" BUCKETED_TOPK_TAIL = "bucketed_topk_tail" + VERL_FORWARD_KL_TOPK = "forward_kl_topk" + + +class LossAggregation(str, Enum): + """How selected token losses form one optimizer-step scalar.""" + + NATIVE_PER_TRAJECTORY = "native_per_trajectory" + TOKEN_MEAN = "token-mean" class Divergence(str, Enum): @@ -422,12 +431,15 @@ class LossConfig(_Base): """Divergence objective and its vocabulary treatment.""" mode: LossMode = LossMode.BUCKETED_TOPK_TAIL + aggregation: LossAggregation = LossAggregation.NATIVE_PER_TRAJECTORY divergence: Divergence = Divergence.REVERSE_KL temperature: float = Field(default=1.0, gt=0.0, le=20.0) scale_by_temperature_squared: bool = True top_k: int = Field(default=64, ge=1, le=262144) jsd_beta: float = Field(default=0.5, ge=0.0, le=1.0) tail_epsilon: float = Field(default=1e-9, gt=0.0, lt=1e-2) + log_prob_min_clamp: float | None = Field(default=None, le=0.0) + loss_max_clamp: float | None = Field(default=None, gt=0.0) #: Number of selected prediction positions projected through the LM head at #: once. Purely a memory/throughput knob -- it does not change the loss. chunk_size: int = Field(default=256, ge=1, le=65536) @@ -846,6 +858,25 @@ def _validate_combination(self) -> RunConfig: # side effect of merely constructing a RunConfig. self.loss = self.loss.model_copy(update={"top_k": 1}) + if self.loss.mode is LossMode.VERL_FORWARD_KL_TOPK: + if self.loss.divergence is not Divergence.FORWARD_KL: + raise ValueError("loss.mode=forward_kl_topk requires loss.divergence=forward_kl") + if self.loss.aggregation is not LossAggregation.TOKEN_MEAN: + raise ValueError( + "loss.mode=forward_kl_topk requires loss.aggregation=token-mean " + "for the supported verl v0.8 profile" + ) + if self.loss.temperature != 1.0 or self.loss.scale_by_temperature_squared: + raise ValueError( + "loss.mode=forward_kl_topk uses upstream logits directly; set " + "temperature=1.0 and scale_by_temperature_squared=false" + ) + if self.loss.sampled_token_nll_weight != 0.0: + raise ValueError( + "loss.mode=forward_kl_topk does not mix sampled-token NLL in the " + "supported verl v0.8 profile" + ) + if mode is TrainingMode.SFT and self.loss.sampled_token_nll_weight not in (0.0, 1.0): raise ValueError( "run.mode=sft trains with oracle cross-entropy only; " diff --git a/src/miniverl/losses/chunked.py b/src/miniverl/losses/chunked.py index 941b4c9..b1a6aee 100644 --- a/src/miniverl/losses/chunked.py +++ b/src/miniverl/losses/chunked.py @@ -45,6 +45,7 @@ "ChunkTargetProvider", "ExactTargetProvider", "BucketedTargetProvider", + "VerlTopKTargetProvider", "LossOutput", "chunked_selected_position_loss", ] @@ -137,6 +138,42 @@ def teacher_entropy(self, start: int, end: int) -> torch.Tensor: ) +@dataclass +class VerlTopKTargetProvider: + """Official verl v0.8 top-k-only teacher supervision.""" + + topk_indices: torch.Tensor + topk_log_probs: torch.Tensor + log_prob_min_clamp: float | None = None + loss_max_clamp: float | None = None + kind: str = "verl_forward_kl_topk" + diagnostics: list[dict[str, torch.Tensor]] = field(default_factory=list) + + def divergence(self, start: int, end: int, student_logits: torch.Tensor) -> torch.Tensor: + from miniverl.losses.verl_topk import verl_forward_kl_topk + + output = verl_forward_kl_topk( + student_logits, + self.topk_log_probs[start:end], + self.topk_indices[start:end], + log_prob_min_clamp=self.log_prob_min_clamp, + loss_max_clamp=self.loss_max_clamp, + ) + self.diagnostics.append( + { + "student_mass": output.student_mass.detach().to("cpu"), + "teacher_mass": output.teacher_mass.detach().to("cpu"), + "overlap_count": output.overlap_count.detach().to("cpu"), + "overlap_token_advantage": output.overlap_token_advantage.detach().to("cpu"), + } + ) + return output.loss + + def teacher_entropy(self, start: int, end: int) -> torch.Tensor: + """Entropy is undefined without the omitted tail distribution.""" + return torch.full((end - start,), float("nan"), dtype=torch.float32) + + @dataclass class LossOutput: """Result of one chunked objective evaluation.""" diff --git a/src/miniverl/losses/verl_topk.py b/src/miniverl/losses/verl_topk.py new file mode 100644 index 0000000..68ed639 --- /dev/null +++ b/src/miniverl/losses/verl_topk.py @@ -0,0 +1,99 @@ +"""Pinned verl v0.8 ``forward_kl_topk`` semantics. + +This objective intentionally ignores the probability tail. It is therefore +separate from miniVERL's native ``bucketed_topk_tail`` coarse-grained KL. The +implementation follows official verl v0.8.0 at commit +``7aed6b230776f963fa09509c10d9c3a767d1102c``; the optional conformance gate +loads that source directly and compares values and gradients. +""" + +from __future__ import annotations + +from dataclasses import dataclass + +import torch +import torch.nn.functional as F + +__all__ = ["VERL_TOPK_SCORE_IMPLEMENTATION", "VerlTopKOutput", "verl_forward_kl_topk"] + +VERL_TOPK_SCORE_IMPLEMENTATION = "verl-v0.8.0-forward-kl-topk-v1" + + +@dataclass(frozen=True) +class VerlTopKOutput: + """Per-position loss and official OPD diagnostics.""" + + loss: torch.Tensor + student_mass: torch.Tensor + teacher_mass: torch.Tensor + overlap_count: torch.Tensor + overlap_token_advantage: torch.Tensor + + +def verl_forward_kl_topk( + student_logits: torch.Tensor, + teacher_topk_log_probs: torch.Tensor, + teacher_topk_ids: torch.Tensor, + *, + log_prob_min_clamp: float | None = None, + loss_max_clamp: float | None = None, +) -> VerlTopKOutput: + """Compute the supported direct-supervision verl v0.8 top-k objective. + + Mass and student top-k identities are derived before log-prob clamps. Both + gathered student and teacher log-probabilities are then minimum-clamped, + the unnormalised teacher top-k terms are summed, and the per-token value is + clamped non-negative. The optional symmetric maximum clamp is the later + official distillation-wrapper stage. + """ + if student_logits.ndim < 2: + raise ValueError("student_logits must have shape [..., vocab]") + expected = (*student_logits.shape[:-1], teacher_topk_ids.shape[-1]) + if tuple(teacher_topk_ids.shape) != expected: + raise ValueError( + f"teacher_topk_ids shape {tuple(teacher_topk_ids.shape)} does not match {expected}" + ) + if teacher_topk_log_probs.shape != teacher_topk_ids.shape: + raise ValueError("teacher top-k IDs and log-probabilities must have identical shapes") + if teacher_topk_ids.dtype not in {torch.int32, torch.int64}: + raise ValueError("teacher_topk_ids must be an integer tensor") + if log_prob_min_clamp is not None and not torch.isfinite(torch.tensor(log_prob_min_clamp)): + raise ValueError("log_prob_min_clamp must be finite or None") + if loss_max_clamp is not None and ( + loss_max_clamp <= 0 or not torch.isfinite(torch.tensor(loss_max_clamp)) + ): + raise ValueError("loss_max_clamp must be positive and finite or None") + + student_log_probs = F.log_softmax(student_logits, dim=-1) + k = teacher_topk_ids.shape[-1] + student_topk_ids = torch.topk(student_log_probs, k=k, dim=-1).indices + student_selected = torch.gather(student_log_probs, dim=-1, index=teacher_topk_ids) + student_mass = student_selected.exp().sum(dim=-1) + teacher_mass = teacher_topk_log_probs.exp().sum(dim=-1) + + teacher_used = teacher_topk_log_probs.float() + student_used = student_selected.float() + if log_prob_min_clamp is not None: + teacher_used = teacher_used.clamp_min(log_prob_min_clamp) + student_used = student_used.clamp_min(log_prob_min_clamp) + per_teacher_token = teacher_used.exp() * (teacher_used - student_used) + loss = per_teacher_token.sum(dim=-1).clamp_min(0.0) + if loss_max_clamp is not None: + loss = loss.clamp(max=loss_max_clamp) + + overlap_mask = (teacher_topk_ids.unsqueeze(-1) == student_topk_ids.unsqueeze(-2)).any(dim=-1) + overlap_count = overlap_mask.sum(dim=-1) + advantage_sum = (-per_teacher_token * overlap_mask).sum(dim=-1) + overlap_token_advantage = advantage_sum / overlap_count.clamp_min(1) + overlap_token_advantage = torch.where( + overlap_count > 0, + overlap_token_advantage, + torch.zeros_like(overlap_token_advantage), + ) + return VerlTopKOutput( + loss=loss, + student_mass=student_mass, + teacher_mass=teacher_mass, + overlap_count=overlap_count, + overlap_token_advantage=overlap_token_advantage, + ) diff --git a/src/miniverl/schemas/cache.py b/src/miniverl/schemas/cache.py index 67b6d7c..8d359dc 100644 --- a/src/miniverl/schemas/cache.py +++ b/src/miniverl/schemas/cache.py @@ -45,6 +45,9 @@ class CacheEntryMeta(BaseModel): checksum: str selected_span_types: dict[str, int] = Field(default_factory=dict) ordered_span_types: list[str] | None = None + prompt_row_digest: str | None = None + actor_response_token_ids: list[int] | None = None + binding_checksum: str | None = None class CacheShardMeta(BaseModel): @@ -74,6 +77,7 @@ class CacheIndex(BaseModel): top_k: int = Field(ge=1) temperature: float = Field(gt=0.0) loss_mode: str + score_implementation_version: str | None = None dtype: str = "float32" entries_per_shard: int = Field(default=32, ge=1, le=4096) entries: dict[str, CacheEntryMeta] = Field(default_factory=dict) @@ -127,6 +131,8 @@ class TeacherTargetBatch: temperature: float = 1.0 top_k: int = 0 span_types: list[str] = field(default_factory=list) + prompt_row_digest: str | None = None + actor_response_token_ids: list[int] | None = None class CacheCompressionStats(BaseModel): diff --git a/src/miniverl/teachers/local.py b/src/miniverl/teachers/local.py index fc59b93..36a5327 100644 --- a/src/miniverl/teachers/local.py +++ b/src/miniverl/teachers/local.py @@ -10,7 +10,11 @@ from miniverl.config.models import LossConfig, LossMode from miniverl.errors import AlignmentError, ConfigError, TokenizerMismatchError from miniverl.losses.bucketed import bucketed_teacher_entropy, teacher_topk_targets -from miniverl.losses.chunked import BucketedTargetProvider, ExactTargetProvider +from miniverl.losses.chunked import ( + BucketedTargetProvider, + ExactTargetProvider, + VerlTopKTargetProvider, +) from miniverl.losses.exact import exact_teacher_entropy from miniverl.models.base import CausalLMBackend from miniverl.schemas.alignment import AlignmentMap @@ -96,15 +100,24 @@ def score( trajectory_id=student.trajectory_id, policy_version=student.policy_version, shape="bucketed", - provider=BucketedTargetProvider( - topk_indices=torch.zeros(0, 1, dtype=torch.long), - topk_log_probs=torch.zeros(0, 1), - tail_log_prob=empty, - divergence_name=self.loss.divergence.value, - temperature=self.loss.temperature, - scale_by_temperature_squared=self.loss.scale_by_temperature_squared, - jsd_beta=self.loss.jsd_beta, - tail_epsilon=self.loss.tail_epsilon, + provider=( + VerlTopKTargetProvider( + topk_indices=torch.zeros(0, 1, dtype=torch.long), + topk_log_probs=torch.zeros(0, 1), + log_prob_min_clamp=self.loss.log_prob_min_clamp, + loss_max_clamp=self.loss.loss_max_clamp, + ) + if self.loss.mode is LossMode.VERL_FORWARD_KL_TOPK + else BucketedTargetProvider( + topk_indices=torch.zeros(0, 1, dtype=torch.long), + topk_log_probs=torch.zeros(0, 1), + tail_log_prob=empty, + divergence_name=self.loss.divergence.value, + temperature=self.loss.temperature, + scale_by_temperature_squared=self.loss.scale_by_temperature_squared, + jsd_beta=self.loss.jsd_beta, + tail_epsilon=self.loss.tail_epsilon, + ) ), target_token_ids=target_ids, weights=weights, @@ -198,16 +211,33 @@ def score( temperature=self.loss.temperature, top_k=top_k, span_types=list(alignment.span_types), + prompt_row_digest=student.metadata.get("row_digest"), + actor_response_token_ids=[ + token_id + for token_id, generated in zip( + student.token_ids, student.model_generated_mask, strict=True + ) + if generated + ], ) - provider = BucketedTargetProvider( - topk_indices=topk_indices, - topk_log_probs=topk_log_probs, - tail_log_prob=tail_log_prob, - divergence_name=self.loss.divergence.value, - temperature=self.loss.temperature, - scale_by_temperature_squared=self.loss.scale_by_temperature_squared, - jsd_beta=self.loss.jsd_beta, - tail_epsilon=self.loss.tail_epsilon, + provider = ( + VerlTopKTargetProvider( + topk_indices=topk_indices, + topk_log_probs=topk_log_probs, + log_prob_min_clamp=self.loss.log_prob_min_clamp, + loss_max_clamp=self.loss.loss_max_clamp, + ) + if self.loss.mode is LossMode.VERL_FORWARD_KL_TOPK + else BucketedTargetProvider( + topk_indices=topk_indices, + topk_log_probs=topk_log_probs, + tail_log_prob=tail_log_prob, + divergence_name=self.loss.divergence.value, + temperature=self.loss.temperature, + scale_by_temperature_squared=self.loss.scale_by_temperature_squared, + jsd_beta=self.loss.jsd_beta, + tail_epsilon=self.loss.tail_epsilon, + ) ) covered = torch.logsumexp(topk_log_probs, dim=-1).exp() return TeacherScoreResult( @@ -237,6 +267,11 @@ def describe(self) -> dict[str, Any]: "revision": getattr(self.backend, "model_revision", None), "capabilities": self.backend.capabilities.to_dict(), "loss_mode": self.loss.mode.value, + "score_implementation_version": ( + "verl-v0.8.0-forward-kl-topk-v1" + if self.loss.mode is LossMode.VERL_FORWARD_KL_TOPK + else "miniverl-native-v1" + ), "top_k": self._effective_top_k(), "temperature": self.loss.temperature, } diff --git a/src/miniverl/training/trainer.py b/src/miniverl/training/trainer.py index 4ddba57..dfe9709 100644 --- a/src/miniverl/training/trainer.py +++ b/src/miniverl/training/trainer.py @@ -37,6 +37,7 @@ from miniverl.alignment.workflow import build_alignment_stage_plan from miniverl.cache.store import TeacherCache from miniverl.config.models import ( + LossAggregation, LossMode, MemoryStrategy, ModelBackend, @@ -1267,6 +1268,8 @@ def _collect( def _open_cache(self) -> TeacherCache: if self._cache is None: + from miniverl.losses.verl_topk import VERL_TOPK_SCORE_IMPLEMENTATION + assert self.teacher is not None top_k = ( self.student.vocab_size @@ -1284,6 +1287,11 @@ def _open_cache(self) -> TeacherCache: "top_k": top_k, "temperature": self.config.loss.temperature, "loss_mode": self.config.loss.mode.value, + "score_implementation_version": ( + VERL_TOPK_SCORE_IMPLEMENTATION + if self.config.loss.mode is LossMode.VERL_FORWARD_KL_TOPK + else "miniverl-native-v1" + ), "dtype": self.config.cache.dtype, } if (path / "index.json").is_file(): @@ -1459,7 +1467,7 @@ def _attach_persisted_offline_dataset(self) -> None: def _load_offline_dataset(self, *, expected_digest: str) -> None: from miniverl.losses.bucketed import bucketed_teacher_entropy - from miniverl.losses.chunked import BucketedTargetProvider + from miniverl.losses.chunked import BucketedTargetProvider, VerlTopKTargetProvider from miniverl.teachers.base import TeacherScoreResult from miniverl.training.offline_dataset import load_offline_dataset @@ -1514,15 +1522,24 @@ def _load_offline_dataset(self, *, expected_digest: str) -> None: raise CheckpointError( f"offline span order changed for {trajectory.trajectory_id!r}" ) - provider = BucketedTargetProvider( - topk_indices=batch.topk_indices, - topk_log_probs=batch.topk_log_probs, - tail_log_prob=batch.tail_log_prob, - divergence_name=self.config.loss.divergence.value, - temperature=self.config.loss.temperature, - scale_by_temperature_squared=(self.config.loss.scale_by_temperature_squared), - jsd_beta=self.config.loss.jsd_beta, - tail_epsilon=self.config.loss.tail_epsilon, + provider = ( + VerlTopKTargetProvider( + topk_indices=batch.topk_indices, + topk_log_probs=batch.topk_log_probs, + log_prob_min_clamp=self.config.loss.log_prob_min_clamp, + loss_max_clamp=self.config.loss.loss_max_clamp, + ) + if self.config.loss.mode is LossMode.VERL_FORWARD_KL_TOPK + else BucketedTargetProvider( + topk_indices=batch.topk_indices, + topk_log_probs=batch.topk_log_probs, + tail_log_prob=batch.tail_log_prob, + divergence_name=self.config.loss.divergence.value, + temperature=self.config.loss.temperature, + scale_by_temperature_squared=(self.config.loss.scale_by_temperature_squared), + jsd_beta=self.config.loss.jsd_beta, + tail_epsilon=self.config.loss.tail_epsilon, + ) ) teacher_score = TeacherScoreResult( trajectory_id=trajectory.trajectory_id, @@ -1698,7 +1715,13 @@ def _compute_group_gradients( divergence_available = False sampled_nll_available = False oracle_ce_available = False + verl_student_mass: list[Any] = [] + verl_teacher_mass: list[Any] = [] + verl_overlap_count: list[Any] = [] + verl_overlap_advantage: list[Any] = [] group_scale = 1.0 / max(len(group), 1) + token_mean = config.loss.aggregation is LossAggregation.TOKEN_MEAN + group_weight_total = sum(sum(sample.alignment.token_weights) for sample in group) requested_batch_size = config.train.trajectory_batch_size physical_batch_size = ( len(group) if requested_batch_size == "auto" else int(requested_batch_size) @@ -1757,12 +1780,16 @@ def _compute_group_gradients( "one padded trajectory batch cannot mix SFT and distillation targets" ) ce_weight = 1.0 if provider is None else config.loss.sampled_token_nll_weight - microbatch_scale = len(samples) * group_scale + microbatch_scale = 1.0 if token_mean else len(samples) * group_scale + effective_weights = ( + torch.cat(weight_rows) if token_mean else normalize_trajectory_weights(weight_rows) + ) + weight_normalizer = group_weight_total if token_mean else float(len(samples)) output = chunked_selected_position_loss( hidden_states=hidden, lm_head=self.student.project, - weights=normalize_trajectory_weights(weight_rows), - weight_normalizer=float(len(samples)), + weights=effective_weights, + weight_normalizer=weight_normalizer, provider=provider, target_token_ids=targets, ce_weight=ce_weight, @@ -1770,6 +1797,16 @@ def _compute_group_gradients( backward=True, loss_scale=microbatch_scale, ) + for target_provider in providers: + diagnostics = getattr(target_provider, "diagnostics", None) + if not isinstance(diagnostics, list): + continue + for values in diagnostics: + verl_student_mass.append(values["student_mass"]) + verl_teacher_mass.append(values["teacher_mass"]) + verl_overlap_count.append(values["overlap_count"]) + verl_overlap_advantage.append(values["overlap_token_advantage"]) + diagnostics.clear() loss_total += float(output.loss) * microbatch_scale positions_total += output.num_positions for sample_index, (sample, weight_tensor) in enumerate( @@ -1784,26 +1821,22 @@ def _compute_group_gradients( entry = span_losses.setdefault(name, [0.0, 0.0]) entry[0] += numerator entry[1] += denominator - tensor_denominator = weight_tensor.sum().clamp_min(1e-12) + tensor_denominator = ( + torch.tensor(group_weight_total, device=device).clamp_min(1e-12) + if token_mean + else weight_tensor.sum().clamp_min(1e-12) + ) if output.per_token_divergence is not None: divergence_available = True - divergence_total += ( - float( - ( - output.per_token_divergence[start:end].to(device) * weight_tensor - ).sum() - / tensor_denominator - ) - * group_scale - ) + divergence_total += float( + (output.per_token_divergence[start:end].to(device) * weight_tensor).sum() + / tensor_denominator + ) * (1.0 if token_mean else group_scale) if output.per_token_ce is not None: - component = ( - float( - (output.per_token_ce[start:end].to(device) * weight_tensor).sum() - / tensor_denominator - ) - * group_scale - ) + component = float( + (output.per_token_ce[start:end].to(device) * weight_tensor).sum() + / tensor_denominator + ) * (1.0 if token_mean else group_scale) if provider is None: oracle_ce_available = True oracle_ce_total += component @@ -1815,12 +1848,13 @@ def _compute_group_gradients( entropy_count += int(sample.teacher.teacher_entropy.numel()) del hidden, output - return { + result = { "loss": loss_total, "selected_positions": positions_total, "trajectories_in_step": len(group), "physical_trajectory_batches": len(batch_indices), "padded_trajectory_batch_size": physical_batch_size, + "loss_aggregation": config.loss.aggregation.value, "teacher_entropy_mean": (entropy_sum / entropy_count) if entropy_count else None, "divergence_loss": divergence_total if divergence_available else None, "sampled_token_nll_loss": (sampled_nll_total if sampled_nll_available else None), @@ -1830,6 +1864,27 @@ def _compute_group_gradients( for name, (numerator, denominator) in sorted(span_losses.items()) }, } + if verl_student_mass: + student_mass = torch.cat(verl_student_mass).float() + teacher_mass = torch.cat(verl_teacher_mass).float() + overlap_count = torch.cat(verl_overlap_count).float() + overlap_advantage = torch.cat(verl_overlap_advantage).float() + overlap_positions = overlap_count > 0 + result["verl_forward_kl_topk"] = { + "student_mass_mean": float(student_mass.mean()), + "student_mass_min": float(student_mass.min()), + "student_mass_max": float(student_mass.max()), + "teacher_mass_mean": float(teacher_mass.mean()), + "teacher_mass_min": float(teacher_mass.min()), + "teacher_mass_max": float(teacher_mass.max()), + "overlap_ratio": float(overlap_count.mean() / config.loss.top_k), + "overlap_token_advantage": ( + float(overlap_advantage[overlap_positions].mean()) + if bool(overlap_positions.any()) + else 0.0 + ), + } + return result def _commit_update(self) -> dict[str, float]: """Non-retryable optimizer commit; ``step`` is invoked at most once.""" @@ -2493,7 +2548,7 @@ def _write_token_analysis(self, samples: list[TrainSample]) -> int: def _reload_targets_from_cache(self, samples: list[TrainSample]) -> list[TrainSample]: """Re-attach providers from the on-disk cache after the teacher is gone.""" - from miniverl.losses.chunked import BucketedTargetProvider + from miniverl.losses.chunked import BucketedTargetProvider, VerlTopKTargetProvider cache = self._open_cache() config = self.config @@ -2505,17 +2560,36 @@ def _reload_targets_from_cache(self, samples: list[TrainSample]) -> list[TrainSa expect_policy_version=( sample.trajectory.policy_version if config.cache.strict_policy_version else None ), + expect_prompt_row_digest=sample.trajectory.metadata.get("row_digest"), + expect_actor_response_token_ids=[ + token_id + for token_id, generated in zip( + sample.trajectory.token_ids, + sample.trajectory.model_generated_mask, + strict=True, + ) + if generated + ], device=self.plan.device, ) - sample.teacher.provider = BucketedTargetProvider( - topk_indices=batch.topk_indices, - topk_log_probs=batch.topk_log_probs, - tail_log_prob=batch.tail_log_prob, - divergence_name=config.loss.divergence.value, - temperature=config.loss.temperature, - scale_by_temperature_squared=config.loss.scale_by_temperature_squared, - jsd_beta=config.loss.jsd_beta, - tail_epsilon=config.loss.tail_epsilon, + sample.teacher.provider = ( + VerlTopKTargetProvider( + topk_indices=batch.topk_indices, + topk_log_probs=batch.topk_log_probs, + log_prob_min_clamp=config.loss.log_prob_min_clamp, + loss_max_clamp=config.loss.loss_max_clamp, + ) + if config.loss.mode is LossMode.VERL_FORWARD_KL_TOPK + else BucketedTargetProvider( + topk_indices=batch.topk_indices, + topk_log_probs=batch.topk_log_probs, + tail_log_prob=batch.tail_log_prob, + divergence_name=config.loss.divergence.value, + temperature=config.loss.temperature, + scale_by_temperature_squared=config.loss.scale_by_temperature_squared, + jsd_beta=config.loss.jsd_beta, + tail_epsilon=config.loss.tail_epsilon, + ) ) return samples diff --git a/tests/conformance/test_verl_v08_loss.py b/tests/conformance/test_verl_v08_loss.py new file mode 100644 index 0000000..8597207 --- /dev/null +++ b/tests/conformance/test_verl_v08_loss.py @@ -0,0 +1,141 @@ +"""Numerical conformance against the exact installed official verl source.""" + +from __future__ import annotations + +import importlib.metadata +import importlib.util +import json +import subprocess +import sys +import types +from pathlib import Path +from types import SimpleNamespace +from urllib.parse import unquote, urlparse + +import pytest +import torch + +from miniverl.bridge.opd_v08 import VERL_COMMIT +from miniverl.losses.verl_topk import verl_forward_kl_topk + +pytestmark = [pytest.mark.torch, pytest.mark.verl_conformance] + + +def _official_module(): # type: ignore[no-untyped-def] + try: + distribution = importlib.metadata.distribution("verl") + except importlib.metadata.PackageNotFoundError: + pytest.skip("official verl v0.8.0 is not installed") + direct_url_text = distribution.read_text("direct_url.json") + assert direct_url_text is not None, "official verl install has no VCS provenance" + direct_url = json.loads(direct_url_text) + vcs_info = direct_url.get("vcs_info") + if vcs_info is not None: + assert vcs_info["commit_id"] == VERL_COMMIT + else: + parsed = urlparse(direct_url["url"]) + checkout = Path(unquote(parsed.path.lstrip("/"))) + completed = subprocess.run( + ["git", "rev-parse", "HEAD"], + cwd=checkout, + check=True, + capture_output=True, + text=True, + ) + assert completed.stdout.strip() == VERL_COMMIT + root = Path(distribution.locate_file("")) + source = root / "verl" / "trainer" / "distillation" / "fsdp" / "losses.py" + assert source.is_file(), source + + # Load the official file directly so the conformance environment does not + # need Ray or the distributed worker stack imported by verl.__init__. + ulysses = types.ModuleType("verl.utils.ulysses") + ulysses.get_ulysses_sequence_parallel_world_size = lambda: 1 + ulysses.slice_input_tensor = lambda value, dim: value + config = types.ModuleType("verl.workers.config") + config.DistillationConfig = type("DistillationConfig", (), {}) + config.DistillationLossConfig = type("DistillationLossConfig", (), {}) + saved = { + name: sys.modules.get(name) + for name in ( + "verl", + "verl.utils", + "verl.utils.ulysses", + "verl.workers", + "verl.workers.config", + ) + } + sys.modules["verl"] = types.ModuleType("verl") + sys.modules["verl.utils"] = types.ModuleType("verl.utils") + sys.modules["verl.utils.ulysses"] = ulysses + sys.modules["verl.workers"] = types.ModuleType("verl.workers") + sys.modules["verl.workers.config"] = config + try: + spec = importlib.util.spec_from_file_location("_official_verl_v08_fsdp_losses", source) + assert spec is not None and spec.loader is not None + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + finally: + for name, previous in saved.items(): + if previous is None: + sys.modules.pop(name, None) + else: + sys.modules[name] = previous + + +def _nested(rows: torch.Tensor) -> torch.Tensor: + offsets = torch.tensor([0, rows.shape[0]], dtype=torch.int64) + return torch.nested.nested_tensor_from_jagged(rows, offsets=offsets) + + +@pytest.mark.parametrize("minimum", [None, -10.0]) +def test_values_diagnostics_reduction_and_gradient_match_official_verl(minimum) -> None: + official = _official_module() + logits_data = torch.tensor( + [ + [0.2, -0.4, 1.3, 0.1, 0.7], + [-0.8, 1.1, 0.4, 0.2, 0.9], + [1.4, 0.3, -0.6, 0.8, 0.0], + ], + dtype=torch.float32, + ) + teacher_ids = torch.tensor([[2, 4], [1, 4], [0, 3]], dtype=torch.int64) + teacher_log_probs = torch.log( + torch.tensor([[0.62, 0.21], [0.58, 0.25], [0.54, 0.30]], dtype=torch.float32) + ) + + official_logits = logits_data.clone().unsqueeze(0).requires_grad_(True) + official_output = official.compute_forward_kl_topk( + student_logits=official_logits, + teacher_topk_log_probs=_nested(teacher_log_probs), + teacher_topk_ids=_nested(teacher_ids), + config=SimpleNamespace(distillation_loss=SimpleNamespace(log_prob_min_clamp=minimum)), + data_format="thd", + ) + official_loss = official_output["distillation_losses"].clamp_min(0.0) + official_scalar = official_loss.mean() + official_scalar.backward() + + local_logits = logits_data.clone().requires_grad_(True) + local = verl_forward_kl_topk( + local_logits, + teacher_log_probs, + teacher_ids, + log_prob_min_clamp=minimum, + ) + local_scalar = local.loss.mean() + local_scalar.backward() + + torch.testing.assert_close(local.loss, official_loss.squeeze(0), rtol=1e-6, atol=1e-7) + torch.testing.assert_close(local.student_mass, official_output["student_mass"].squeeze(0)) + torch.testing.assert_close(local.teacher_mass, official_output["teacher_mass"].squeeze(0)) + torch.testing.assert_close(local.overlap_count, official_output["overlap_count"].squeeze(0)) + torch.testing.assert_close( + local.overlap_token_advantage, + official_output["overlap_token_advantage"].squeeze(0), + ) + torch.testing.assert_close(local_scalar, official_scalar) + torch.testing.assert_close( + local_logits.grad, official_logits.grad.squeeze(0), rtol=1e-6, atol=1e-7 + ) diff --git a/tests/integration/test_prompt_opd_pipeline.py b/tests/integration/test_prompt_opd_pipeline.py index 6941dc7..ddadaf6 100644 --- a/tests/integration/test_prompt_opd_pipeline.py +++ b/tests/integration/test_prompt_opd_pipeline.py @@ -89,9 +89,13 @@ def test_prompt_opd_trains_without_an_environment_or_reward(tmp_path) -> None: }, "selection": {"selector": "all_model_tokens"}, "loss": { - "mode": "bucketed_topk_tail", + "mode": "forward_kl_topk", "divergence": "forward_kl", + "aggregation": "token-mean", + "temperature": 1.0, + "scale_by_temperature_squared": False, "top_k": 4, + "log_prob_min_clamp": -10.0, "chunk_size": 16, }, "train": { @@ -128,3 +132,19 @@ def test_prompt_opd_trains_without_an_environment_or_reward(tmp_path) -> None: for row in rows ) assert all(row["metadata"]["reward_model"] is None for row in rows) + metrics = [ + json.loads(line) + for line in (result.run_dir / "metrics.jsonl").read_text(encoding="utf-8").splitlines() + ] + update = next(row for row in metrics if row.get("phase") == "opd") + assert update["loss_aggregation"] == "token-mean" + assert set(update["verl_forward_kl_topk"]) == { + "student_mass_mean", + "student_mass_min", + "student_mass_max", + "teacher_mass_mean", + "teacher_mass_min", + "teacher_mass_max", + "overlap_ratio", + "overlap_token_advantage", + } diff --git a/tests/unit/test_cache.py b/tests/unit/test_cache.py index e6f7847..5a3133b 100644 --- a/tests/unit/test_cache.py +++ b/tests/unit/test_cache.py @@ -86,6 +86,54 @@ def test_float32_round_trip_is_exact(tmp_path: Path): assert loaded.span_types == batch.span_types +def test_prompt_target_binding_rejects_a_different_actor_response(tmp_path: Path): + cache = TeacherCache.create( + tmp_path / "tc", + miniverl_version="0.8.0.dev0", + teacher_model_id="toy-teacher", + teacher_model_revision="rev-abc", + tokenizer_fingerprint="fp-1234", + vocab_size=VOCAB, + top_k=TOP_K, + temperature=1.0, + loss_mode="forward_kl_topk", + score_implementation_version="verl-v0.8.0-forward-kl-topk-v1", + ) + batch = _batch("prompt:0", policy_version=7) + batch.prompt_row_digest = "a" * 64 + batch.actor_response_token_ids = [11, 12, 13] + cache.write(batch, selector="all_model_tokens") + cache.flush() + + loaded = cache.read( + "prompt:0", + expect_policy_version=7, + expect_prompt_row_digest="a" * 64, + expect_actor_response_token_ids=[11, 12, 13], + ) + assert loaded.actor_response_token_ids == [11, 12, 13] + with pytest.raises(StaleCacheError, match="exact actor response token IDs"): + cache.read("prompt:0", expect_actor_response_token_ids=[11, 99, 13]) + + +def test_score_implementation_version_is_part_of_cache_identity(tmp_path: Path): + cache = _cache(tmp_path / "tc") + with pytest.raises(StaleCacheError, match="score_implementation_version"): + cache.assert_compatible( + teacher_model_id="toy-teacher", + teacher_model_revision="rev-abc", + tokenizer_fingerprint="fp-1234", + tokenizer_identity={}, + teacher_adapter_provenance=None, + vocab_size=VOCAB, + top_k=TOP_K, + temperature=1.0, + loss_mode="bucketed_topk_tail", + dtype="float32", + score_implementation_version="miniverl-native-v1", + ) + + def test_float16_round_trip_stays_within_its_documented_precision(tmp_path: Path): cache = _cache(tmp_path / "tc", dtype="float16") batch = _batch("t0") diff --git a/tests/unit/test_prompt_source_config.py b/tests/unit/test_prompt_source_config.py index f28759d..bfc60c9 100644 --- a/tests/unit/test_prompt_source_config.py +++ b/tests/unit/test_prompt_source_config.py @@ -74,3 +74,26 @@ def test_plain_string_prompts_require_an_explicit_opt_in() -> None: assert config.source.allow_plain_string_prompts is True assert config.source.truncation.value == "left" + + +def test_verl_forward_kl_topk_requires_its_exact_supported_contract() -> None: + base = { + "models": _models(), + "environment": {"name": "calculator"}, + "loss": { + "mode": "forward_kl_topk", + "divergence": "forward_kl", + "aggregation": "token-mean", + "temperature": 1.0, + "scale_by_temperature_squared": False, + "top_k": 8, + "log_prob_min_clamp": -10.0, + }, + } + config = RunConfig.model_validate(base) + assert config.loss.mode.value == "forward_kl_topk" + assert config.loss.aggregation.value == "token-mean" + + incompatible = {**base, "loss": {**base["loss"], "aggregation": "native_per_trajectory"}} + with pytest.raises(ValidationError, match=r"requires loss\.aggregation=token-mean"): + RunConfig.model_validate(incompatible) diff --git a/tests/unit/test_verl_forward_kl_topk.py b/tests/unit/test_verl_forward_kl_topk.py new file mode 100644 index 0000000..57f2837 --- /dev/null +++ b/tests/unit/test_verl_forward_kl_topk.py @@ -0,0 +1,70 @@ +from __future__ import annotations + +import torch + +from miniverl.losses.verl_topk import verl_forward_kl_topk + + +def test_forward_kl_topk_matches_pinned_formula_and_gradient() -> None: + student = torch.tensor( + [[0.2, -0.1, 1.4, 0.7], [-1.0, 0.5, 0.2, 1.2]], + dtype=torch.float64, + requires_grad=True, + ) + teacher_ids = torch.tensor([[2, 3], [3, 1]]) + teacher_log_probs = torch.log(torch.tensor([[0.60, 0.25], [0.55, 0.30]], dtype=torch.float64)) + + output = verl_forward_kl_topk(student, teacher_log_probs, teacher_ids) + reference_student = torch.log_softmax(student, dim=-1).gather(-1, teacher_ids) + expected = ( + (teacher_log_probs.float().exp() * (teacher_log_probs.float() - reference_student.float())) + .sum(dim=-1) + .clamp_min(0.0) + ) + expected.mean().backward() + expected_grad = student.grad.detach().clone() + student.grad = None + output.loss.mean().backward() + + torch.testing.assert_close(output.loss, expected) + torch.testing.assert_close(student.grad, expected_grad) + torch.testing.assert_close(output.teacher_mass, teacher_log_probs.exp().sum(dim=-1)) + torch.testing.assert_close(output.student_mass, reference_student.exp().sum(dim=-1)) + + +def test_clamps_follow_pinned_order_and_overlap_uses_clamped_terms() -> None: + student = torch.tensor([[9.0, 8.0, -20.0, -30.0]]) + teacher_ids = torch.tensor([[2, 0]]) + teacher_log_probs = torch.tensor([[-30.0, -0.2]]) + + output = verl_forward_kl_topk( + student, + teacher_log_probs, + teacher_ids, + log_prob_min_clamp=-10.0, + loss_max_clamp=0.5, + ) + + unclamped_student = torch.log_softmax(student, dim=-1).gather(-1, teacher_ids) + expected_terms = teacher_log_probs.clamp_min(-10).exp() * ( + teacher_log_probs.clamp_min(-10) - unclamped_student.clamp_min(-10) + ) + expected = expected_terms.sum(dim=-1).clamp_min(0).clamp(max=0.5) + torch.testing.assert_close(output.loss, expected) + assert output.overlap_count.tolist() == [1] + torch.testing.assert_close(output.overlap_token_advantage, -expected_terms[:, 1]) + # Mass is measured before stability clamps in official verl v0.8.0. + torch.testing.assert_close(output.teacher_mass, teacher_log_probs.exp().sum(dim=-1)) + torch.testing.assert_close(output.student_mass, unclamped_student.exp().sum(dim=-1)) + + +def test_topk_objective_remains_distinct_from_tail_bucket_kl() -> None: + student = torch.tensor([[0.1, 0.3, -0.4, 1.2]]) + teacher_ids = torch.tensor([[3, 1]]) + teacher_log_probs = torch.log(torch.tensor([[0.50, 0.20]])) + + output = verl_forward_kl_topk(student, teacher_log_probs, teacher_ids) + + assert output.loss.item() >= 0.0 + assert output.teacher_mass.item() == torch.tensor(0.7).item() + assert output.teacher_mass.item() < 1.0 From 2c48d79bbbc9eaf00f99cd2473a8b76c0be1ed4a Mon Sep 17 00:00:00 2001 From: Daoyuan Li <94409450+DaoyuanLi2816@users.noreply.github.com> Date: Wed, 12 Aug 2026 00:16:04 -0700 Subject: [PATCH 2/3] Keep conformance tests collection-safe --- .github/workflows/verl-bridge.yml | 4 ++-- tests/conformance/test_verl_v08_loss.py | 5 ++++- tests/unit/test_verl_forward_kl_topk.py | 8 +++++++- 3 files changed, 13 insertions(+), 4 deletions(-) diff --git a/.github/workflows/verl-bridge.yml b/.github/workflows/verl-bridge.yml index 5e0ab97..c40c7e7 100644 --- a/.github/workflows/verl-bridge.yml +++ b/.github/workflows/verl-bridge.yml @@ -68,13 +68,13 @@ jobs: - name: Install miniVERL bridge environment run: >- python -m pip install --upgrade pip && - python -m pip install ".[bridge,train]" "hydra-core>=1.3,<2" + python -m pip install ".[bridge,train]" "hydra-core>=1.3,<2" pytest - name: Install the exact official verl source without its distributed stack run: >- python -m pip install --no-deps "git+https://github.com/verl-project/verl.git@7aed6b230776f963fa09509c10d9c3a767d1102c" - name: Compare forward_kl_topk values, diagnostics and gradients - run: pytest -q tests/conformance/test_verl_v08_loss.py -m verl_conformance + run: python -m pytest -q tests/conformance/test_verl_v08_loss.py -m verl_conformance - name: Generate and export the standards smoke bundle run: | python scripts/prepare_verl_bridge_smoke.py --out _verl-smoke-source diff --git a/tests/conformance/test_verl_v08_loss.py b/tests/conformance/test_verl_v08_loss.py index 8597207..bcfc676 100644 --- a/tests/conformance/test_verl_v08_loss.py +++ b/tests/conformance/test_verl_v08_loss.py @@ -2,6 +2,8 @@ from __future__ import annotations +# ruff: noqa: E402 - skip collection before importing torch-backed miniVERL code + import importlib.metadata import importlib.util import json @@ -13,7 +15,8 @@ from urllib.parse import unquote, urlparse import pytest -import torch + +torch = pytest.importorskip("torch") from miniverl.bridge.opd_v08 import VERL_COMMIT from miniverl.losses.verl_topk import verl_forward_kl_topk diff --git a/tests/unit/test_verl_forward_kl_topk.py b/tests/unit/test_verl_forward_kl_topk.py index 57f2837..60352b9 100644 --- a/tests/unit/test_verl_forward_kl_topk.py +++ b/tests/unit/test_verl_forward_kl_topk.py @@ -1,6 +1,12 @@ from __future__ import annotations -import torch +# ruff: noqa: E402 - skip collection before importing torch-backed miniVERL code + +import pytest + +torch = pytest.importorskip("torch") + +pytestmark = pytest.mark.torch from miniverl.losses.verl_topk import verl_forward_kl_topk From ffbeb973d5838e4f31de150e4fb61460c886a3d4 Mon Sep 17 00:00:00 2001 From: Daoyuan Li <94409450+DaoyuanLi2816@users.noreply.github.com> Date: Wed, 12 Aug 2026 00:16:30 -0700 Subject: [PATCH 3/3] Satisfy conformance test lint --- tests/conformance/test_verl_v08_loss.py | 4 ++-- tests/unit/test_verl_forward_kl_topk.py | 4 ++-- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/tests/conformance/test_verl_v08_loss.py b/tests/conformance/test_verl_v08_loss.py index bcfc676..d195053 100644 --- a/tests/conformance/test_verl_v08_loss.py +++ b/tests/conformance/test_verl_v08_loss.py @@ -1,9 +1,9 @@ +# ruff: noqa: E402 - skip collection before importing torch-backed miniVERL code + """Numerical conformance against the exact installed official verl source.""" from __future__ import annotations -# ruff: noqa: E402 - skip collection before importing torch-backed miniVERL code - import importlib.metadata import importlib.util import json diff --git a/tests/unit/test_verl_forward_kl_topk.py b/tests/unit/test_verl_forward_kl_topk.py index 60352b9..873a300 100644 --- a/tests/unit/test_verl_forward_kl_topk.py +++ b/tests/unit/test_verl_forward_kl_topk.py @@ -1,7 +1,7 @@ -from __future__ import annotations - # ruff: noqa: E402 - skip collection before importing torch-backed miniVERL code +from __future__ import annotations + import pytest torch = pytest.importorskip("torch")