Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
55 changes: 55 additions & 0 deletions experimental/train_inference_parity/vllm_plugin/README.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,55 @@
# UniRL train/inference parity vLLM plugin

An opt-in, correctness-first plugin for Qwen3-30B-A3B. It is not a general
performance optimization or a claim of portable bitwise reproducibility.

## Installation

Install from this directory with `pip install -e '.[verified]'`. The `verified`
extra pins vLLM 0.27.0, Torch 2.13.0, and Transformers 5.6.0. UniRL must also be
installed for its native weight-sync worker extension. The parent experiment's
README documents the four-GPU recipe and two-phase verification.

The wheel registers `unirl_train_inference_parity` in `vllm.general_plugins`.
Without `UNIRL_PARITY_ENABLE=1`, registration imports neither Torch nor vLLM and
does not patch anything. An enabled worker requires all of:

```bash
export UNIRL_PARITY_ENABLE=1
export UNIRL_PARITY_PROFILE=public_reference
export UNIRL_PARITY_MODEL=qwen3_moe_30b_a3b
export UNIRL_PARITY_PATCHES=common,qwen3_moe_30b_a3b
export UNIRL_PARITY_STRICT=1
export VLLM_PLUGINS=unirl_train_inference_parity
```

Use the experiment launcher rather than these flags alone for parity validation:
it also configures precision, attention, NCCL, and Ray propagation.

## Ownership

- `common/`: shared numerical providers, precision, attention, and reductions.
- `models/qwen3_moe_30b_a3b/`: Qwen-specific routing, projections, MoE, and reload.
- `compat.py`: pinned versions and private-symbol signature checks.
- `registry.py`: preflight checks and before/after installation evidence.

Training imports the public-reference providers from this package; the plugin
does not import the training experiment. Neither direction may import a sibling
experiment. Core UniRL does not import this plugin.

## Gotchas

- Registration is process-local and idempotent. After installation, changing the
enable flag does not undo CUDA overrides: start a new process instead.
A failed installation also requires a restart, because mutation may be partial.
- Preflight checks detect known upstream signature and provider drift. They are
not proof of numerical equality; rerun both verification phases after upgrades.
- The custom `o_proj` changes row sharding to column sharding. Generic vLLM
meta-staged reload can infer the wrong layout despite equal element counts.
Its dedicated subclass therefore bypasses meta staging and loads the full
canonical weight directly into the existing CuMem allocation.
- Do not wrap the expert Parameter loaders: vLLM introspects them during
layerwise reload. Derived MoE columns are invalidated at module load and before
worker sleep so cached allocations never survive weight replacement.
- Keep the TP routing agreement checks before route-dependent collectives.
Removing them can turn a routing mismatch into a distributed hang.
25 changes: 25 additions & 0 deletions experimental/train_inference_parity/vllm_plugin/pyproject.toml
Original file line number Diff line number Diff line change
@@ -0,0 +1,25 @@
[build-system]
requires = ["setuptools>=61", "wheel"]
build-backend = "setuptools.build_meta"

[project]
name = "unirl-train-inference-parity-vllm"
version = "0.1.0"
description = "Experimental UniRL train/inference parity plugin for vLLM"
readme = "README.md"
requires-python = ">=3.12,<3.14"
license = { text = "Apache-2.0" }

[project.optional-dependencies]
verified = [
"vllm==0.27.0",
"torch==2.13.0",
"transformers==5.6.0",
]

[project.entry-points."vllm.general_plugins"]
unirl_train_inference_parity = "unirl_train_inference_parity_vllm:register"

[tool.setuptools.packages.find]
where = ["src"]
include = ["unirl_train_inference_parity_vllm*"]
Original file line number Diff line number Diff line change
@@ -0,0 +1,91 @@
"""vLLM general-plugin entry point for UniRL train/inference parity."""

from __future__ import annotations

import os
from copy import deepcopy

_RUNTIME_MANIFEST: dict[str, object] | None = None
_REGISTRATION_FAILED = False


def _install_capability_manifest_bridge() -> None:
"""Attach this worker's plugin manifest to UniRL's generic capability RPC."""
from unirl.rollout.engine.vllm.worker_extension import UniRLWeightSyncExtension

original = UniRLWeightSyncExtension.unirl_weight_sync_capabilities
if getattr(original, "_unirl_parity_manifest_bridge", False):
return

def with_parity_manifest(self):
capabilities = dict(original(self))
manifest = runtime_manifest()
if manifest is None:
raise RuntimeError("parity plugin manifest is unavailable after registration")
capabilities["train_inference_parity"] = manifest
return capabilities

