diff --git a/experimental/train_inference_parity/vllm_plugin/README.md b/experimental/train_inference_parity/vllm_plugin/README.md new file mode 100644 index 000000000..0ff965969 --- /dev/null +++ b/experimental/train_inference_parity/vllm_plugin/README.md @@ -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. diff --git a/experimental/train_inference_parity/vllm_plugin/pyproject.toml b/experimental/train_inference_parity/vllm_plugin/pyproject.toml new file mode 100644 index 000000000..cf354d298 --- /dev/null +++ b/experimental/train_inference_parity/vllm_plugin/pyproject.toml @@ -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*"] diff --git a/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/__init__.py b/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/__init__.py new file mode 100644 index 000000000..445710129 --- /dev/null +++ b/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/__init__.py @@ -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"] diff --git a/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/common/__init__.py b/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/common/__init__.py new file mode 100644 index 000000000..a7c6f5a2d --- /dev/null +++ b/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/common/__init__.py @@ -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", +] diff --git a/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/common/attention.py b/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/common/attention.py new file mode 100644 index 000000000..4087ee646 --- /dev/null +++ b/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/common/attention.py @@ -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"] diff --git a/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/common/moe_combine.py b/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/common/moe_combine.py new file mode 100644 index 000000000..2a6bfee4c --- /dev/null +++ b/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/common/moe_combine.py @@ -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"] diff --git a/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/common/norm.py b/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/common/norm.py new file mode 100644 index 000000000..17df9cd7f --- /dev/null +++ b/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/common/norm.py @@ -0,0 +1,52 @@ +"""Route vLLM RMSNorm layers to the public BI primitive.""" + +from __future__ import annotations + +from ..compat import require_symbol +from ..registry import SymbolResult, symbol_result + +_INSTALLED = False + + +def preflight_norm_patch(*, strict: bool) -> None: + require_symbol( + "vllm.model_executor.layers.layernorm", + "RMSNorm.forward_cuda", + parameters=(("self", ""), ("x", ""), ("residual", "None")), + origin="vllm.model_executor.layers.layernorm.RMSNorm.forward_cuda", + strict=strict, + ) + + +def install_norm_patch() -> tuple[SymbolResult, ...]: + global _INSTALLED + if _INSTALLED: + raise RuntimeError("RMSNorm parity patch installed twice") + from vllm.model_executor.layers.layernorm import RMSNorm + + from .providers import rms_norm + + original = RMSNorm.forward_cuda + + def forward_cuda(self, x, residual=None): + if self.variance_size_override is not None: + raise RuntimeError("parity RMSNorm does not support variance_size_override") + if residual is not None: + added = x + residual + return rms_norm(added, self.weight.data, self.variance_epsilon), added + return rms_norm(x, self.weight.data, self.variance_epsilon) + + forward_cuda._unirl_parity_original = original + RMSNorm.forward_cuda = forward_cuda + _INSTALLED = True + return ( + symbol_result( + "vllm.model_executor.layers.layernorm.RMSNorm.forward_cuda", + forward_cuda, + before=original, + actual=RMSNorm.forward_cuda, + ), + ) + + +__all__ = ["install_norm_patch", "preflight_norm_patch"] diff --git a/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/common/precision.py b/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/common/precision.py new file mode 100644 index 000000000..cd732578d --- /dev/null +++ b/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/common/precision.py @@ -0,0 +1,100 @@ +"""Process-global precision settings without enabling vLLM's global BI mode.""" + +from __future__ import annotations + +from ..compat import require_symbol, require_value +from ..registry import SymbolResult, value_result + +_INSTALLED = False + + +def preflight_precision_contract(*, strict: bool) -> None: + for symbol in ( + "matmul.fp32_precision", + "matmul.allow_bf16_reduced_precision_reduction", + "matmul.allow_fp16_reduced_precision_reduction", + ): + require_value( + "torch.backends.cuda", + symbol, + expected_type=str if symbol.endswith("fp32_precision") else bool, + ) + for symbol in ("conv.fp32_precision", "rnn.fp32_precision"): + require_value("torch.backends.cudnn", symbol, expected_type=str) + require_symbol( + "torch.backends.cuda", + "preferred_blas_library", + parameters=(("backend", "None"),), + origin="torch.backends.cuda.preferred_blas_library", + strict=strict, + ) + + +def install_precision_contract() -> tuple[SymbolResult, ...]: + global _INSTALLED + if _INSTALLED: + raise RuntimeError("precision parity contract installed twice") + import torch + + before = { + "matmul_fp32": torch.backends.cuda.matmul.fp32_precision, + "conv_fp32": torch.backends.cudnn.conv.fp32_precision, + "rnn_fp32": torch.backends.cudnn.rnn.fp32_precision, + "bf16_reduction": torch.backends.cuda.matmul.allow_bf16_reduced_precision_reduction, + "fp16_reduction": torch.backends.cuda.matmul.allow_fp16_reduced_precision_reduction, + "blas": torch.backends.cuda.preferred_blas_library(), + } + torch.backends.cuda.matmul.fp32_precision = "ieee" + torch.backends.cudnn.conv.fp32_precision = "ieee" + torch.backends.cudnn.rnn.fp32_precision = "ieee" + torch.backends.cuda.matmul.allow_bf16_reduced_precision_reduction = False + torch.backends.cuda.matmul.allow_fp16_reduced_precision_reduction = False + torch.backends.cuda.preferred_blas_library(backend="cublaslt") + _INSTALLED = True + return ( + value_result( + "torch.backends.cuda.matmul.fp32_precision", + "literal:'ieee'", + before=before["matmul_fp32"], + actual=torch.backends.cuda.matmul.fp32_precision, + verified=torch.backends.cuda.matmul.fp32_precision == "ieee", + ), + value_result( + "torch.backends.cudnn.conv.fp32_precision", + "literal:'ieee'", + before=before["conv_fp32"], + actual=torch.backends.cudnn.conv.fp32_precision, + verified=torch.backends.cudnn.conv.fp32_precision == "ieee", + ), + value_result( + "torch.backends.cudnn.rnn.fp32_precision", + "literal:'ieee'", + before=before["rnn_fp32"], + actual=torch.backends.cudnn.rnn.fp32_precision, + verified=torch.backends.cudnn.rnn.fp32_precision == "ieee", + ), + value_result( + "torch.backends.cuda.matmul.allow_bf16_reduced_precision_reduction", + "literal:false", + before=before["bf16_reduction"], + actual=torch.backends.cuda.matmul.allow_bf16_reduced_precision_reduction, + verified=not torch.backends.cuda.matmul.allow_bf16_reduced_precision_reduction, + ), + value_result( + "torch.backends.cuda.matmul.allow_fp16_reduced_precision_reduction", + "literal:false", + before=before["fp16_reduction"], + actual=torch.backends.cuda.matmul.allow_fp16_reduced_precision_reduction, + verified=not torch.backends.cuda.matmul.allow_fp16_reduced_precision_reduction, + ), + value_result( + "torch.backends.cuda.preferred_blas_library", + "torch._C._BlasBackend.Cublaslt", + before=before["blas"], + actual=torch.backends.cuda.preferred_blas_library(), + verified=torch.backends.cuda.preferred_blas_library().name == "Cublaslt", + ), + ) + + +__all__ = ["install_precision_contract", "preflight_precision_contract"] diff --git a/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/common/providers.py b/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/common/providers.py new file mode 100644 index 000000000..aba737e06 --- /dev/null +++ b/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/common/providers.py @@ -0,0 +1,84 @@ +"""Public vLLM/Torch providers used by the reference profile.""" + +from __future__ import annotations + +import torch + +from ..compat import require_symbol + + +def preflight_providers(*, strict: bool) -> None: + module = "vllm.model_executor.layers.batch_invariant" + contracts = { + "linear_batch_invariant": ( + ("input", ""), + ("weight", ""), + ("bias", "None"), + ), + "rms_norm_batch_invariant": ( + ("input", ""), + ("weight", ""), + ("eps", "1e-06"), + ("residual", "None"), + ), + "log_softmax": (("input", ""), ("dim", "-1")), + "softmax_batch_invariant": ( + ("input", ""), + ("dim", ""), + ("dtype", "None"), + ), + "mean_batch_invariant": ( + ("input", ""), + ("dim", ""), + ("keepdim", "False"), + ("dtype", "None"), + ), + } + for name, parameters in contracts.items(): + require_symbol( + module, + name, + parameters=parameters, + origin=f"{module}.{name}", + strict=strict, + ) + + +def linear(input: torch.Tensor, weight: torch.Tensor, bias=None) -> torch.Tensor: + from vllm.model_executor.layers.batch_invariant import linear_batch_invariant + + return linear_batch_invariant(input.contiguous(), weight.contiguous(), bias) + + +def rms_norm(input: torch.Tensor, weight: torch.Tensor, eps: float) -> torch.Tensor: + from vllm.model_executor.layers.batch_invariant import rms_norm_batch_invariant + + return rms_norm_batch_invariant(input.contiguous(), weight.contiguous(), eps) + + +def log_softmax(input: torch.Tensor, dim: int = -1) -> torch.Tensor: + from vllm.model_executor.layers.batch_invariant import log_softmax as implementation + + return implementation(input, dim=dim) + + +def softmax(input: torch.Tensor, dim: int = -1, dtype=None) -> torch.Tensor: + from vllm.model_executor.layers.batch_invariant import softmax_batch_invariant + + return softmax_batch_invariant(input, dim=dim, dtype=dtype) + + +def mean(input: torch.Tensor, dim, keepdim=False, dtype=None) -> torch.Tensor: + from vllm.model_executor.layers.batch_invariant import mean_batch_invariant + + return mean_batch_invariant(input, dim=dim, keepdim=keepdim, dtype=dtype) + + +__all__ = [ + "linear", + "log_softmax", + "mean", + "preflight_providers", + "rms_norm", + "softmax", +] diff --git a/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/common/reductions.py b/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/common/reductions.py new file mode 100644 index 000000000..478f220ab --- /dev/null +++ b/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/common/reductions.py @@ -0,0 +1,117 @@ +"""Install vLLM public batch-invariant reductions as CUDA ATen providers.""" + +from __future__ import annotations + +import inspect + +from ..compat import require_symbol +from ..registry import SymbolResult, symbol_result + +_LIBRARY = None +_NATIVE_CUDA_OPS = ( + "aten::_log_softmax", + "aten::softmax", + "aten::_softmax", + "aten::mean.dim", +) + + +def _cuda_registration(torch, operator: str) -> str | None: + return next( + (line for line in torch._C._dispatch_dump_table(operator).splitlines() if line.startswith("CUDA:")), + None, + ) + + +def preflight_reduction_patch(*, strict: bool) -> None: + import torch + + require_symbol( + "vllm.v1.worker.gpu.sample.logprob", + "compute_token_logprobs", + parameters=(("logits", ""), ("token_ids", "")), + origin=("vllm.v1.worker.gpu.sample.logprob.compute_token_logprobs"), + strict=strict, + ) + if strict: + conflicts = { + operator: registration + for operator in _NATIVE_CUDA_OPS + if (registration := _cuda_registration(torch, operator)) is not None and "/pytorch/" not in registration + } + if conflicts: + raise RuntimeError(f"unexpected pre-existing CUDA reduction providers: {conflicts}") + + +def install_reduction_patch() -> tuple[SymbolResult, ...]: + global _LIBRARY + if _LIBRARY is not None: + raise RuntimeError("reduction parity patch installed twice") + import torch + + from .providers import log_softmax, mean, softmax + + before_registrations = {operator: _cuda_registration(torch, operator) for operator in _NATIVE_CUDA_OPS} + import vllm.v1.worker.gpu.sample.logprob as sample_logprob + + original_compute_token_logprobs = sample_logprob.compute_token_logprobs + + def aten_log_softmax(input, dim, half_to_float): + return log_softmax(input.float() if half_to_float else input, dim=dim) + + library = torch.library.Library("aten", "IMPL") + library.impl("aten::_log_softmax", aten_log_softmax, "CUDA") + library.impl("aten::softmax", softmax, "CUDA") + + def aten_softmax(input, dim, half_to_float): + return softmax(input.float() if half_to_float else input, dim=dim) + + library.impl("aten::_softmax", aten_softmax, "CUDA") + library.impl("aten::mean.dim", mean, "CUDA") + + def compute_token_logprobs(logits, token_ids): + return log_softmax(logits.float(), dim=-1).gather( + -1, + token_ids.to(torch.int64), + ) + + sample_logprob.compute_token_logprobs = compute_token_logprobs + _LIBRARY = library + registered = ( + ("aten::_log_softmax[CUDA]", aten_log_softmax, "aten::_log_softmax"), + ("aten::softmax[CUDA]", softmax, "aten::softmax"), + ("aten::_softmax[CUDA]", aten_softmax, "aten::_softmax"), + ("aten::mean.dim[CUDA]", mean, "aten::mean.dim"), + ) + results = [] + for symbol, provider, operator in registered: + registration = _cuda_registration(torch, operator) + results.append( + SymbolResult( + symbol=symbol, + before_provider=before_registrations[operator] or "", + before_signature=None, + before_identity=None, + after_provider=registration or "", + after_signature=str(inspect.signature(provider)), + after_identity=id(provider), + verified=( + registration is not None + and "unirl_train_inference_parity_vllm/common/reductions.py" in registration + ), + ) + ) + if not all(result.verified for result in results): + raise RuntimeError(f"failed to verify CUDA reduction providers: {results}") + results.append( + symbol_result( + "vllm.v1.worker.gpu.sample.logprob.compute_token_logprobs", + compute_token_logprobs, + before=original_compute_token_logprobs, + actual=sample_logprob.compute_token_logprobs, + ) + ) + return tuple(results) + + +__all__ = ["install_reduction_patch", "preflight_reduction_patch"] diff --git a/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/compat.py b/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/compat.py new file mode 100644 index 000000000..1d1f20898 --- /dev/null +++ b/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/compat.py @@ -0,0 +1,131 @@ +"""Pinned runtime and private-symbol contracts for the verified stack.""" + +from __future__ import annotations + +import importlib +import inspect +from dataclasses import asdict, dataclass +from importlib.metadata import PackageNotFoundError, version + +_VERSION_ALLOWLIST = { + "vllm": ("0.27.0",), + "torch": ("2.13.0",), + "transformers": ("5.6.0",), +} +_REQUIRED = "" + + +@dataclass(frozen=True) +class RuntimeVersion: + package: str + installed: str + allowed: tuple[str, ...] + + def to_manifest(self) -> dict[str, object]: + return asdict(self) + + +def validate_runtime_versions() -> tuple[RuntimeVersion, ...]: + results = [] + failures = [] + for package, allowed in _VERSION_ALLOWLIST.items(): + try: + installed = version(package) + except PackageNotFoundError: + failures.append(f"{package}=; allowed={list(allowed)}") + continue + public_version = installed.split("+", 1)[0] + if public_version not in allowed: + failures.append(f"{package}={installed}; allowed={list(allowed)}") + results.append(RuntimeVersion(package=package, installed=installed, allowed=allowed)) + if failures: + raise RuntimeError("unsupported UniRL parity runtime: " + "; ".join(failures)) + return tuple(results) + + +def _signature_shape( + value: object, +) -> tuple[tuple[str, str, str], ...]: + parameters = inspect.signature(value).parameters.values() + return tuple( + ( + parameter.name, + parameter.kind.name, + _REQUIRED if parameter.default is inspect.Parameter.empty else repr(parameter.default), + ) + for parameter in parameters + ) + + +def require_symbol( + module_name: str, + symbol_path: str, + *, + parameters: tuple[tuple[str, str] | tuple[str, str, str], ...] | None = None, + origin: str | None = None, + strict: bool, +) -> object: + """Resolve a required symbol and optionally enforce its exact signature.""" + try: + value: object = importlib.import_module(module_name) + except ImportError as error: + raise RuntimeError(f"required parity module {module_name!r} could not be imported") from error + for component in symbol_path.split("."): + try: + value = getattr(value, component) + except AttributeError as error: + raise RuntimeError(f"required parity symbol {module_name}.{symbol_path} is missing") from error + + qualified_name = ".".join( + part + for part in ( + getattr(value, "__module__", None), + getattr(value, "__qualname__", None), + ) + if part + ) + if origin is not None and qualified_name != origin: + raise RuntimeError( + f"conflicting provider at {module_name}.{symbol_path}: " + f"expected {origin}, got {qualified_name or type(value).__name__}" + ) + + if parameters is not None: + try: + actual = _signature_shape(value) + except (TypeError, ValueError) as error: + raise RuntimeError(f"cannot inspect required parity symbol {module_name}.{symbol_path}") from error + expected = tuple( + ( + parameter[0], + (inspect.Parameter.POSITIONAL_OR_KEYWORD.name if len(parameter) == 2 else parameter[1]), + parameter[-1], + ) + for parameter in parameters + ) + if strict and actual != expected: + raise RuntimeError(f"signature drift at {module_name}.{symbol_path}: expected {expected}, got {actual}") + return value + + +def require_value( + module_name: str, + symbol_path: str, + *, + expected_type: type, +) -> object: + value = require_symbol(module_name, symbol_path, strict=True) + if not isinstance(value, expected_type): + raise RuntimeError( + f"invalid parity symbol {module_name}.{symbol_path}: " + f"expected {expected_type.__name__}, got {type(value).__name__}" + ) + return value + + +__all__ = [ + "RuntimeVersion", + "require_symbol", + "require_value", + "validate_runtime_versions", +] diff --git a/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/config.py b/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/config.py new file mode 100644 index 000000000..eb20b4c4c --- /dev/null +++ b/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/config.py @@ -0,0 +1,60 @@ +"""Environment-only configuration available in every spawned vLLM process.""" + +from __future__ import annotations + +import os +from dataclasses import dataclass + +PUBLIC_REFERENCE = "public_reference" +QWEN3_MOE_MODEL = "qwen3_moe_30b_a3b" +REQUIRED_PATCHES = ("common", QWEN3_MOE_MODEL) + + +@dataclass(frozen=True) +class PluginConfig: + profile: str + model: str + patches: tuple[str, ...] + strict: bool + + +def _required_environment(name: str) -> str: + try: + value = os.environ[name].strip() + except KeyError as error: + raise RuntimeError(f"{name} is required when UNIRL_PARITY_ENABLE=1") from error + if not value: + raise RuntimeError(f"{name} must be non-empty when UNIRL_PARITY_ENABLE=1") + return value + + +def load_config() -> PluginConfig: + profile = _required_environment("UNIRL_PARITY_PROFILE").lower() + if profile != PUBLIC_REFERENCE: + raise ValueError(f"unknown UNIRL_PARITY_PROFILE={profile!r}") + model = _required_environment("UNIRL_PARITY_MODEL").lower() + if model != QWEN3_MOE_MODEL: + raise ValueError(f"unsupported UNIRL_PARITY_MODEL={model!r}") + raw_patches = _required_environment("UNIRL_PARITY_PATCHES") + patch_values = tuple(value.strip() for value in raw_patches.split(",")) + if any(not value for value in patch_values) or len(patch_values) != len(set(patch_values)): + raise ValueError("UNIRL_PARITY_PATCHES must contain unique, non-empty patch names") + if patch_values != REQUIRED_PATCHES: + raise ValueError( + f"UNIRL_PARITY_PATCHES must enable the complete ordered contract {REQUIRED_PATCHES!r}; got {patch_values!r}" + ) + patches = patch_values + raw_strict = _required_environment("UNIRL_PARITY_STRICT") + if raw_strict != "1": + raise ValueError("UNIRL_PARITY_STRICT must be exactly '1' for the public-reference contract") + strict = True + return PluginConfig(profile=profile, model=model, patches=patches, strict=strict) + + +__all__ = [ + "PUBLIC_REFERENCE", + "QWEN3_MOE_MODEL", + "REQUIRED_PATCHES", + "PluginConfig", + "load_config", +] diff --git a/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/models/__init__.py b/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/models/__init__.py new file mode 100644 index 000000000..455439daf --- /dev/null +++ b/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/models/__init__.py @@ -0,0 +1 @@ +"""Model-specific vLLM parity patches.""" diff --git a/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/models/qwen3_moe_30b_a3b/__init__.py b/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/models/qwen3_moe_30b_a3b/__init__.py new file mode 100644 index 000000000..8eeb19b76 --- /dev/null +++ b/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/models/qwen3_moe_30b_a3b/__init__.py @@ -0,0 +1,5 @@ +"""Qwen3-30B-A3B vLLM model patch.""" + +from .patch import install_qwen3_moe_patch, preflight_qwen3_moe_patch + +__all__ = ["install_qwen3_moe_patch", "preflight_qwen3_moe_patch"] diff --git a/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/models/qwen3_moe_30b_a3b/dense.py b/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/models/qwen3_moe_30b_a3b/dense.py new file mode 100644 index 000000000..c74bd327b --- /dev/null +++ b/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/models/qwen3_moe_30b_a3b/dense.py @@ -0,0 +1,119 @@ +"""Public-reference QKV, o_proj and LM-head providers.""" + +from __future__ import annotations + +import types + +import torch +import torch.nn as nn + +from ...common.providers import linear + +_PARITY_O_PROJECTION_CLASSES: dict[type, type] = {} + + +def _use_direct_layerwise_reload(layer) -> None: + """Load the column-sharded o-proj without meta staging; see the plugin README.md.""" + from vllm.model_executor.model_loader.reload.meta import SKIP_MODULES + + base_class = type(layer) + parity_class = _PARITY_O_PROJECTION_CLASSES.get(base_class) + if parity_class is None: + parity_class = type("UniRLParityOProjection", (base_class,), {}) + _PARITY_O_PROJECTION_CLASSES[base_class] = parity_class + layer.__class__ = parity_class + SKIP_MODULES.add(parity_class.__name__) + + +def _column_weight_loader(self, parameter, loaded_weight): + rank = int(self._unirl_parity_tp_rank) + world = int(self._unirl_parity_tp_world) + rows = int(loaded_weight.shape[0]) + if rows % world or tuple(loaded_weight.shape) != (self.output_size, self.input_size): + raise ValueError(f"invalid full o-proj weight shape: {tuple(loaded_weight.shape)} for TP{world}") + shard = rows // world + loaded_shard = loaded_weight.narrow(0, rank * shard, shard).to( + device=parameter.device, + dtype=parameter.dtype, + ) + parameter.data.copy_(loaded_shard) + + +def _reshape_row_to_column(layer) -> None: + from vllm.distributed import ( + get_tensor_model_parallel_rank, + get_tensor_model_parallel_world_size, + ) + + rank = get_tensor_model_parallel_rank() + world = get_tensor_model_parallel_world_size() + _use_direct_layerwise_reload(layer) + old = layer.weight + replacement = nn.Parameter( + torch.empty( + layer.output_size // world, + layer.input_size, + device=old.device, + dtype=old.dtype, + ), + requires_grad=False, + ) + from vllm.model_executor.utils import set_weight_attrs + + layer._unirl_parity_tp_rank = int(rank) + layer._unirl_parity_tp_world = int(world) + loader = types.MethodType(_column_weight_loader, layer) + set_weight_attrs(replacement, {"weight_loader": loader}) + layer.weight = replacement + layer._unirl_parity_output_size = layer.output_size + + +def _qkv_forward(self, hidden): + return linear(hidden, self.weight, None), None + + +def _o_forward(self, hidden): + from vllm.distributed import get_tp_group + + group = get_tp_group() + full_input = group.all_gather(hidden.contiguous(), dim=-1) + local_output = linear(full_input, self.weight, None) + output = local_output if group.world_size == 1 else group.all_gather(local_output.contiguous(), dim=-1) + return output, None + + +def make_attention(base_class): + class ParityAttention(base_class): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.qkv_proj.forward = types.MethodType( + _qkv_forward, + self.qkv_proj, + ) + _reshape_row_to_column(self.o_proj) + self.o_proj.forward = types.MethodType( + _o_forward, + self.o_proj, + ) + + ParityAttention.__name__ = "ParityQwen3MoeAttention" + return ParityAttention + + +def make_logits_processor(base_class): + class ParityLogitsProcessor(base_class): + def _get_logits(self, hidden_states, lm_head, embedding_bias=None): + if embedding_bias is not None: + raise ValueError("Qwen3-MoE parity LM head does not support bias") + from vllm.distributed import get_tp_group + + group = get_tp_group() + local_logits = linear(hidden_states, lm_head.weight, None) + logits = local_logits if group.world_size == 1 else group.all_gather(local_logits.contiguous(), dim=-1) + return logits[..., : self.org_vocab_size] + + ParityLogitsProcessor.__name__ = "ParityQwen3MoeLogitsProcessor" + return ParityLogitsProcessor + + +__all__ = ["make_attention", "make_logits_processor"] diff --git a/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/models/qwen3_moe_30b_a3b/moe.py b/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/models/qwen3_moe_30b_a3b/moe.py new file mode 100644 index 000000000..74bd55f02 --- /dev/null +++ b/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/models/qwen3_moe_30b_a3b/moe.py @@ -0,0 +1,252 @@ +"""Correctness-first Qwen3-MoE path using vLLM BI linear and Torch combine.""" + +from __future__ import annotations + +import types + +import torch +import torch.nn.functional as F + +from ...common.moe_combine import moe_combine +from ...common.providers import linear +from .router import install_gate, install_hf_router + + +def _routed_experts(experts): + return getattr(experts, "routed_experts", experts) + + +def _validate_tp_routing( + counts: torch.Tensor, + topk_ids: torch.Tensor, + *, + num_experts: int, + group, + world: int, +) -> None: + """Collectively fail before route-dependent expert collectives diverge.""" + shape = tuple(topk_ids.shape) + counts_valid = bool( + counts.dim() == 1 + and counts.numel() == num_experts + and torch.all(counts >= 0).item() + and int(counts.sum().item()) == int(topk_ids.numel()) + ) + ids_valid = bool( + topk_ids.dim() == 2 + and (topk_ids.numel() == 0 or (int(topk_ids.min().item()) >= 0 and int(topk_ids.max().item()) < num_experts)) + ) + if world == 1: + if not counts_valid or not ids_valid: + raise RuntimeError("Qwen3-MoE parity routing metadata is invalid") + return + metadata = torch.tensor( + ( + topk_ids.dim(), + topk_ids.numel(), + shape[0] if len(shape) > 0 else -1, + shape[1] if len(shape) > 1 else -1, + counts.numel(), + num_experts, + int(counts_valid), + int(ids_valid), + ), + dtype=torch.int64, + device=topk_ids.device, + ) + gathered_metadata = group.all_gather(metadata.contiguous(), dim=0).view( + world, + metadata.numel(), + ) + reference_metadata = gathered_metadata[0] + metadata_matches = bool(torch.all(gathered_metadata == reference_metadata).item()) + metadata_valid = bool( + torch.all(gathered_metadata[:, 0] == 2).item() + and torch.all(gathered_metadata[:, 1] == gathered_metadata[:, 2] * gathered_metadata[:, 3]).item() + and torch.all(gathered_metadata[:, 4] == num_experts).item() + and torch.all(gathered_metadata[:, 5] == num_experts).item() + and torch.all(gathered_metadata[:, 6] == 1).item() + and torch.all(gathered_metadata[:, 7] == 1).item() + ) + if not metadata_matches or not metadata_valid: + observed = gathered_metadata.detach().cpu().tolist() + raise RuntimeError(f"Qwen3-MoE parity TP routing metadata mismatch before expert collectives: {observed}") + + payload = torch.cat( + ( + counts.detach().to(device=topk_ids.device, dtype=torch.int64), + topk_ids.detach().reshape(-1).to(dtype=torch.int64), + ) + ).contiguous() + gathered_payload = group.all_gather(payload, dim=0).view(world, payload.numel()) + mismatch_ranks = ( + torch.any(gathered_payload != gathered_payload[0], dim=1).nonzero().flatten().detach().cpu().tolist() + ) + if mismatch_ranks: + raise RuntimeError( + "Qwen3-MoE parity TP routing mismatch before expert collectives; " + f"counts/topk_ids differ on ranks {mismatch_ranks}" + ) + + +def _get_w2_column(experts) -> torch.Tensor: + cached = getattr(experts, "_unirl_parity_w2_column", None) + if cached is not None: + return cached + from vllm.distributed import ( + get_tensor_model_parallel_rank, + get_tensor_model_parallel_world_size, + get_tp_group, + ) + + world = get_tensor_model_parallel_world_size() + rank = get_tensor_model_parallel_rank() + stock = _routed_experts(experts).w2_weight + full = stock if world == 1 else get_tp_group().all_gather(stock.contiguous(), dim=-1) + hidden_local = int(full.shape[1]) // world + cached = full[:, rank * hidden_local : (rank + 1) * hidden_local].contiguous() + experts._unirl_parity_w2_column = cached + return cached + + +def _install_reload_invalidation(experts) -> None: + original = experts.load_weights + if getattr(original, "_unirl_parity_wrapper", False): + return + + def load_weights(self, weights): + self._unirl_parity_w2_column = None + return original(weights) + + load_weights._unirl_parity_wrapper = True + experts.load_weights = types.MethodType(load_weights, experts) + # Do not wrap w13/w2 Parameter loaders. vLLM's layerwise WTE reload + # introspects and replays those exact loaders; a variadic invalidation + # wrapper caused all 96 expert tensors to be reconstructed incorrectly. + # The module hook above and the worker pre-sleep hook cover cache lifetime. + + +def _expert_forward(block, hidden_states: torch.Tensor) -> torch.Tensor: + from vllm.distributed import ( + get_tensor_model_parallel_world_size, + get_tp_group, + ) + from vllm.model_executor.layers.fused_moe.moe_permute_unpermute import ( + moe_permute, + ) + + input_was_1d = hidden_states.dim() == 1 + hidden = hidden_states.reshape(-1, hidden_states.shape[-1]).contiguous() + router_logits, _ = block.gate(hidden) + topk_weights, topk_ids = block.experts.router.select_experts( + hidden_states=hidden, + router_logits=router_logits, + ) + num_experts = int(block.n_routed_experts) + permuted, _scale, offsets, inverse, _indices = moe_permute( + hidden, + None, + topk_ids.to(torch.int32), + num_experts, + ) + counts_tensor = offsets[1:] - offsets[:-1] + world = get_tensor_model_parallel_world_size() + group = get_tp_group() + _validate_tp_routing( + counts_tensor, + topk_ids, + num_experts=num_experts, + group=group, + world=world, + ) + counts = counts_tensor.detach().cpu().tolist() + w2_column = _get_w2_column(block.experts) + routed_experts = _routed_experts(block.experts) + slots = int(topk_ids.numel()) + if slots == 0: + return torch.zeros_like(hidden) + max_count = max(int(value) for value in counts) + local_gate_width = int(routed_experts.w13_weight.shape[1]) + local_intermediate = local_gate_width // 2 + padded_gate_up = [] + offset = 0 + for expert, count_value in enumerate(counts): + count = int(count_value) + rows = permuted[offset : offset + count].contiguous() + offset += count + if count == 0: + local_gate_up = hidden.new_zeros((max_count, local_gate_width)) + else: + active = linear( + rows, + routed_experts.w13_weight[expert], + None, + ) + padding = hidden.new_zeros((max_count - count, local_gate_width)) + local_gate_up = torch.cat((active, padding), dim=0) + padded_gate_up.append(local_gate_up) + + local_gate_up = torch.stack(padded_gate_up, dim=0).contiguous() + gathered_gate_up = local_gate_up if world == 1 else group.all_gather(local_gate_up, dim=-1) + rank_packed = gathered_gate_up.view( + num_experts, + max_count, + world, + 2, + local_intermediate, + ) + gate = rank_packed[:, :, :, 0].reshape(num_experts, max_count, -1) + up = rank_packed[:, :, :, 1].reshape(num_experts, max_count, -1) + activation = F.silu(gate) * up + + local_hidden_width = int(w2_column.shape[1]) + padded_down = [] + for expert, count_value in enumerate(counts): + if int(count_value) == 0: + local_down = hidden.new_zeros((max_count, local_hidden_width)) + else: + local_down = linear( + activation[expert].contiguous(), + w2_column[expert], + None, + ) + padded_down.append(local_down) + local_down = torch.stack(padded_down, dim=0).contiguous() + gathered_down = local_down if world == 1 else group.all_gather(local_down, dim=-1) + expert_output = torch.cat( + [gathered_down[expert, : int(count)] for expert, count in enumerate(counts) if int(count) > 0], + dim=0, + ) + permuted_weights = torch.zeros( + slots, + dtype=topk_weights.dtype, + device=topk_weights.device, + ) + permuted_weights[inverse.long()] = topk_weights.reshape(-1) + contributions = (expert_output * permuted_weights[:, None]).to(torch.bfloat16) + result = moe_combine( + contributions, + inverse.view_as(topk_ids).to(torch.int32), + ) + return result.squeeze(0) if input_was_1d else result + + +def make_moe_block(base_class): + class ParityMoeBlock(base_class): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + install_gate(self.gate) + install_hf_router(self.experts) + _install_reload_invalidation(self.experts) + + def forward(self, hidden_states): + return _expert_forward(self, hidden_states) + + def _unirl_before_sleep(self): + self.experts._unirl_parity_w2_column = None + + ParityMoeBlock.__name__ = "ParityQwen3MoeSparseMoeBlock" + return ParityMoeBlock + + +__all__ = ["make_moe_block"] diff --git a/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/models/qwen3_moe_30b_a3b/patch.py b/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/models/qwen3_moe_30b_a3b/patch.py new file mode 100644 index 000000000..243919037 --- /dev/null +++ b/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/models/qwen3_moe_30b_a3b/patch.py @@ -0,0 +1,304 @@ +"""Install the public-reference Qwen3-MoE vLLM patch.""" + +from __future__ import annotations + +import importlib + +import torch + +from ...common.providers import preflight_providers +from ...compat import require_symbol +from ...registry import PatchResult, SymbolResult, symbol_result, value_result +from .dense import make_attention, make_logits_processor +from .moe import make_moe_block + +_INSTALLED = False + + +def _install_rope() -> tuple[SymbolResult, ...]: + import vllm.model_executor.layers.rotary_embedding.base as rope + + original_inverse_frequency = rope.RotaryEmbeddingBase._compute_inv_freq + original_forward_cuda = rope.RotaryEmbedding.forward_cuda + + def inverse_frequency(self, base): + exponent = torch.arange(0, self.rotary_dim, 2, dtype=torch.float32) + return 1.0 / (base ** (exponent / self.rotary_dim)) + + def apply(self, positions, query, key=None): + positions = positions.flatten() + cache = self._match_cos_sin_cache_dtype(query).index_select(0, positions) + cosine, sine = cache.chunk(2, dim=-1) + tokens = int(positions.shape[0]) + + def rotate(value): + original_shape = value.shape + value = value.view(tokens, -1, self.head_size) + first, second = torch.chunk(value[..., : self.rotary_dim], 2, dim=-1) + cos = cosine.unsqueeze(-2).to(value.dtype) + sin = sine.unsqueeze(-2).to(value.dtype) + rotated = torch.cat( + (first * cos - second * sin, second * cos + first * sin), + dim=-1, + ) + return torch.cat( + (rotated, value[..., self.rotary_dim :]), + dim=-1, + ).reshape(original_shape) + + return rotate(query), None if key is None else rotate(key) + + rope.RotaryEmbeddingBase._compute_inv_freq = inverse_frequency + rope.RotaryEmbedding.forward_cuda = apply + return ( + symbol_result( + ("vllm.model_executor.layers.rotary_embedding.base.RotaryEmbeddingBase._compute_inv_freq"), + inverse_frequency, + before=original_inverse_frequency, + actual=rope.RotaryEmbeddingBase._compute_inv_freq, + ), + symbol_result( + ("vllm.model_executor.layers.rotary_embedding.base.RotaryEmbedding.forward_cuda"), + apply, + before=original_forward_cuda, + actual=rope.RotaryEmbedding.forward_cuda, + ), + ) + + +def _install_sleep_cache_invalidation() -> SymbolResult: + """Drop parity-owned CUDA caches before vLLM releases their allocations.""" + import vllm.v1.worker.gpu_worker as gpu_worker + + original_sleep = gpu_worker.Worker.sleep + + def sleep(self, level: int = 1) -> None: + model = getattr(getattr(self, "model_runner", None), "model", None) + if model is not None: + for module in model.modules(): + before_sleep = getattr(module, "_unirl_before_sleep", None) + if callable(before_sleep): + before_sleep() + original_sleep(self, level=level) + + gpu_worker.Worker.sleep = sleep + return symbol_result( + "vllm.v1.worker.gpu_worker.Worker.sleep", + sleep, + before=original_sleep, + actual=gpu_worker.Worker.sleep, + ) + + +def preflight_qwen3_moe_patch(strict: bool) -> None: + preflight_providers(strict=strict) + contracts = ( + ( + "vllm.model_executor.layers.rotary_embedding.base", + "RotaryEmbeddingBase._compute_inv_freq", + (("self", ""), ("base", "")), + ("vllm.model_executor.layers.rotary_embedding.base.RotaryEmbeddingBase._compute_inv_freq"), + ), + ( + "vllm.model_executor.layers.rotary_embedding.base", + "RotaryEmbedding.forward_cuda", + ( + ("self", ""), + ("positions", ""), + ("query", ""), + ("key", "None"), + ), + ("vllm.model_executor.layers.rotary_embedding.base.RotaryEmbedding.forward_cuda"), + ), + ( + "vllm.model_executor.models.qwen3_moe", + "Qwen3MoeSparseMoeBlock", + (("vllm_config", ""), ("prefix", "''")), + "vllm.model_executor.models.qwen3_moe.Qwen3MoeSparseMoeBlock", + ), + ( + "vllm.model_executor.models.qwen3_moe", + "Qwen3MoeSparseMoeBlock.forward", + (("self", ""), ("hidden_states", "")), + ("vllm.model_executor.models.qwen3_moe.Qwen3MoeSparseMoeBlock.forward"), + ), + ( + "vllm.model_executor.models.qwen3_moe", + "Qwen3MoeAttention", + ( + ("hidden_size", ""), + ("num_heads", ""), + ("num_kv_heads", ""), + ("rope_parameters", ""), + ("max_position_embeddings", "8192"), + ("head_dim", "None"), + ("rms_norm_eps", "1e-06"), + ("qkv_bias", "False"), + ("cache_config", "None"), + ("quant_config", "None"), + ("prefix", "''"), + ("dual_chunk_attention_config", "None"), + ), + "vllm.model_executor.models.qwen3_moe.Qwen3MoeAttention", + ), + ( + "vllm.model_executor.models.qwen3_moe", + "Qwen3MoeAttention.forward", + ( + ("self", ""), + ("positions", ""), + ("hidden_states", ""), + ), + ("vllm.model_executor.models.qwen3_moe.Qwen3MoeAttention.forward"), + ), + ( + "vllm.model_executor.models.qwen3_moe", + "LogitsProcessor", + ( + ("vocab_size", ""), + ("org_vocab_size", "None"), + ("scale", "1.0"), + ("logits_as_input", "False"), + ("soft_cap", "None"), + ), + "vllm.model_executor.layers.logits_processor.LogitsProcessor", + ), + ( + "vllm.model_executor.models.qwen3_moe", + "LogitsProcessor._get_logits", + ( + ("self", ""), + ("hidden_states", ""), + ("lm_head", ""), + ("embedding_bias", ""), + ), + ("vllm.model_executor.layers.logits_processor.LogitsProcessor._get_logits"), + ), + ( + "vllm.model_executor.layers.fused_moe.moe_permute_unpermute", + "moe_permute", + ( + ("hidden_states", ""), + ("a1q_scale", ""), + ("topk_ids", ""), + ("n_expert", ""), + ("n_local_expert", "-1"), + ("expert_map", "None"), + ("permuted_hidden_states", "None"), + ("scratch", "None"), + ), + ("vllm.model_executor.layers.fused_moe.moe_permute_unpermute.moe_permute"), + ), + ( + "vllm.distributed.parallel_state", + "GroupCoordinator.all_gather", + ( + ("self", ""), + ("input_", ""), + ("dim", "-1"), + ), + "vllm.distributed.parallel_state.GroupCoordinator.all_gather", + ), + ( + "vllm.model_executor.layers.linear", + "ReplicatedLinear.forward", + (("self", ""), ("x", "")), + "vllm.model_executor.layers.linear.ReplicatedLinear.forward", + ), + ( + "vllm.model_executor.layers.fused_moe.router.base_router", + "BaseRouter.select_experts", + ( + ("self", ""), + ("hidden_states", ""), + ("router_logits", ""), + ("topk_indices_dtype", "None"), + ("input_ids", "KEYWORD_ONLY", "None"), + ), + "vllm.model_executor.layers.fused_moe.router.fused_moe_router.FusedMoERouter.select_experts", + ), + ( + "vllm.model_executor.layers.fused_moe.router.fused_topk_router", + "FusedTopKRouter._compute_routing", + ( + ("self", ""), + ("hidden_states", ""), + ("router_logits", ""), + ("indices_type", ""), + ("input_ids", "KEYWORD_ONLY", "None"), + ), + ("vllm.model_executor.layers.fused_moe.router.fused_topk_router.FusedTopKRouter._compute_routing"), + ), + ( + "vllm.v1.worker.gpu_worker", + "Worker.sleep", + (("self", ""), ("level", "1")), + "vllm.v1.worker.gpu_worker.Worker.sleep", + ), + ) + for module, symbol, parameters, origin in contracts: + require_symbol( + module, + symbol, + parameters=parameters, + origin=origin, + strict=strict, + ) + + +def install_qwen3_moe_patch(*, strict: bool) -> PatchResult: + if not strict: + raise ValueError("the Qwen3-MoE public-reference installer requires strict=True") + global _INSTALLED + if _INSTALLED: + raise RuntimeError("Qwen3-MoE parity patch installed twice") + module = importlib.import_module("vllm.model_executor.models.qwen3_moe") + original_moe_block = module.Qwen3MoeSparseMoeBlock + original_attention = module.Qwen3MoeAttention + original_logits_processor = module.LogitsProcessor + original_marker = getattr(module, "_unirl_train_inference_parity", "") + moe_block = make_moe_block(original_moe_block) + attention = make_attention(original_attention) + logits_processor = make_logits_processor(original_logits_processor) + rope_results = _install_rope() + sleep_result = _install_sleep_cache_invalidation() + module.Qwen3MoeSparseMoeBlock = moe_block + module.Qwen3MoeAttention = attention + module.LogitsProcessor = logits_processor + module._unirl_train_inference_parity = True + _INSTALLED = True + return PatchResult( + name="qwen3_moe_30b_a3b", + symbols=( + *rope_results, + sleep_result, + symbol_result( + ("vllm.model_executor.models.qwen3_moe.Qwen3MoeSparseMoeBlock"), + moe_block, + before=original_moe_block, + actual=module.Qwen3MoeSparseMoeBlock, + ), + symbol_result( + "vllm.model_executor.models.qwen3_moe.Qwen3MoeAttention", + attention, + before=original_attention, + actual=module.Qwen3MoeAttention, + ), + symbol_result( + "vllm.model_executor.models.qwen3_moe.LogitsProcessor", + logits_processor, + before=original_logits_processor, + actual=module.LogitsProcessor, + ), + value_result( + ("vllm.model_executor.models.qwen3_moe._unirl_train_inference_parity"), + "literal:true", + before=original_marker, + actual=module._unirl_train_inference_parity, + verified=module._unirl_train_inference_parity is True, + ), + ), + ) + + +__all__ = ["install_qwen3_moe_patch", "preflight_qwen3_moe_patch"] diff --git a/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/models/qwen3_moe_30b_a3b/router.py b/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/models/qwen3_moe_30b_a3b/router.py new file mode 100644 index 000000000..521ebb8dc --- /dev/null +++ b/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/models/qwen3_moe_30b_a3b/router.py @@ -0,0 +1,45 @@ +"""HF-ordered router implemented with public vLLM BI providers.""" + +from __future__ import annotations + +import types + +import torch + +from ...common.providers import linear, softmax + + +def _gate_forward(self, hidden_states): + if self.bias is not None: + raise ValueError("Qwen3-MoE parity router gate requires bias=False") + hidden = hidden_states.reshape(-1, hidden_states.shape[-1]).contiguous() + return linear(hidden, self.weight, None), None + + +def install_gate(gate) -> None: + gate.forward = types.MethodType(_gate_forward, gate) + + +def _compute_routing( + self, + hidden_states, + router_logits, + indices_type, + **kwargs, +): + del hidden_states, kwargs + probabilities = softmax(router_logits.float().contiguous(), dim=-1) + weights, indices = torch.topk(probabilities, self.top_k, dim=-1) + if self.renormalize: + weights = weights / weights.sum(dim=-1, keepdim=True) + return weights.float(), indices.to(indices_type or torch.int32) + + +def install_hf_router(experts) -> None: + router = experts.router + if router.scoring_func != "softmax" or router.renormalize is not True: + raise ValueError("Qwen3-MoE parity requires softmax routing with renormalize=true") + router._compute_routing = types.MethodType(_compute_routing, router) + + +__all__ = ["install_gate", "install_hf_router"] diff --git a/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/registry.py b/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/registry.py new file mode 100644 index 000000000..2af2aa257 --- /dev/null +++ b/experimental/train_inference_parity/vllm_plugin/src/unirl_train_inference_parity_vllm/registry.py @@ -0,0 +1,142 @@ +"""Fail-closed, two-phase registry for parity plugin installers.""" + +from __future__ import annotations + +import inspect +from collections.abc import Callable, Mapping +from dataclasses import asdict, dataclass + + +@dataclass(frozen=True) +class SymbolResult: + """Before/after evidence for one concrete patched symbol.""" + + symbol: str + before_provider: str + before_signature: str | None + before_identity: int | None + after_provider: str + after_signature: str | None + after_identity: int | None + verified: bool + + +@dataclass(frozen=True) +class PatchResult: + """Structured result returned by every patch installer.""" + + name: str + symbols: tuple[SymbolResult, ...] + + def to_manifest(self) -> dict[str, object]: + return asdict(self) + + +@dataclass(frozen=True) +class Installer: + preflight: Callable[[bool], None] + install: Callable[[bool], PatchResult] + + +_INSTALLED: dict[str, PatchResult] = {} + + +def _provider_name(value: object) -> str: + module = getattr(value, "__module__", None) + qualname = getattr(value, "__qualname__", None) + if module and qualname: + return f"{module}.{qualname}" + if isinstance(value, str): + return value + return repr(value) + + +def _signature(value: object) -> str | None: + try: + return str(inspect.signature(value)) + except (TypeError, ValueError): + return None + + +def symbol_result( + symbol: str, + provider: object, + *, + before: object, + actual: object, +) -> SymbolResult: + """Describe and verify the provider currently bound to ``symbol``.""" + verified = actual is provider + if not verified: + raise RuntimeError( + f"parity provider verification failed for {symbol}: expected identity {provider!r}, got {actual!r}" + ) + return SymbolResult( + symbol=symbol, + before_provider=_provider_name(before), + before_signature=_signature(before), + before_identity=id(before), + after_provider=_provider_name(provider), + after_signature=_signature(provider), + after_identity=id(actual), + verified=True, + ) + + +def value_result( + symbol: str, + provider: str, + *, + before: object, + actual: object, + verified: bool, +) -> SymbolResult: + if not verified: + raise RuntimeError(f"parity value verification failed for {symbol}") + return SymbolResult( + symbol=symbol, + before_provider=_provider_name(before), + before_signature=None, + before_identity=id(before), + after_provider=provider, + after_signature=None, + after_identity=id(actual), + verified=True, + ) + + +def install_selected( + names: tuple[str, ...], + *, + strict: bool, + installers: Mapping[str, Installer], +) -> tuple[PatchResult, ...]: + unknown = sorted(set(names) - installers.keys()) + if unknown: + raise ValueError(f"unknown parity patches {unknown}; known={sorted(installers)}") + + pending = tuple(name for name in names if name not in _INSTALLED) + # Complete every compatibility check before mutating process-global state. + for name in pending: + installers[name].preflight(strict=strict) + for name in pending: + result = installers[name].install(strict=strict) + if result.name != name: + raise RuntimeError(f"parity installer {name!r} returned result for {result.name!r}") + _INSTALLED[name] = result + return tuple(_INSTALLED[name] for name in names) + + +def installed_manifest() -> tuple[PatchResult, ...]: + return tuple(_INSTALLED[name] for name in sorted(_INSTALLED)) + + +__all__ = [ + "Installer", + "PatchResult", + "SymbolResult", + "install_selected", + "installed_manifest", + "symbol_result", + "value_result", +]