Skip to content

RFC: UniRL Experimental Plugin Framework for Train–Inference Parity #431

Description

@cboss6

Summary

This RFC introduces model-level train–inference parity in UniRL as the foundation for true on-policy RL. Given the same policy snapshot and sampled trajectory, rollout and training replay must produce identical objective-facing policy log-probabilities.

Without this guarantee, an apparently on-policy run may optimize against a numerically different or stale policy. UniRL's parity capability turns true-on-policy behavior from an implicit assumption into an explicitly verified runtime property.

The first implementation establishes this contract for autoregressive GRPO with Qwen3-30B-A3B, an FSDP actor, and direct vLLM rollout. Support for diffusion models, Megatron training, and performance-optimized kernels are follow-up work.

Motivation

RL rollouts and differentiable replay are often executed by different software stacks. Rollout log-probabilities define the behavior policy, while replay log-probabilities are used by policy ratios, clipping, and KL-related terms. If they differ, training optimizes a different objective even when generated outputs look the same. If rollout weights are stale or partially synchronized, the two paths no longer represent the same policy snapshot.

Train–inference parity is therefore an industry requirement. The framework work is still necessary because UniRL must make that requirement enforceable across independent training and inference engines: align their objective-facing observables, validate all parity affecting settings, fail before backward, and record reproducible evidence.

Proposed Plugin Framework

flowchart TB
    Profile["Hydra parity profile<br/>model · revision · topology · precision · sampling"]

    subgraph Control["Control plane: validate before launch"]
        Contract["ParityRuntimeContract<br/>single source of truth"]
        GlobalCheck{"Global preflight passed?"}
        EarlyFail["Fail before Ray connection<br/>or model construction"]
        Propagate["Propagate the validated environment<br/>to every Ray worker"]

        Contract --> GlobalCheck
        GlobalCheck -- No --> EarlyFail
        GlobalCheck -- Yes --> Propagate
    end

    Profile --> Contract

    subgraph Providers["Shared parity provider contract"]
        Common["Common reference operators<br/>attention · norm · reductions · precision"]
        ModelPatch["Model-specific adapter<br/>RoPE · logits · routing · MoE"]
        Common --> ModelPatch
    end

    subgraph Training["Training worker plugin: FSDP replay"]
        TrainEntry["UniRL external_libs entry"]
        TrainOptIn{"Explicit parity opt-in?"}
        TrainPreflight["Validate versions, symbols,<br/>and operator availability"]
        TrainInstall["Install training-side providers"]
        ExactContext["Enter exact_context<br/>through replay and backward"]
        Replay["Differentiable replay<br/>selected-token log-probabilities in FP32"]

        TrainEntry --> TrainOptIn
        TrainOptIn -- No --> TrainFail["Reject parity profile"]
        TrainOptIn -- Yes --> TrainPreflight
        TrainPreflight --> TrainInstall
        TrainInstall --> ExactContext --> Replay
    end

    subgraph Rollout["Rollout worker plugin: direct vLLM"]
        VllmEntry["vllm.general_plugins entry point"]
        Enable{"UNIRL_PARITY_ENABLE=1?"}
        NoOp["No-op<br/>normal vLLM behavior"]
        PluginConfig["Load profile, model,<br/>strict mode, and ordered patch set"]
        PluginPreflight["Phase 1: preflight all installers<br/>versions · signatures · compatibility"]
        InstallCommon["Phase 2a: install common providers"]
        InstallModel["Phase 2b: install model adapter"]
        Manifest["Verify installed symbol identities<br/>and emit runtime manifest"]
        Generate["Generate rollout<br/>selected-token log-probabilities in FP32"]

        VllmEntry --> Enable
        Enable -- No --> NoOp
        Enable -- Yes --> PluginConfig --> PluginPreflight
        PluginPreflight --> InstallCommon --> InstallModel --> Manifest --> Generate
    end

    Propagate --> TrainEntry
    Propagate --> VllmEntry
    ModelPatch -. provider implementations .-> TrainInstall
    Common -. provider implementations .-> InstallCommon
    ModelPatch -. provider implementations .-> InstallModel

    subgraph Verification["True-on-policy verification loop"]
        Snapshot["Committed policy snapshot<br/>model_version + fingerprint"]
        Compare["Collective-safe exact parity gate<br/>shape · dtype · finite · torch.equal"]
        Gate{"Exact parity?"}
        Stop["All ranks fail together<br/>before backward"]
        Update["Backward and optimizer update"]
        Sync["Transactional FSDP-to-vLLM weight sync"]
        Commit["All TP workers acknowledge load<br/>commit model version · reset caches"]
        Artifact["Atomic verification artifact<br/>initial_checkpoint + post_update_reload"]

        Snapshot --> Generate
        Snapshot --> Replay
        Generate --> Compare
        Replay --> Compare
        Compare --> Gate
        Gate -- No --> Stop
        Gate -- Yes --> Update --> Sync --> Commit
        Commit --> Snapshot
        Compare -. metrics .-> Artifact
        Manifest -. installed providers .-> Artifact
        Commit -. weight receipts .-> Artifact
    end
