[WS2] Ablation Matrix API v2.0 on top of PR230 - #288
Conversation
Rollout and training compute logprobs for the same tokens with the same weights and still disagree. This turns "which of the dozens of possible causes is it" into switches that can be flipped one at a time and attributed to a side. Ships the framework only. Operators are claimed and written separately, so operator_checks/ is empty by design and adding one changes nothing outside its own directory. Layout, in dependency order: schema/ pure data types, no behaviour pipeline/ the seven execution steps, free functions only engines/ the two sides under test reference_adapters/ wiring reference implementations in model_meta/ model shape and module correspondence operator_checks/ plugins, one directory per operator (empty) Three design decisions worth stating: Four variants, not two. A single swap cannot attribute a side: only a one-sided swap says which side is at fault, and only the two-sided swap proves the reference itself is sound. That last arm is the self-check gate -- without it one wrong reference quietly steers every attribution, which is worse than having no framework at all. Four gates before the matrix. "Not measured" and "measured and clean" are different, and confusing them is the mistake this kind of framework is most likely to make. A silently reverted switch, missing evidence, an incomplete set of logprob shards, or a failed pitfall guard each block a verdict rather than passing through as a clean result. Convergence is judged on clip_fraction, not dlogp_mean. At every production floor the mean sits far below the GRPO clip edge, so judging on it would mark almost every factor NOT_THIS_FACTOR while gradient signal is being discarded in the tail. Thresholds are code constants keyed by (model family, noise floor), not configuration: a tunable threshold is one somebody can tune until the test passes, and it has to enter the execution fingerprint so changing it invalidates historical results. 39 framework tests, all on CPU via a synthetic scoring backend that can reproduce the failure modes above. Nothing here claims anything about real Megatron or vLLM numerics.
The framework shipped without operators. This adds the three interfaces that show the three shapes a factor can take, so each can be claimed and implemented independently: attention/rope_fusion implementation swap against a SHARED_BACKEND gemm/forward_reduce collective communication, SELF_WRITTEN reference logprob/precision_downcast parameter sweep, no reference Every adapter method raises NotImplementedError; the declaration layer works, so `list` and `plan` verify a factor is wired before anything is implemented. engines/ gets megatron.py and vllm.py placeholders whose docstrings carry the settings that must be pinned and the readback path for each. Since engines/ is for the two sides under test, the CPU scoring harness moves out to tests/mismatch_cpu_backend.py -- satisfying ScoringBackend does not make something a side under test. Drops MismatchFactor.owner and .tracked_by: ownership belongs in the issue tracker, not in every factor declaration. Adds README.md and two tutorials covering how to add a kernel factor and how to add a communication feature. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01QjJLV6aQZTtm4X6pjex8Ri
Applied Clean Code chapter 4 across the package. The docstrings had grown into design documents: rationale essays on types, restatements of the signature, and system-wide explanation attached to one local declaration -- the "too much information" and "nonlocal information" smells. That material belongs in README.md and docs/, where it already is. Removed roughly 530 lines. What survives is what the code cannot say for itself: why a field is RECORD_ONLY, why a silent fallback is more dangerous than an error, why thresholds are constants rather than configuration. One executable change: declared_collectives() drops an intermediate variable that only repeated the function name. Everything else is comments -- verified by comparing every module's AST with docstrings stripped. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01QjJLV6aQZTtm4X6pjex8Ri
The README, both tutorials and the PR description each restated the four arms, the gates and the noise floors. Duplicated prose goes stale in the copies nobody edits, so each fact now lives in exactly one place: README holds the concepts, the tutorials hold the steps, and the comm tutorial covers only what differs from the kernel one. Also dropped the code that the repository already shows. A tutorial that pastes an entire factor declaration is a second copy to keep in sync; pointing at operator_checks/attention/factors/rope_fusion.py is not. 1176 lines to 629. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01QjJLV6aQZTtm4X6pjex8Ri
📝 WalkthroughWalkthroughAdded a mismatch-diagnosis framework with immutable schemas, planning and scoring pipelines, operator plugins, reference adapters, CLI commands, model metadata, reporting, documentation, and comprehensive tests. ChangesMismatch diagnosis framework
Estimated code review effort: 5 (Critical) | ~120 minutes Sequence Diagram(s)sequenceDiagram
participant CLI
participant PluginRegistry
participant Planner
participant Runner
participant Diagnosis
participant Report
CLI->>PluginRegistry: load operator plugins
PluginRegistry-->>CLI: return registered factors
CLI->>Planner: request runnable variants
Planner->>Runner: provide ordered variants
Runner->>Diagnosis: provide results and evidence
Diagnosis->>Report: build mismatch report
Possibly related PRs
Suggested reviewers: 🚥 Pre-merge checks | ✅ 4 | ❌ 1❌ Failed checks (1 warning)
✅ Passed checks (4 passed)
✨ Finishing Touches🧪 Generate unit tests (beta)
Thanks for using CodeRabbit! It's free for OSS, and your support helps us grow. If you like it, consider giving us a shout-out. Comment |
There was a problem hiding this comment.
Actionable comments posted: 20
🧹 Nitpick comments (9)
rl_engine/mismatch/__main__.py (1)
52-63: 🎯 Functional Correctness | 🔵 Trivial | ⚡ Quick winValidate
--operatoragainst the registered operators.Both commands accept any string. If a user mistypes an operator name,
listprints nothing and returns 0, andplanreports "nothing to plan". The exit status stays 0, so a typo looks like an empty framework.Resolve the operator name after the plugins load and fail with a clear message.
♻️ Proposed fix
def command_list(operator: str | None) -> int: operators = load_operator_plugins() + if operator is not None and operator not in operators: + print(f"unknown operator {operator!r}; registered: {', '.join(operators) or 'none'}") + return 2 if not operators:Apply the same check in
command_plan.🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the rest with a brief reason, keep changes minimal, and validate. In `@rl_engine/mismatch/__main__.py` around lines 52 - 63, Validate the optional operator argument in both command_list and command_plan after plugins are loaded, resolving it against the registered operators. If the name is provided but unrecognized, emit a clear error and return a nonzero exit status instead of treating it as an empty result; preserve existing behavior when omitted or valid.tests/test_mismatch_framework.py (2)
249-252: 📐 Maintainability & Code Quality | 🔵 Trivial | ⚡ Quick win
make_factor(reference=None)does not produce a reference-free factor.Line 141 of the helper replaces
Nonewithmake_reference(). Thereference=Noneargument on line 249 therefore has no effect, and line 250 does the real work. If a later change removes line 250, this test silently becomes a swap test that still passes the name assertion only by accident.Let the helper express "no reference" with a sentinel default, then drop the
__dict__rebuild.♻️ Proposed fix
+_UNSET = object() + def make_factor( factor_id: str = "fixture.swap", *, - reference: ReferenceImplementation | None = None, + reference: ReferenceImplementation | None = _UNSET, # type: ignore[assignment] ... - reference=reference if reference is not None else make_reference(), + reference=make_reference() if reference is _UNSET else reference, )- factor = make_factor(reference=None, allowed_values=(1, 2, 4)) - factor = MismatchFactor(**{**factor.__dict__, "reference": None}) + factor = make_factor(reference=None, allowed_values=(1, 2, 4)) variants = build_variants(factor)🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the rest with a brief reason, keep changes minimal, and validate. In `@tests/test_mismatch_framework.py` around lines 249 - 252, Update make_factor to use a sentinel default that distinguishes an omitted reference from an explicit reference=None, while preserving automatic reference creation for omitted arguments. In the test around build_variants, remove the MismatchFactor __dict__ reconstruction and rely directly on make_factor(reference=None) to create the reference-free factor.
462-471: 📐 Maintainability & Code Quality | 🔵 Trivial | 💤 Low valuePatch the arm by name, not by index.
four_arms()returns results in the insertion order of a local dict. Line 463 replaces index 2 and assumes it istraining_reference_only. Line 490 makes the same assumption for index 0. If the helper gains an arm or reorders one, these tests still pass while measuring a different arm.Select the entry by
variant.nameinstead.♻️ Proposed fix
results = four_arms() - results[2] = make_result( + index = next(i for i, r in enumerate(results) if r.variant.name == "training_reference_only") + results[index] = make_result( "training_reference_only",🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the rest with a brief reason, keep changes minimal, and validate. In `@tests/test_mismatch_framework.py` around lines 462 - 471, Update the test setup around four_arms() to locate and replace the result whose variant.name matches the intended arm, rather than assigning by numeric index. Apply the same name-based selection to both replacements currently using indices 2 and 0, preserving the existing make_result values for training_reference_only and the other targeted arm.rl_engine/mismatch/docs/add-a-kernel-factor.md (1)
23-29: 📐 Maintainability & Code Quality | 🔵 Trivial | 💤 Low valueFenced code blocks in the new Markdown files omit a language. markdownlint reports MD040 at four places across two files. Add a language tag such as
textto each plain block.
rl_engine/mismatch/docs/add-a-kernel-factor.md#L23-L29: tag the reference-authority block at line 23, the directory tree at line 42, and the skipped-prerequisites output at line 165 withtext.rl_engine/mismatch/README.md#L72-L82: tag the package layout tree at line 72 withtext.🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the rest with a brief reason, keep changes minimal, and validate. In `@rl_engine/mismatch/docs/add-a-kernel-factor.md` around lines 23 - 29, Add the text language tag to all four plain fenced code blocks: the reference-authority block at rl_engine/mismatch/docs/add-a-kernel-factor.md lines 23-29, the directory tree at lines 42-47, the skipped-prerequisites output at lines 165-171, and the package layout tree at rl_engine/mismatch/README.md lines 72-82. No other content changes are needed.Source: Linters/SAST tools
rl_engine/mismatch/pipeline/registry.py (3)
178-185: 📐 Maintainability & Code Quality | 🔵 Trivial | 💤 Low valueSort
__all__to satisfy RUF022.Ruff reports that
__all__is not sorted in isort style. MoveOPERATOR_CHECKSbefore the class names.🧹 Proposed ordering
__all__ = [ - "FactorDiscoveryError", "OPERATOR_CHECKS", + "FactorDiscoveryError", "OperatorChecks", "PluginRegistry", "RegistrationError", "discover_factors", ]🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the rest with a brief reason, keep changes minimal, and validate. In `@rl_engine/mismatch/pipeline/registry.py` around lines 178 - 185, Sort the `__all__` entries in isort style by moving `OPERATOR_CHECKS` before `FactorDiscoveryError`, while preserving all existing exports.Source: Linters/SAST tools
167-172: 📐 Maintainability & Code Quality | 🔵 Trivial | 💤 Low valueHandle factor ids that contain more than one dot.
factor.id.split(".", 1)[-1]keeps every dot after the first one. For an id such aslogprob.precision.downcast,expected_suffixbecomesprecision.downcast, which no module name can equal. Discovery then fails with a message that asks for a file name containing a dot.Use the last path segment, or state the single-dot convention in the docstring and reject ids with more than one dot at registration.
♻️ Proposed suffix handling
- expected_suffix = factor.id.split(".", 1)[-1] + expected_suffix = factor.id.rsplit(".", 1)[-1]🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the rest with a brief reason, keep changes minimal, and validate. In `@rl_engine/mismatch/pipeline/registry.py` around lines 167 - 172, Update the factor filename validation around expected_suffix to use only the final dot-separated segment of factor.id, so ids such as logprob.precision.downcast expect downcast.py. Preserve the existing mismatch error and factor discovery flow.
91-115: 🚀 Performance & Scalability | 🔵 Trivial | ⚡ Quick winCache each plugin's declared factors at registration.
_check_factor_conflictscallsdeclare_factors()three times for every already-registered plugin, andfactors_forcalls it again on each query. For the plugins in this PR,declare_factors()runsdiscover_factors, which performs apkgutil.iter_modulesscan, repeatedimportlib.import_modulelookups, and a sort on every call. Registration cost grows quadratically with the plugin count, and the same work repeats at query time.Store the tuple once at registration and reuse it.
♻️ Proposed caching of declared factors
def __init__(self) -> None: self._plugins: dict[str, OperatorChecks] = {} + self._factors: dict[str, tuple[MismatchFactor, ...]] = {} def register(self, plugin_cls: type) -> type: """Instantiate and register a plugin, checking for conflicts.""" plugin = plugin_cls() name = getattr(plugin, "operator", "") if not name: raise RegistrationError(f"{plugin_cls.__name__} does not declare an operator name") if name in self._plugins: raise RegistrationError(f"operator {name!r} is already registered") - self._check_factor_conflicts(plugin) + declared = tuple(plugin.declare_factors()) + self._check_factor_conflicts(plugin.operator, declared) self._plugins[name] = plugin + self._factors[name] = declared return plugin_clsThen read
self._factorsinside_check_factor_conflicts,factors_for, and clear it inclear().🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the rest with a brief reason, keep changes minimal, and validate. In `@rl_engine/mismatch/pipeline/registry.py` around lines 91 - 115, Cache each plugin’s declared factors once during registration as a tuple in self._factors, then reuse that cache in _check_factor_conflicts and factors_for instead of calling declare_factors repeatedly. Ensure newly registered plugins are added to the cache only after conflict validation succeeds, and clear the cached factors in clear() alongside the existing registry state.rl_engine/mismatch/pipeline/comparison.py (2)
145-155: 🩺 Stability & Availability | 🔵 Trivial | ⚡ Quick win
strengthraisesKeyErrorfor a newDeterminismLevelmember.The map lists four members. If someone adds a fifth member to
DeterminismLevel,strength[left.determinism]raisesKeyErrorinside comparison, and the whole run fails rather than reporting an issue.The mapped integers are also only used for inequality, so the ordering they encode is never read. Compare the members directly, or attach the ordering to the enum so that a new member cannot be omitted.
♻️ Proposed simplification
- strength = { - DeterminismLevel.NONE: 0, - DeterminismLevel.STABLE_WITHIN_PROCESS: 1, - DeterminismLevel.STABLE_ACROSS_RUNS: 2, - DeterminismLevel.STABLE_ACROSS_TOPOLOGY: 3, - } issues: list[ComparisonIssue] = [] # strict=False: the two sides may declare different numbers of collectives. paired = zip(rollout.collectives, training.collectives, strict=False) for index, (left, right) in enumerate(paired): - if strength[left.determinism] != strength[right.determinism]: + if left.determinism is not right.determinism:🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the rest with a brief reason, keep changes minimal, and validate. In `@rl_engine/mismatch/pipeline/comparison.py` around lines 145 - 155, Remove the local strength mapping from the comparison loop and compare left.determinism and right.determinism directly for inequality. Update the condition in the comparison function while preserving the existing issue-reporting behavior for differing determinism levels, so newly added DeterminismLevel members cannot cause a KeyError.
152-154: 🎯 Functional Correctness | 🔵 Trivial | ⚡ Quick winA collective-count difference is never reported.
strict=Falsedrops the unpaired tail. If the rollout side declares two collectives and the training side declares one, the second collective is not compared and no issue is emitted. A missing collective on one side is a diagnostically relevant difference for this framework.Emit a
RECORD_ONLY-style issue, or at minimum record the count difference in the report.🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the rest with a brief reason, keep changes minimal, and validate. In `@rl_engine/mismatch/pipeline/comparison.py` around lines 152 - 154, Update the collective comparison loop around paired and the rollout.collectives/training.collectives lists to detect unequal lengths instead of silently dropping the unpaired tail. Emit the existing RECORD_ONLY-style issue for each missing collective, or otherwise record the count difference in the comparison report, while preserving comparisons for paired collectives.
🤖 Prompt for all review comments with AI agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.
Inline comments:
In `@rl_engine/mismatch/model_meta/qwen3.py`:
- Around line 62-63: Update QWEN3_0B5_SHAPE to the published Qwen3 0.6B
configuration dimensions, replacing the incorrect layer, hidden-size,
query-head, and key/value-head values. Then align QWEN3_SINGLE_LAYER_SHAPE with
the same per-layer dimensions while keeping its layer count at one.
In `@rl_engine/mismatch/operator_checks/attention/_common.py`:
- Around line 34-36: Update the pinned_libraries configuration to add an exact
FlashInfer package version or commit alongside the existing transformer_engine
LibraryPin, covering the flashinfer dependency used by rollout_impl and
apply_rope. Keep the Transformer Engine pin unchanged and ensure FlashInfer
cannot resolve from an unbounded version range.
In `@rl_engine/mismatch/operator_checks/gemm/_common.py`:
- Line 96: Update the torch version in the gemm reference’s pinned_libraries
configuration to align with the pyproject.toml requirement: either pin it to the
declared minimum 2.4.1 or raise the manifest minimum to 2.6.0, keeping both
declarations consistent.
In `@rl_engine/mismatch/operator_checks/logprob/factors/precision_downcast.py`:
- Around line 36-48: The precision downcast factor must vary the downcast
placement it claims to test. Update the factor definition around the visible
switch and comparison_rules to include variants that set DOWNCAST_POINTS and
allow precision.downcast_at to differ, or split the logic into separate
head-dtype and downcast-placement factors while preserving the existing scope
for each.
In `@rl_engine/mismatch/pipeline/comparison.py`:
- Around line 64-69: Update _values_equal to give bitwise and semantic float
comparisons distinct, intentional behavior: keep NaN equality only for the
semantic path, ensure MUST_MATCH_BITWISE does not treat separate NaN values as
equal, and configure a non-zero tolerance for semantic comparison if that is the
intended contract; otherwise remove the redundant float comparison branch while
preserving exact equality.
In `@rl_engine/mismatch/pipeline/diagnosis.py`:
- Around line 120-125: Update the shard validation around world_size and rank
collection to require every shard’s world_size to match, and require the rank
set to equal set(range(world_size)). Return Diagnosis.INSUFFICIENT_EVIDENCE with
the existing diagnostic outcome whenever either condition fails, rather than
validating only the first shard’s size and the number of distinct ranks.
In `@rl_engine/mismatch/pipeline/planner.py`:
- Around line 116-125: Update rl_engine/mismatch/pipeline/planner.py lines
116-125 in build_variants() to emit a baseline plus controlled parameter-sweep
variants using names and switch values recognized by diagnosis, rather than only
value_* arms; update rl_engine/mismatch/pipeline/diagnosis.py lines 147-170 in
_run_matrix() to support that parameter-sweep contract, or explicitly reject
parameter-sweep factors before entering the four-arm matrix. Ensure
logp.precision_downcast no longer incorrectly produces INSUFFICIENT_EVIDENCE.
- Around line 88-91: Update the package validation loop in the
prerequisite-planning logic to parse each full requirement, resolve the
installed distribution version, and enforce its declared version constraint
rather than only checking import availability. Continue appending an
UnmetPrerequisite for missing distributions or incompatible versions, while
preserving the existing package-name normalization for module discovery.
In `@rl_engine/mismatch/pipeline/report.py`:
- Around line 40-47: Update the filtering logic in the report-building flow
around proven, explained, and kept so proven correspondences are retained only
when they match a finding, and findings marked equivalent are excluded
consistently. In build_report, preserve and pass only the filtered reports to
trace_root_causes instead of tracing every report, ensuring unrelated false
positives and hypotheses for filtered equivalences are omitted.
In `@rl_engine/mismatch/pipeline/runner.py`:
- Around line 208-216: Update the repeat-processing flow around expand_repeats
so scores and readbacks are captured only from the first environment, while
repeats[role] continues collecting every run for topology-independence checks.
Ensure compute_metrics and effective_config consume the explicitly selected
first-environment results rather than values overwritten by later iterations.
- Around line 116-122: Update the mismatch metric calculation around the ratios
list and k3 estimator to avoid exponential overflow and logarithm domain errors:
clamp each delta before calling math.exp, while computing the k3 term with delta
directly instead of math.log(ratio). Update the worst-ratio selection near the
ratio_max calculation to use delta directly as well, preserving the existing
mismatch reporting behavior for extreme values.
- Around line 209-213: The score input in runner.py around the PolicyRole loop
must prevent repeat_under environment keys from overwriting required variant
settings: pass repeat environments through a separate environment channel, or
exclude those repeat keys from required-setting readback verification. Apply the
corresponding handling at forward_reduce.py lines 45-46, preserving required
settings while still supporting process restart configuration.
In `@rl_engine/mismatch/reference_adapters/settings.py`:
- Around line 92-102: Update the readback lookup logic to use setting.readback
rather than setting.key when checking membership in readback and retrieving
actual. Continue reporting setting.key in mismatch and unobservable results,
while preserving the existing observability and constraint handling.
In `@rl_engine/mismatch/schema/contracts.py`:
- Around line 37-68: Deep-freeze all caller-provided schema mappings during
construction: update ComparisonIssue.values and OperatorContract.extra in
contracts.py, MismatchAgent.comparison_rules in factors.py, and
FactorVariant.switch_values, replace_on, and nested repeat_under values in
variants.py. Use immutable mapping snapshots recursively so later mutations to
the original dictionaries or nested values cannot alter the frozen records.
In `@rl_engine/mismatch/schema/factors.py`:
- Around line 161-177: Reorder the exports in __all__ in factors.py according to
Ruff’s configured isort-style ordering to resolve RUF022, preserving every
existing exported symbol and its spelling.
In `@rl_engine/mismatch/schema/fingerprints.py`:
- Around line 35-66: The fingerprint schema at
rl_engine/mismatch/schema/fingerprints.py:35-66 must replace caller-owned
identity mappings with a recursively immutable, canonically normalized
representation. At rl_engine/mismatch/schema/fingerprints.py:83-87, reject
unsupported values or normalize them explicitly before hashing, and remove
default=str. At rl_engine/mismatch/schema/metrics.py:82-109, snapshot metadata
and effective configuration into that same immutable representation before
archiving results so VariantRecord.content_hash remains stable after source
dictionaries mutate.
In `@rl_engine/mismatch/schema/rollout_context.py`:
- Around line 61-70: Validate the sequence identity invariants in the rollout
context model: require response_token_ids, active_mask, and position_ids to have
equal lengths, and require group_size to match len(group.rollout_ids). Add these
checks at the context validation/construction boundary while preserving valid
rollout behavior.
In `@rl_engine/mismatch/schema/thresholds.py`:
- Around line 100-103: The tolerance_floor function currently omits
routing_replay when resolving expected_range, causing production MoE lookups to
fail. Add a routing_replay parameter to tolerance_floor, pass it through to
expected_range, and update _run_matrix to provide the effective routing state
when calling tolerance_floor.
In `@rl_engine/mismatch/schema/values.py`:
- Around line 88-91: Enforce the exact-version contract in the LibraryPin
constructor by rejecting specifier or range expressions in version, such as
“>=2.0”, and accepting only concrete observed versions. Keep the existing
package, commit, and container_digest fields unchanged.
- Around line 138-142: Update positive_int() to validate that value is a plain
int before calling int(value), rejecting booleans and numeric fractions;
preserve the existing positive-value check and ValueError behavior for
non-positive integers.
---
Nitpick comments:
In `@rl_engine/mismatch/__main__.py`:
- Around line 52-63: Validate the optional operator argument in both
command_list and command_plan after plugins are loaded, resolving it against the
registered operators. If the name is provided but unrecognized, emit a clear
error and return a nonzero exit status instead of treating it as an empty
result; preserve existing behavior when omitted or valid.
In `@rl_engine/mismatch/docs/add-a-kernel-factor.md`:
- Around line 23-29: Add the text language tag to all four plain fenced code
blocks: the reference-authority block at
rl_engine/mismatch/docs/add-a-kernel-factor.md lines 23-29, the directory tree
at lines 42-47, the skipped-prerequisites output at lines 165-171, and the
package layout tree at rl_engine/mismatch/README.md lines 72-82. No other
content changes are needed.
In `@rl_engine/mismatch/pipeline/comparison.py`:
- Around line 145-155: Remove the local strength mapping from the comparison
loop and compare left.determinism and right.determinism directly for inequality.
Update the condition in the comparison function while preserving the existing
issue-reporting behavior for differing determinism levels, so newly added
DeterminismLevel members cannot cause a KeyError.
- Around line 152-154: Update the collective comparison loop around paired and
the rollout.collectives/training.collectives lists to detect unequal lengths
instead of silently dropping the unpaired tail. Emit the existing
RECORD_ONLY-style issue for each missing collective, or otherwise record the
count difference in the comparison report, while preserving comparisons for
paired collectives.
In `@rl_engine/mismatch/pipeline/registry.py`:
- Around line 178-185: Sort the `__all__` entries in isort style by moving
`OPERATOR_CHECKS` before `FactorDiscoveryError`, while preserving all existing
exports.
- Around line 167-172: Update the factor filename validation around
expected_suffix to use only the final dot-separated segment of factor.id, so ids
such as logprob.precision.downcast expect downcast.py. Preserve the existing
mismatch error and factor discovery flow.
- Around line 91-115: Cache each plugin’s declared factors once during
registration as a tuple in self._factors, then reuse that cache in
_check_factor_conflicts and factors_for instead of calling declare_factors
repeatedly. Ensure newly registered plugins are added to the cache only after
conflict validation succeeds, and clear the cached factors in clear() alongside
the existing registry state.
In `@tests/test_mismatch_framework.py`:
- Around line 249-252: Update make_factor to use a sentinel default that
distinguishes an omitted reference from an explicit reference=None, while
preserving automatic reference creation for omitted arguments. In the test
around build_variants, remove the MismatchFactor __dict__ reconstruction and
rely directly on make_factor(reference=None) to create the reference-free
factor.
- Around line 462-471: Update the test setup around four_arms() to locate and
replace the result whose variant.name matches the intended arm, rather than
assigning by numeric index. Apply the same name-based selection to both
replacements currently using indices 2 and 0, preserving the existing
make_result values for training_reference_only and the other targeted arm.
🪄 Autofix
Fix all unresolved CodeRabbit comments on this PR:
- Push a commit to this branch (recommended)
- Create a new PR with the fixes
ℹ️ Review info
⚙️ Run configuration
Configuration used: defaults
Review profile: CHILL
Plan: Pro Plus
Run ID: 41b8780e-85fe-44bf-a9ee-76c9eead348e
📒 Files selected for processing (50)
rl_engine/mismatch/README.mdrl_engine/mismatch/__init__.pyrl_engine/mismatch/__main__.pyrl_engine/mismatch/docs/README.mdrl_engine/mismatch/docs/add-a-comm-feature.mdrl_engine/mismatch/docs/add-a-kernel-factor.mdrl_engine/mismatch/engines/__init__.pyrl_engine/mismatch/engines/megatron.pyrl_engine/mismatch/engines/vllm.pyrl_engine/mismatch/model_meta/__init__.pyrl_engine/mismatch/model_meta/qwen3.pyrl_engine/mismatch/operator_checks/__init__.pyrl_engine/mismatch/operator_checks/attention/__init__.pyrl_engine/mismatch/operator_checks/attention/_common.pyrl_engine/mismatch/operator_checks/attention/adapter.pyrl_engine/mismatch/operator_checks/attention/factors/__init__.pyrl_engine/mismatch/operator_checks/attention/factors/rope_fusion.pyrl_engine/mismatch/operator_checks/gemm/__init__.pyrl_engine/mismatch/operator_checks/gemm/_common.pyrl_engine/mismatch/operator_checks/gemm/adapter.pyrl_engine/mismatch/operator_checks/gemm/factors/__init__.pyrl_engine/mismatch/operator_checks/gemm/factors/forward_reduce.pyrl_engine/mismatch/operator_checks/logprob/__init__.pyrl_engine/mismatch/operator_checks/logprob/_common.pyrl_engine/mismatch/operator_checks/logprob/adapter.pyrl_engine/mismatch/operator_checks/logprob/factors/__init__.pyrl_engine/mismatch/operator_checks/logprob/factors/precision_downcast.pyrl_engine/mismatch/pipeline/__init__.pyrl_engine/mismatch/pipeline/comparison.pyrl_engine/mismatch/pipeline/diagnosis.pyrl_engine/mismatch/pipeline/planner.pyrl_engine/mismatch/pipeline/registry.pyrl_engine/mismatch/pipeline/report.pyrl_engine/mismatch/pipeline/runner.pyrl_engine/mismatch/reference_adapters/__init__.pyrl_engine/mismatch/reference_adapters/settings.pyrl_engine/mismatch/schema/__init__.pyrl_engine/mismatch/schema/collectives.pyrl_engine/mismatch/schema/contracts.pyrl_engine/mismatch/schema/factors.pyrl_engine/mismatch/schema/fingerprints.pyrl_engine/mismatch/schema/metrics.pyrl_engine/mismatch/schema/pitfalls.pyrl_engine/mismatch/schema/rollout_context.pyrl_engine/mismatch/schema/thresholds.pyrl_engine/mismatch/schema/tracing.pyrl_engine/mismatch/schema/values.pyrl_engine/mismatch/schema/variants.pytests/mismatch_cpu_backend.pytests/test_mismatch_framework.py
| QWEN3_0B5_SHAPE = "L=24,H=896,Hq=14,Hkv=2,D=64" | ||
| QWEN3_SINGLE_LAYER_SHAPE = "L=1,H=896,Hq=14,Hkv=2,D=64" |
There was a problem hiding this comment.
🎯 Functional Correctness | 🟠 Major | ⚡ Quick win
🧩 Analysis chain
🏁 Script executed:
#!/bin/bash
set -euo pipefail
curl -fsSL 'https://huggingface.co/Qwen/Qwen3-0.6B/raw/main/config.json' |
jq '{num_hidden_layers, hidden_size, num_attention_heads, num_key_value_heads, head_dim}'Repository: RL-Align/RL-Kernel
Length of output: 285
🏁 Script executed:
#!/bin/bash
set -euo pipefail
echo "== file excerpt =="
sed -n '1,120p' rl_engine/mismatch/model_meta/qwen3.py
echo
echo "== references to QWEN3_0B5_SHAPE and QWEN3_SINGLE_LAYER_SHAPE =="
rg -n "QWEN3_(0B5|0B6|SINGLE_LAYER)_SHAPE|Qwen3|0\.6" rl_engine tests .github 2>/dev/null || true
echo
echo "== nearby config names in model_meta =="
fd -a '.*\.py$' rl_engine/mismatch/model_meta 2>/dev/null | sed 's#^\./##' | sort
echo
echo "== inspect HuggingFace config summary =="
curl -fsSL 'https://huggingface.co/Qwen/Qwen3-0.6B/raw/main/config.json' |
jq '{model_type, num_hidden_layers, hidden_size, num_attention_heads, num_key_value_heads, head_dim, intermediate_size, rope_theta}'Repository: RL-Align/RL-Kernel
Length of output: 8055
🏁 Script executed:
#!/bin/bash
set -euo pipefail
echo "== model_meta imports and schema =="
sed -n '1,120p' rl_engine/mismatch/model_meta/__init__.py
sed -n '1,220p' rl_engine/mismatch/schema.py
echo
echo "== shape constant usages in tests =="
rg -n "QWEN3_(0B5|0B6|SINGLE_LAYER)_SHAPE|ShapeString|qwen3|Qwen3" tests rl_engine/mismatch -g '*.py' | sed -n '1,220p'
echo
echo "== model_meta tests =="
rg -n "test_.*qwen|QWEN3|model_meta|qwen3" tests/test_mismatch_model_meta.py -g '*.py' 2>/dev/null || true
sed -n '1,260p' tests/test_mismatch_model_meta.pyRepository: RL-Align/RL-Kernel
Length of output: 775
🏁 Script executed:
#!/bin/bash
set -euo pipefail
echo "== tracked files around mismatch =="
git ls-files rl_engine/mismatch | sort | sed -n '1,160p'
echo
echo "== search schema/constants =="
rg -n "ShapeString|parse_shape|QWEN3_(0B5|0B6|0B|8B|SINGLE_LAYER)_SHAPE|model_meta" rl_engine tests -g '*.py' | sed -n '1,220p'
echo
echo "== model_meta qwen3 test file if present =="
if [ -f tests/test_mismatch_model_meta.py ]; then
sed -n '1,280p' tests/test_mismatch_model_meta.py
else
echo "tests/test_mismatch_model_meta.py not found"
fiRepository: RL-Align/RL-Kernel
Length of output: 3498
Correct the small Qwen3 metadata constant.
QWEN3_0B5_SHAPE uses dimensions that do not match the published Qwen3 0.6B config. Use the real Qwen3 small-model layers and dimensions if this constant represents the published model, and align QWEN3_SINGLE_LAYER_SHAPE with those per-layer values.
🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.
In `@rl_engine/mismatch/model_meta/qwen3.py` around lines 62 - 63, Update
QWEN3_0B5_SHAPE to the published Qwen3 0.6B configuration dimensions, replacing
the incorrect layer, hidden-size, query-head, and key/value-head values. Then
align QWEN3_SINGLE_LAYER_SHAPE with the same per-layer dimensions while keeping
its layer count at one.
| pinned_libraries=( | ||
| LibraryPin("transformer_engine", "2.9.0.dev0", commit="8260f49"), | ||
| ), |
There was a problem hiding this comment.
🗄️ Data Integrity & Integration | 🟠 Major | ⚡ Quick win
🧩 Analysis chain
🏁 Script executed:
#!/bin/bash
set -euo pipefail
# Find the dependency declaration and lock metadata for the FlashInfer package.
fd -t f -i '^(pyproject\.toml|setup\.py|requirements.*|.*lock)$' . \
-x rg -n -i -C 2 'flashinfer|transformer_engine' {}Repository: RL-Align/RL-Kernel
Length of output: 458
🏁 Script executed:
#!/bin/bash
set -euo pipefail
echo "== candidate files =="
git ls-files | rg '(^|/)(pyproject\.toml|setup\.py|requirements.*|.*lock|poetry\.lock|uv\.lock|Pipfile\.lock|pixi\.lock|conda\.lock)$' || true
echo "== search flashinfer/transformer_engine in tracked files =="
rg -n -i -C 2 'flashinfer|transformer_engine|LibraryPin|rollout_impl|pinned_libraries' . || true
echo "== find _common.py =="
fd -t f -i '_common.py|attention' .Repository: RL-Align/RL-Kernel
Length of output: 33089
🏁 Script executed:
#!/bin/bash
set -euo pipefail
echo "== exact dependency declarations =="
python3 - <<'PY'
import tomllib, ast
from pathlib import Path
for path in [Path("pyproject.toml"), Path("setup.py")]:
print(f"\n--- {path} ---")
text = path.read_text()
if path.suffix == ".toml":
data = tomllib.loads(text)
for section in ("project.dependencies", "project.optional-dependencies.cuda"):
print(section, data.get(section, data.get(section.split(".")[-1], "N/A")))
else:
tree = ast.parse(text)
for node in ast.walk(tree):
if isinstance(node, ast.keyword) and node.arg == "extras_require":
print(node.value)
PY
echo "== lock/config files relevant to pip/uv/poetry =="
git ls-files | rg '(^|/)(uv\.lock|pdm\.lock|poetry\.lock|Pipfile\.lock|requirements.*\.in|.*requirements.*\.txt|pyproject\.lock)$' || true
echo "== LibraryPin schema and comparison =="
cat -n rl_engine/mismatch/schema/values.py | sed -n '70,95p'Repository: RL-Align/RL-Kernel
Length of output: 312
🏁 Script executed:
#!/bin/bash
set -euo pipefail
echo "== pyproject dependency lines =="
awk '/^\[project(\.dependencies)?\]|^\[project\.optional-dependencies\]|cuda = /{p=1} p{print NR": "$0} /^\[/{if(!/dependencies/&&!/optional/)p=0}' pyproject.toml
echo "== setup.py extras_require section =="
awk '/extras_require\s*=/{p=1} p{print NR": "$0} /^\}/{if(p)exit}' setup.py
echo "== lock/config files relevant to pip/uv/poetry =="
git ls-files | rg '(^|/)(uv\.lock|pdm\.lock|poetry\.lock|Pipfile\.lock|requirements.*\.in|.*requirements.*\.txt|pyproject\.lock)$' || true
echo "== LibraryPin schema and comparison =="
cat -n rl_engine/mismatch/schema/values.py | sed -n '70,95p'Repository: RL-Align/RL-Kernel
Length of output: 1977
Pin the FlashInfer rollout reference.
rollout_impl uses flashinfer.rope.apply_rope, and the docs say Transformer Engine/FlashInfer can change kernel selection across versions. pinned_libraries only records Transformer Engine, while the dependency is unbound (flashinfer-python>=0.1.6 / flashinfer). Add an exact FlashInfer package version or lock, together with the Transformer Engine pin, so the RoPE factor cannot flip based on a FlashInfer upgrade.
🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.
In `@rl_engine/mismatch/operator_checks/attention/_common.py` around lines 34 -
36, Update the pinned_libraries configuration to add an exact FlashInfer package
version or commit alongside the existing transformer_engine LibraryPin, covering
the flashinfer dependency used by rollout_impl and apply_rope. Keep the
Transformer Engine pin unchanged and ensure FlashInfer cannot resolve from an
unbounded version range.
| guards="nccl_algo_unpinned", | ||
| ), | ||
| ), | ||
| pinned_libraries=(LibraryPin("torch", "2.6.0"),), |
There was a problem hiding this comment.
🗄️ Data Integrity & Integration | 🟡 Minor | ⚡ Quick win
🧩 Analysis chain
🏁 Script executed:
#!/bin/bash
# Description: Compare the declared torch version with the pin used by the gemm reference.
set -euo pipefail
# Manifest-declared torch versions.
fd -H -t f -e toml -e txt -e cfg . | xargs rg -n 'torch\s*[=<>~!]' || true
# Where LibraryPin values are consumed.
rg -n -C 4 'pinned_libraries|LibraryPin'Repository: RL-Align/RL-Kernel
Length of output: 226
🏁 Script executed:
#!/bin/bash
set -euo pipefail
echo "== tracked files matching common / gemm =="
git ls-files | rg '(^|/)(pyproject\.toml|requirements.*\.txt|_common\.py)$|gemm/|operator_checks/mismatch' || true
echo "== _common.py relevant area =="
cat -n rl_engine/mismatch/operator_checks/gemm/_common.py | sed -n '1,180p'
echo "== LibraryPin / pinned_libraries occurrences =="
rg -n -C 3 'class LibraryPin|LibraryPin|pinned_libraries|verify|raise|check' rl_engine pyproject.toml requirements.txt 2>/dev/null || trueRepository: RL-Align/RL-Kernel
Length of output: 50375
🏁 Script executed:
#!/bin/bash
set -euo pipefail
echo "== exact matches for pinned_libraries and LibraryPin =="
rg -n -C 3 '\bpinned_libraries\b|\bLibraryPin\b' .
echo "== git diff summary and torch pin context =="
git diff --stat
git diff -- rl_engine/mismatch/operator_checks/gemm/_common.py pyproject.toml requirements.txt | sed -n '1,220p'
echo "== _common.py lines 70-130 =="
cat -n rl_engine/mismatch/operator_checks/gemm/_common.py | sed -n '70,130p'Repository: RL-Align/RL-Kernel
Length of output: 12773
Align the gemm torch pin with the project manifest.
pyproject.toml declares torch>=2.4.1, while this gemm reference uses LibraryPin("torch", "2.6.0"). Use the declared minimum as the pinned library version, or raise the manifest requirement so the reference can run under the declared toolchain.
🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.
In `@rl_engine/mismatch/operator_checks/gemm/_common.py` at line 96, Update the
torch version in the gemm reference’s pinned_libraries configuration to align
with the pyproject.toml requirement: either pin it to the declared minimum 2.4.1
or raise the manifest minimum to 2.6.0, keeping both declarations consistent.
| switch=Switch( | ||
| path="logp.head_dtype", | ||
| rebind_cost=RebindCost.PER_REQUEST, | ||
| # vLLM computes logits at the model dtype, so only training can vary. | ||
| applies_to=(PolicyRole.TRAINING,), | ||
| allowed_values=tuple(HEAD_DTYPES), | ||
| ), | ||
| comparison_rules={ | ||
| "precision.lm_head": ComparisonRule.MUST_MATCH_BITWISE, | ||
| "precision.accumulate": ComparisonRule.MUST_MATCH_SEMANTICALLY, | ||
| "precision.downcast_at": ComparisonRule.MUST_MATCH_SEMANTICALLY, | ||
| "extra.logprobs_mode": ComparisonRule.RECORD_ONLY, | ||
| }, |
There was a problem hiding this comment.
🎯 Functional Correctness | 🟠 Major | 🏗️ Heavy lift
Ablate the downcast placement or narrow the factor scope.
This factor claims to test where fp32 values are written back. Its variants can only change logp.head_dtype. DOWNCAST_POINTS is not used, and precision.downcast_at must match. The planner therefore cannot test the stated downcast hypothesis.
Add variants that set the downcast placement, or split this into a head-dtype factor and a downcast-placement factor.
🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.
In `@rl_engine/mismatch/operator_checks/logprob/factors/precision_downcast.py`
around lines 36 - 48, The precision downcast factor must vary the downcast
placement it claims to test. Update the factor definition around the visible
switch and comparison_rules to include variants that set DOWNCAST_POINTS and
allow precision.downcast_at to differ, or split the logic into separate
head-dtype and downcast-placement factors while preserving the existing scope
for each.
| def _values_equal(left: Any, right: Any, *, bitwise: bool) -> bool: | ||
| if isinstance(left, float) and isinstance(right, float): | ||
| if math.isnan(left) and math.isnan(right): | ||
| return True | ||
| return left == right if bitwise else math.isclose(left, right, rel_tol=0.0, abs_tol=0.0) | ||
| return left == right |
There was a problem hiding this comment.
🎯 Functional Correctness | 🟠 Major | ⚡ Quick win
The bitwise flag has no effect on float comparison.
math.isclose(left, right, rel_tol=0.0, abs_tol=0.0) reduces to exact equality. Both branches of _values_equal therefore evaluate the same result for floats, and MUST_MATCH_SEMANTICALLY compares floats exactly despite the name.
The NaN rule also runs under MUST_MATCH_BITWISE, so two NaN values compare as equal on a rule that claims bitwise identity.
Decide the intended semantics and encode them. If MUST_MATCH_SEMANTICALLY needs a tolerance, set a non-zero rel_tol. If it does not, delete the float branch and keep only the NaN handling that the bitwise rule needs.
🔧 Proposed explicit semantics
-def _values_equal(left: Any, right: Any, *, bitwise: bool) -> bool:
+SEMANTIC_REL_TOL = 1e-9
+
+
+def _values_equal(left: Any, right: Any, *, bitwise: bool) -> bool:
if isinstance(left, float) and isinstance(right, float):
- if math.isnan(left) and math.isnan(right):
- return True
- return left == right if bitwise else math.isclose(left, right, rel_tol=0.0, abs_tol=0.0)
+ if bitwise:
+ return math.copysign(1.0, left) == math.copysign(1.0, right) and (
+ left == right or (math.isnan(left) and math.isnan(right))
+ )
+ if math.isnan(left) and math.isnan(right):
+ return True
+ return math.isclose(left, right, rel_tol=SEMANTIC_REL_TOL, abs_tol=0.0)
return left == rightAlso applies to: 113-121
🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.
In `@rl_engine/mismatch/pipeline/comparison.py` around lines 64 - 69, Update
_values_equal to give bitwise and semantic float comparisons distinct,
intentional behavior: keep NaN equality only for the semantic path, ensure
MUST_MATCH_BITWISE does not treat separate NaN values as equal, and configure a
non-zero tolerance for semantic comparison if that is the intended contract;
otherwise remove the redundant float comparison branch while preserving exact
equality.
| class EnvironmentFingerprint: | ||
| """The execution environment. Change this layer and every number is stale.""" | ||
|
|
||
| python_version: str | ||
| torch_version: str | ||
| torch_build_hash: str # hash of the build config (cuda/hip build, op set) | ||
| driver_version: str | ||
| device_model: str | ||
| libraries: tuple[LibraryPin, ...] | ||
| determinism_env: Mapping[str, str] # NVTE_* / CUBLAS_* / NCCL_* / torch backends | ||
| source_revision: str # this framework's own version | ||
|
|
||
|
|
||
| @dataclass(frozen=True) | ||
| class ExecutionFingerprint: | ||
| """One execution's full identity. Any part differing makes two runs | ||
| incomparable. | ||
|
|
||
| What goes in is the value read back, never the value requested: asking for | ||
| ``num_splits=1`` and the backend using 1 are two different facts. Thresholds | ||
| go in too, so changing one makes every historical pass/fail stale -- which is | ||
| why thresholds are code constants, a configurable value cannot be pinned into | ||
| an identity. | ||
| """ | ||
|
|
||
| identity: str # fingerprint of the ComparisonIdentity | ||
| environment: EnvironmentFingerprint | ||
| switch_binding: str # effective switch values, read back | ||
| implementation: Mapping[PolicyRole, str] # what each side actually instantiated | ||
| model_state: Mapping[PolicyRole, str] # each side's weights | ||
| collectives: tuple[str, ...] # fingerprints of the collectives that ran | ||
| threshold_table: str # fingerprint of EXPECTED_RANGES |
There was a problem hiding this comment.
🗄️ Data Integrity & Integration | 🟠 Major | 🏗️ Heavy lift
Make archived execution data deeply immutable and canonically serializable.
frozen=True prevents attribute rebinding but does not freeze caller-owned dictionaries. Also, default=str hashes unsupported values from their string representation. A later dictionary mutation can make a VariantRecord.content_hash no longer describe its record. Unsupported values can also produce non-canonical fingerprints.
rl_engine/mismatch/schema/fingerprints.py#L35-L66: replace identity mappings with a recursively immutable normalized representation.rl_engine/mismatch/schema/fingerprints.py#L83-L87: reject unsupported values or normalize them explicitly before hashing; do not usedefault=str.rl_engine/mismatch/schema/metrics.py#L82-L109: snapshot metadata and effective configuration into the same immutable representation before archiving results.
📍 Affects 2 files
rl_engine/mismatch/schema/fingerprints.py#L35-L66(this comment)rl_engine/mismatch/schema/fingerprints.py#L83-L87rl_engine/mismatch/schema/metrics.py#L82-L109
🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.
In `@rl_engine/mismatch/schema/fingerprints.py` around lines 35 - 66, The
fingerprint schema at rl_engine/mismatch/schema/fingerprints.py:35-66 must
replace caller-owned identity mappings with a recursively immutable, canonically
normalized representation. At rl_engine/mismatch/schema/fingerprints.py:83-87,
reject unsupported values or normalize them explicitly before hashing, and
remove default=str. At rl_engine/mismatch/schema/metrics.py:82-109, snapshot
metadata and effective configuration into that same immutable representation
before archiving results so VariantRecord.content_hash remains stable after
source dictionaries mutate.
| prompt_token_ids: tuple[int, ...] | ||
| response_token_ids: tuple[int, ...] | ||
| active_mask: tuple[bool, ...] # loss mask: which tokens participate | ||
| position_ids: tuple[int, ...] | ||
| checkpoint_id: str | ||
| checkpoint_revision: str | ||
| model_shape: str # a trimmed model is a different model | ||
| group: RolloutGroup | ||
| batch_placement: BatchPlacement | ||
| sampling_decision: DynamicSamplingDecision |
There was a problem hiding this comment.
🎯 Functional Correctness | 🟡 Minor | ⚡ Quick win
Validate the sequence identity invariants.
Require response_token_ids, active_mask, and position_ids to have the same length. Require group_size to equal len(group.rollout_ids). The runner derives positions from active_mask and indexes token IDs with them. An invalid identity can therefore fail during scoring or report the wrong worst token.
Proposed fix
`@dataclass`(frozen=True)
class ComparisonIdentity:
@@
sampling_decision: DynamicSamplingDecision
+
+ def __post_init__(self) -> None:
+ sequence_length = len(self.response_token_ids)
+ if len(self.active_mask) != sequence_length or len(self.position_ids) != sequence_length:
+ raise ValueError("response_token_ids, active_mask, and position_ids must align")
+ if self.group.group_size != len(self.group.rollout_ids):
+ raise ValueError("group_size must match rollout_ids")📝 Committable suggestion
‼️ IMPORTANT
Carefully review the code before committing. Ensure that it accurately replaces the highlighted code, contains no missing lines, and has no issues with indentation. Thoroughly test & benchmark the code to ensure it meets the requirements.
| prompt_token_ids: tuple[int, ...] | |
| response_token_ids: tuple[int, ...] | |
| active_mask: tuple[bool, ...] # loss mask: which tokens participate | |
| position_ids: tuple[int, ...] | |
| checkpoint_id: str | |
| checkpoint_revision: str | |
| model_shape: str # a trimmed model is a different model | |
| group: RolloutGroup | |
| batch_placement: BatchPlacement | |
| sampling_decision: DynamicSamplingDecision | |
| prompt_token_ids: tuple[int, ...] | |
| response_token_ids: tuple[int, ...] | |
| active_mask: tuple[bool, ...] # loss mask: which tokens participate | |
| position_ids: tuple[int, ...] | |
| checkpoint_id: str | |
| checkpoint_revision: str | |
| model_shape: str # a trimmed model is a different model | |
| group: RolloutGroup | |
| batch_placement: BatchPlacement | |
| sampling_decision: DynamicSamplingDecision | |
| def __post_init__(self) -> None: | |
| sequence_length = len(self.response_token_ids) | |
| if len(self.active_mask) != sequence_length or len(self.position_ids) != sequence_length: | |
| raise ValueError("response_token_ids, active_mask, and position_ids must align") | |
| if self.group.group_size != len(self.group.rollout_ids): | |
| raise ValueError("group_size must match rollout_ids") |
🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.
In `@rl_engine/mismatch/schema/rollout_context.py` around lines 61 - 70, Validate
the sequence identity invariants in the rollout context model: require
response_token_ids, active_mask, and position_ids to have equal lengths, and
require group_size to match len(group.rollout_ids). Add these checks at the
context validation/construction boundary while preserving valid rollout
behavior.
| def tolerance_floor(model_family: str, noise_floor: NoiseFloor) -> float: | ||
| """The floor below which a difference is not treated as a signal.""" | ||
|
|
||
| return expected_range(model_family, noise_floor).suspect_above |
There was a problem hiding this comment.
🎯 Functional Correctness | 🟠 Major | ⚡ Quick win
Propagate routing_replay into the tolerance lookup.
For model_family="moe" or "large_moe" at NoiseFloor.PRODUCTION, this function calls expected_range() with routing_replay=None. The table has no such entry, so diagnosis raises ThresholdLookupError instead of classifying the factor. Add a routing_replay parameter here and pass the effective routing state from _run_matrix.
🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.
In `@rl_engine/mismatch/schema/thresholds.py` around lines 100 - 103, The
tolerance_floor function currently omits routing_replay when resolving
expected_range, causing production MoE lookups to fail. Add a routing_replay
parameter to tolerance_floor, pass it through to expected_range, and update
_run_matrix to provide the effective routing state when calling tolerance_floor.
| package: str | ||
| version: str # exact; ranges are not accepted | ||
| commit: str | None = None | ||
| container_digest: str | None = None # the only truly reproducible anchor |
There was a problem hiding this comment.
🗄️ Data Integrity & Integration | 🟠 Major | ⚡ Quick win
Enforce exact LibraryPin.version values.
The constructor accepts version ranges such as ">=2.0" despite the exact-version contract. A range does not identify the installed kernel version. The same execution fingerprint can then describe runs with different libraries. Reject specifiers here, or populate this field only from observed package metadata.
🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.
In `@rl_engine/mismatch/schema/values.py` around lines 88 - 91, Enforce the
exact-version contract in the LibraryPin constructor by rejecting specifier or
range expressions in version, such as “>=2.0”, and accepting only concrete
observed versions. Keep the existing package, commit, and container_digest
fields unchanged.
| def positive_int(value: Any) -> int: | ||
| parsed = int(value) | ||
| if parsed <= 0: | ||
| raise ValueError(f"expected a positive integer, got {value!r}") | ||
| return parsed |
There was a problem hiding this comment.
🎯 Functional Correctness | 🟡 Minor | ⚡ Quick win
🧩 Analysis chain
🏁 Script executed:
#!/bin/bash
set -euo pipefail
# Trace callers that use this parser for configuration values.
rg -n -C 3 '\bpositive_int\b|gpu[_-]?count' rl_engine testsRepository: RL-Align/RL-Kernel
Length of output: 9165
🏁 Script executed:
#!/bin/bash
set -euo pipefail
# Inspect the parser definitions and exports with line numbers.
sed -n '1,175p' rl_engine/mismatch/schema/values.py
# Locate positive_int usage more broadly.
rg -n '\bpositive_int\b|SettingChannel|RequiredSetting|Prerequisites|min_gpu_count' rl_engine testsRepository: RL-Align/RL-Kernel
Length of output: 10419
🏁 Script executed:
#!/bin/bash
set -euo pipefail
python3 - <<'PY'
from typing import Any
def positive_int_current(value: Any) -> int:
parsed = int(value)
if parsed <= 0:
raise ValueError(f"expected a positive integer, got {value!r}")
return parsed
def positive_int_rejected(value: Any) -> int:
if not isinstance(value, int) or isinstance(value, bool):
raise TypeError(f"expected a positive integer, got {value!r}")
if value <= 0:
raise ValueError(f"expected a positive integer, got {value!r}")
return value
for candidate in [1, True, 1.9, "1", -1, 0, "2.0", "x"]:
current = None
rejected = None
try:
current = positive_int_current(candidate)
except Exception as exc:
current = type(exc).__name__ + ": " + str(exc)
try:
rejected = positive_int_rejected(candidate)
except Exception as exc:
rejected = type(exc).__name__ + ": " + str(exc)
print(
f"{candidate!r}: current={current!r}; rejected={rejected!r}; "
f"same={current==rejected}; current_ok_accepts_nonint={current == 1 and candidate not in (1, -1, 0, 42)}"
)
PYRepository: RL-Align/RL-Kernel
Length of output: 1272
Reject non-integer inputs before converting.
positive_int() currently accepts booleans like True and numeric fractions like 1.9, and truncates them to positive values. Reject values that are not plain int objects before calling int(value), or otherwise enforce strict integer input.
🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.
In `@rl_engine/mismatch/schema/values.py` around lines 138 - 142, Update
positive_int() to validate that value is a plain int before calling int(value),
rejecting booleans and numeric fractions; preserve the existing positive-value
check and ValueError behavior for non-positive integers.
Signed-off-by: Zhang Jian <jianmusings@gmail.com>
There was a problem hiding this comment.
Actionable comments posted: 1
🧹 Nitpick comments (1)
rl_engine/mismatch/schema/__init__.py (1)
115-194: 📐 Maintainability & Code Quality | 🔵 Trivial | ⚡ Quick winSort
__all__to satisfy RUF022.Ruff reports that this export list is not sorted. Sort the names in
__all__, or apply the configured Ruff auto-fix.🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the rest with a brief reason, keep changes minimal, and validate. In `@rl_engine/mismatch/schema/__init__.py` around lines 115 - 194, Sort the exported names in __all__ alphabetically to satisfy Ruff rule RUF022, preserving every existing export and its spelling.Source: Linters/SAST tools
🤖 Prompt for all review comments with AI agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.
Inline comments:
In `@rl_engine/mismatch/pipeline/__init__.py`:
- Around line 53-83: The __all__ list in the module is not in isort-style order,
triggering Ruff RUF022. Reorder the existing exported names alphabetically
without adding, removing, or renaming any entries.
---
Nitpick comments:
In `@rl_engine/mismatch/schema/__init__.py`:
- Around line 115-194: Sort the exported names in __all__ alphabetically to
satisfy Ruff rule RUF022, preserving every existing export and its spelling.
🪄 Autofix
Fix all unresolved CodeRabbit comments on this PR:
- Push a commit to this branch (recommended)
- Create a new PR with the fixes
ℹ️ Review info
⚙️ Run configuration
Configuration used: defaults
Review profile: CHILL
Plan: Pro Plus
Run ID: 12978c19-b58e-47a4-a0dc-3e80119eb80f
📒 Files selected for processing (10)
rl_engine/mismatch/__init__.pyrl_engine/mismatch/model_meta/__init__.pyrl_engine/mismatch/operator_checks/attention/_common.pyrl_engine/mismatch/operator_checks/attention/adapter.pyrl_engine/mismatch/pipeline/__init__.pyrl_engine/mismatch/reference_adapters/__init__.pyrl_engine/mismatch/schema/__init__.pyrl_engine/mismatch/schema/metrics.pytests/mismatch_cpu_backend.pytests/test_mismatch_framework.py
🚧 Files skipped from review as they are similar to previous changes (8)
- rl_engine/mismatch/reference_adapters/init.py
- rl_engine/mismatch/model_meta/init.py
- rl_engine/mismatch/operator_checks/attention/_common.py
- rl_engine/mismatch/operator_checks/attention/adapter.py
- rl_engine/mismatch/init.py
- tests/test_mismatch_framework.py
- rl_engine/mismatch/schema/metrics.py
- tests/mismatch_cpu_backend.py
| __all__ = [ | ||
| "compare_contracts", | ||
| "resolve_field_path", | ||
| "CONVERGENCE_RATIO", | ||
| "diagnose", | ||
| "ContradictoryFactor", | ||
| "UnmetPrerequisite", | ||
| "build_variants", | ||
| "missing_prerequisites", | ||
| "order_cases_by_rebind_cost", | ||
| "reject_contradictory_factors", | ||
| "suggested_floor_is_lowest", | ||
| "OPERATOR_CHECKS", | ||
| "FactorDiscoveryError", | ||
| "OperatorChecks", | ||
| "PluginRegistry", | ||
| "RegistrationError", | ||
| "discover_factors", | ||
| "build_report", | ||
| "filter_known_equivalences", | ||
| "render_summary", | ||
| "trace_root_causes", | ||
| "ReadOnlyViolation", | ||
| "RunContext", | ||
| "ScoringBackend", | ||
| "assert_comparison_is_read_only", | ||
| "assert_order_is_topology_independent", | ||
| "compute_metrics", | ||
| "expand_repeats", | ||
| "run_variant", | ||
| ] |
There was a problem hiding this comment.
📐 Maintainability & Code Quality | 🟡 Minor | ⚡ Quick win
Sort __all__ to clear Ruff RUF022.
Ruff reports that this export list is not sorted. Apply isort-style ordering without changing the exported names.
🧰 Tools
🪛 Ruff (0.16.1)
[warning] 53-83: __all__ is not sorted
Apply an isort-style sorting to __all__
(RUF022)
🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.
In `@rl_engine/mismatch/pipeline/__init__.py` around lines 53 - 83, The __all__
list in the module is not in isort-style order, triggering Ruff RUF022. Reorder
the existing exported names alphabetically without adding, removing, or renaming
any entries.
Source: Linters/SAST tools
Latest Status [9 Aug 2026]
Ready for review.
Motivation
This PR supercedes (is an updated version of) #230. It did not delete any feature from it. Now GEMM, Attention and Logprob can use this interface to do cross-config alignment.
Hey @CyberSecurityErial, can you add your all reduce kernel onto this interface, see if it can work?
Hey @inaniloquentee, can you try update #236, see if it can work?
You can directly push to this branch to change it.
If this can work, I will cc Siru and KJ to add GEMM and Logprob.
Training-inference cross-config alignment framework
Rollout and training compute logprobs for the same tokens with the same weights and still disagree. As such, @CyberSecurityErial developed a framework that will automatically detect weak spots in your settings that caused this inconsistency. You can then seamlessly switch these settings to our framework's implementation, which guarantees consistency.
Here we provide the interface for GEMM, Attention and Logprob, with one example mismatch factor registered for each. The four adapter methods behind them still raise
NotImplementedError.Details on how this framework runs and how to add possible mismatch factors is in
rl_engine/mismatch/README.md.What is in it
Tests are in
tests/test_mismatch_framework.py— 39 of them, covering planning, execution, the four gates, diagnosis, reporting and thresholds. They run on CPU in under a second, driven bytests/mismatch_cpu_backend.py, a scoring harness that can fake the failure modes the framework exists to catch: a one-sided bias, a switch that silently does nothing, and output that changes with the environment.Three Kernel API: GEMM, Attention, LogProb
For each kernel we registered one example factor showing how a new one is added:
attn.rope_fusion,gemm.forward_reduce,logp.precision_downcast. They are deliberately three different shapes — an implementation swap against a shared backend, a collective communication factor, and a parameter sweep — so a new factor can be matched to the nearest one.You can list the registered factors, and expand them into the cases that would run at a given noise floor:
--noise-floorsays how much noise the run itself carries, which decides how small a difference is resolvable: atsingle_layer_anchorthe expectation is bitwise equality, while atproductionadlogp_meanof 0.002–0.008 is normal. Reading a low-floor result against the production band would call a definite operator bug "normal", so every threshold is keyed on it.Example: how much does RoPE fusion move the numbers?
This is not the framework running. The three adapters raise
NotImplementedError, so nothing in this PR can produce these numbers yet. They come from calling the kernels directly, and they are whatattn.rope_fusionis declared to measure once its adapter is filled in. For clean-ness, I deleted the testing script.Environment: RTX PRO 6000 Blackwell (sm_120), driver 580.126.09, CUDA 13.0,
torch 2.13.0+cu130, TransformerEngine 2.19.0.dev0+8260f49 built with
NVTE_FRAMEWORK=pytorch NVTE_CUDA_ARCHS=120 NVTE_WITH_NCCL_EP=0,flashinfer-python 0.6.16.post3. Random weights, no trained checkpoint.
attn.rope_fusion— S=512, Hq=14, D=64, θ=1e6Row 3 shows the fused kernel in bf16 is bitwise identical to computing in fp32 and downcasting, so it already accumulates in fp32. The 9.55e-4 in row 1 therefore comes from the unfused path, which does the arithmetic in bf16. Row 5 rules out nondeterminism. Row 4 is the only cross-framework comparison here: TransformerEngine on the training side against FlashInfer, which is what vLLM dispatches to when it is available.
Summary by CodeRabbit
New Features
Documentation
Tests