with_parity_manifest._unirl_parity_manifest_bridge = True
UniRLWeightSyncExtension.unirl_weight_sync_capabilities = with_parity_manifest


def register() -> tuple[str, ...]:
"""Install once per dedicated process; changing the opt-in requires a restart."""
global _RUNTIME_MANIFEST, _REGISTRATION_FAILED
if _REGISTRATION_FAILED:
raise RuntimeError("parity plugin registration previously failed; restart this process")
if os.environ.get("UNIRL_PARITY_ENABLE") != "1":
if _RUNTIME_MANIFEST is not None:
raise RuntimeError("parity overrides cannot be disabled in-process; restart without the opt-in")
return ()

import json

from .compat import validate_runtime_versions
from .config import load_config
from .registry import Installer, install_selected

config = load_config()
if _RUNTIME_MANIFEST is not None:
return config.patches
versions = validate_runtime_versions()
from .common import install_common, preflight_common
from .models.qwen3_moe_30b_a3b import (
install_qwen3_moe_patch,
preflight_qwen3_moe_patch,
)

installers = {
"common": Installer(preflight=preflight_common, install=install_common),
"qwen3_moe_30b_a3b": Installer(
preflight=preflight_qwen3_moe_patch,
install=install_qwen3_moe_patch,
),
}
_REGISTRATION_FAILED = True
installed = install_selected(config.patches, strict=config.strict, installers=installers)
manifest = {
"entrypoint": "unirl_train_inference_parity",
"model": config.model,
"patches": [result.to_manifest() for result in installed],
"pid": os.getpid(),
"profile": config.profile,
"runtime": [result.to_manifest() for result in versions],
"source_file": __file__,
"strict": config.strict,
}
_RUNTIME_MANIFEST = manifest
_install_capability_manifest_bridge()
_REGISTRATION_FAILED = False
print(
"[unirl.parity.vllm] manifest=" + json.dumps(manifest, sort_keys=True, separators=(",", ":")),
flush=True,
)
return tuple(result.name for result in installed)


def runtime_manifest() -> dict[str, object] | None:
return deepcopy(_RUNTIME_MANIFEST)


__all__ = ["register", "runtime_manifest"]
Original file line number Diff line number Diff line change
@@ -0,0 +1,36 @@
"""Model-independent vLLM parity patches."""

from ..registry import PatchResult
from .attention import install_attention_contract, preflight_attention_contract
from .moe_combine import moe_combine
from .norm import install_norm_patch, preflight_norm_patch
from .precision import install_precision_contract, preflight_precision_contract
from .providers import preflight_providers
from .reductions import install_reduction_patch, preflight_reduction_patch


def preflight_common(strict: bool) -> None:
preflight_providers(strict=strict)
preflight_precision_contract(strict=strict)
preflight_reduction_patch(strict=strict)
preflight_norm_patch(strict=strict)
preflight_attention_contract(strict=strict)


def install_common(*, strict: bool) -> PatchResult:
if not strict:
raise ValueError("the public-reference common installer requires strict=True")
symbols = (
*install_precision_contract(),
*install_reduction_patch(),
*install_norm_patch(),
*install_attention_contract(),
)
return PatchResult(name="common", symbols=symbols)


__all__ = [
"install_common",
"moe_combine",
"preflight_common",
]
Original file line number Diff line number Diff line change
@@ -0,0 +1,152 @@
"""Activate only vLLM's attention batch-invariance gates."""

from __future__ import annotations

import importlib
import types

from ..compat import require_symbol, require_value
from ..registry import SymbolResult, symbol_result, value_result

_INSTALLED = False
_ENV_MODULES = (
"vllm.model_executor.layers.attention.attention",
"vllm.v1.attention.backends.flash_attn",
"vllm.v1.attention.backends.fa_utils",
)
_TRITON_MODULE = "vllm.v1.attention.ops.triton_unified_attention"
_FLASH_MODULE = "vllm.v1.attention.backends.flash_attn"


class _EnvProxy:
def __init__(self, original):
self._original = original

def __getattr__(self, name):
if name == "VLLM_BATCH_INVARIANT":
return True
return getattr(self._original, name)


def _probe_fa3() -> tuple[object, type]:
backend = require_symbol(
_FLASH_MODULE,
"FlashAttentionBackend",
origin=f"{_FLASH_MODULE}.FlashAttentionBackend",
strict=True,
)
get_version = require_symbol(
_FLASH_MODULE,
"get_flash_attn_version",
parameters=(
("requires_alibi", "False"),
("head_size", "None"),
("head_size_v", "None"),
("has_sinks", "False"),
("requires_local_attention", "False"),
),
origin="vllm.v1.attention.backends.fa_utils.get_flash_attn_version",
strict=True,
)
if get_version() != 3:
raise RuntimeError(f"parity requires vLLM FlashAttention 3, got {get_version()!r}")
if backend.get_name() != "FLASH_ATTN" or backend.supports_batch_invariance() is not True:
raise RuntimeError("vLLM FlashAttention backend does not satisfy the batch-invariant contract")
implementation = backend.get_impl_cls()
if (
getattr(implementation, "__module__", None) != _FLASH_MODULE
or getattr(implementation, "__name__", None) != "FlashAttentionImpl"
):
raise RuntimeError(f"unexpected vLLM FlashAttention implementation {implementation!r}")
facade = importlib.import_module("vllm.vllm_flash_attn")
varlen = getattr(facade, "flash_attn_varlen_func", None)
if not callable(varlen) or getattr(varlen, "__module__", None) != ("vllm.vllm_flash_attn.flash_attn_interface"):
raise RuntimeError(f"unexpected vLLM FA3 varlen provider {varlen!r}")
return varlen, implementation


def preflight_attention_contract(*, strict: bool) -> None:
if not strict:
raise ValueError("the parity attention contract requires strict=True")
for module_name in _ENV_MODULES:
envs = require_value(
module_name,
"envs",
expected_type=types.ModuleType,
)
if envs.__name__ != "vllm.envs":
raise RuntimeError(
f"conflicting attention env provider at {module_name}.envs: expected vllm.envs, got {envs.__name__}"
)
require_value(
module_name,
"envs.VLLM_BATCH_INVARIANT",
expected_type=bool,
)
require_value(
_TRITON_MODULE,
"is_batch_invariant",
expected_type=bool,
)
_probe_fa3()


def install_attention_contract() -> tuple[SymbolResult, ...]:
global _INSTALLED
if _INSTALLED:
raise RuntimeError("attention parity contract installed twice")
fa3_varlen, fa3_implementation = _probe_fa3()
results = []
for module_name in _ENV_MODULES:
module = importlib.import_module(module_name)
original = module.envs
proxy = _EnvProxy(original)
module.envs = proxy
results.append(
symbol_result(
f"{module_name}.envs",
proxy,
before=original,
actual=module.envs,
)
)
triton_module = importlib.import_module(_TRITON_MODULE)
original_batch_invariant = triton_module.is_batch_invariant
triton_module.is_batch_invariant = True
results.append(
value_result(
f"{_TRITON_MODULE}.is_batch_invariant",
"literal:true",
before=original_batch_invariant,
actual=triton_module.is_batch_invariant,
verified=triton_module.is_batch_invariant is True,
)
)
results.extend(
(
symbol_result(
"vllm.vllm_flash_attn.flash_attn_varlen_func",
fa3_varlen,
before=fa3_varlen,
actual=fa3_varlen,
),
symbol_result(
f"{_FLASH_MODULE}.FlashAttentionImpl",
fa3_implementation,
before=fa3_implementation,
actual=fa3_implementation,
),
value_result(
f"{_FLASH_MODULE}.get_flash_attn_version()",
"literal:3",
before=3,
actual=3,
verified=True,
),
)
)
_INSTALLED = True
return tuple(results)


__all__ = ["install_attention_contract", "preflight_attention_contract"]
Original file line number Diff line number Diff line change
@@ -0,0 +1,31 @@
"""Shared correctness-first MoE output combination."""

from __future__ import annotations

import torch


def moe_combine(
contributions: torch.Tensor,
row_map: torch.Tensor,
) -> torch.Tensor:
"""Combine top-k expert rows with the parity rounding contract."""
if contributions.dim() != 2 or row_map.dim() != 2:
raise ValueError("moe_combine requires 2D contributions and row_map tensors")
tokens, topk = row_map.shape
hidden = int(contributions.shape[-1])
total = torch.zeros(
(tokens, hidden),
dtype=torch.float32,
device=contributions.device,
)
zero = torch.zeros_like(total)
for slot in range(topk):
rows = row_map[:, slot].long()
contribution = contributions.index_select(0, rows.clamp_min(0)).float()
total = total + torch.where((rows >= 0)[:, None], contribution, zero)
total = total.to(torch.bfloat16).float()
return total.to(torch.bfloat16)


__all__ = ["moe_combine"]
Loading
Loading