Loading

1. Profile-driven runtime contract

A profile freezes the model revision, topology, package versions, precision, sampling, engine switches, and weight-sync path. Preflight aggregates contract violations before Ray connection or model construction, then propagates the same explicit environment to every worker.

2. Training-side adapter

UniRL's external_libs hook loads a model adapter into dedicated FSDP workers. The adapter installs reference operators after explicit opt-in and exposes exact_context() for replay and activation-checkpoint recomputation.

3. Rollout-side vLLM plugin

The rollout adapter is a separately installable vllm.general_plugins package and is a no-op unless UNIRL_PARITY_ENABLE=1. Once enabled, its two-phase registry first validates the profile, versions, patch set, and private symbol signatures, then installs common providers followed by model-specific patches. It fails closed and records the before/after providers in a runtime manifest.

New model support is added as another profile plus a small model adapter. Shared numerical operators stay in the plugin's common layer.

4. Parity gate and evidence

The algorithm compares the RL observable before backward. The first strict profile requires identical finite FP32 selected-token log-probabilities:

torch.equal(replay_selected_token_logp, rollout_selected_token_logp)

The gate is collective-safe, so all FSDP ranks continue or fail together. A dedicated FSDP-to-vLLM sync records per-TP-worker weight receipts. The JSON artifact must cover both initial_checkpoint and post_update_reload, together with parity metrics, model versions, fingerprints, runtime information, and the plugin manifest.

First Implementation

The initial profile intentionally covers only:

  • autoregressive GRPO;
  • Qwen/Qwen3-30B-A3B at one frozen revision;
  • tensor-parallel configuration: a four-rank FSDP actor with a direct vLLM TP4 rollout engine on a single four-GPU node;
  • expert-parallel configuration: a four-rank FSDP actor with a direct vLLM EP4 rollout engine using the DeepEP high-throughput backend on a single four-GPU node;
  • BF16 model compute and FP32 selected-token processed log-probabilities;
  • correctness-first implementations for attention, RoPE, normalization,
    reductions, logits, and Qwen3-MoE routing/expert-combine paths.

Future Scope

  • Add diffusion models and more AR models.
  • Add a Megatron weight exporter while reusing the receiver, gate, and artifact.
  • Admit optimized kernels only after conformance with the reference providers.

Acceptance Criteria for the Initial Profile

  • With the enable flag absent, normal UniRL and vLLM execution is unchanged.
  • Invalid configuration or incompatible symbols fail before model launch or partial plugin installation.
  • Rollout and replay compare all selected-token FP32 log-probabilities before backward, with zero mismatches.
  • A post-update rollout uses the committed model version on every vLLM TP rank.
  • The verification artifact contains both required phases and enough provenance to reproduce the claim.

Expected Behavior

The following Qwen3-30B-A3B run uses vLLM TP4 for rollout and FSDP for training, demonstrating exact equality between the rollout and replay selected-token log-probabilities:

Image

Related Work

feat(experimental): add bitwise Qwen3-MoE train–inference parity case based on GRPO + vLLM(TP>1) + FSDP

Decision Requested

Approve the experimental plugin direction, the narrow first profile, and the rule that future model, engine, and backend support must enter through explicit profiles with reproducible parity evidence.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions