From 8d37dd35c099f6ca341a651e135e7706c26a5260 Mon Sep 17 00:00:00 2001 From: JiaryCoder Date: Fri, 4 Sep 2026 08:04:35 +0000 Subject: [PATCH 1/2] feat(compilation): freeze FL dispatch before torch.compile and route hot ops through FlagGems PT2 contracts (vLLM 0.24) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit With `--compilation-config '{"custom_ops": ["all"]}'` vLLM 0.24 traces `CustomOp.forward_oot` fullgraph. The FL dispatch control plane (manager lookup, policy epochs, RLock, IO dump, exception-driven fallback) and the raw FlagGems launchers (LibEntry, logger.debug) are not Dynamo-traceable, so profile_run aborted at the first RMSNorm (Unsupported method call `__getitem__` in LibEntry.run; NameError('ext') in the Inductor subprocess). * dispatch/__init__.py: `freeze_dispatch()` resolves every `CachedOp` once before the first trace; frozen calls are a plain `impl.fn(*args)`. `FrozenDispatchManifest` (+ `FrozenOpSelection`) records policy fingerprint and per-op impl_id/kind/vendor/callable; sha256 of that is the manifest fingerprint. `thaw_dispatch()`, `is_dispatch_frozen()`, `get_frozen_dispatch_manifest()`. Post-fork reset via `os.register_at_fork`. * dispatch/policy.py, dispatch/manager.py: reject policy mutation while frozen; `reset_default_manager()` thaws. * compilation/dispatch.py: `freeze_dispatch_for_compile(vllm_config)` — binds the manifest fingerprint into `additional_config` (=> vLLM compile-cache identity) and registers `logging.Logger.{debug,info,warning,error}` in `torch._dynamo.config.ignore_logging_functions`. * worker/model_runner.py: call it in `load_model()` before any compile path. * backends/flaggems/impl/{normalization,activation,rotary,fused_moe}.py: on NVIDIA route to `flag_gems.pt2.*` (same kernel objects; eager unchanged). Verified on NVIDIA, TP=2, custom_ops=all: MiniCPM5-1B serve (14.2 s compile, 51/51 PIECEWISE + 51/51 FULL CUDA graphs), Qwen3.5-0.8B cold compile, Qwen3.6-35B-A3B MoE eager/compile tie-aware cross-phase 16/16 exact. Requires the matching FlagGems PR (flag_gems.pt2). Signed-off-by: JiaryCoder --- vllm_fl/compilation/__init__.py | 2 + vllm_fl/compilation/dispatch.py | 114 ++++++++ vllm_fl/dispatch/__init__.py | 270 +++++++++++++++++- .../backends/flaggems/impl/activation.py | 46 ++- .../backends/flaggems/impl/fused_moe.py | 78 +++-- .../backends/flaggems/impl/normalization.py | 26 +- .../dispatch/backends/flaggems/impl/rotary.py | 20 +- vllm_fl/dispatch/manager.py | 5 + vllm_fl/dispatch/policy.py | 17 ++ vllm_fl/worker/model_runner.py | 9 + 10 files changed, 550 insertions(+), 37 deletions(-) create mode 100644 vllm_fl/compilation/dispatch.py diff --git a/vllm_fl/compilation/__init__.py b/vllm_fl/compilation/__init__.py index 594495300..9afb71a5a 100644 --- a/vllm_fl/compilation/__init__.py +++ b/vllm_fl/compilation/__init__.py @@ -6,10 +6,12 @@ eager_break_during_capture, is_breakable_cudagraph_enabled, ) +from vllm_fl.compilation.dispatch import freeze_dispatch_for_compile __all__ = [ "is_breakable_cudagraph_enabled", "eager_break_during_capture", "BreakableCUDAGraphCapture", "BreakableCUDAGraphWrapper", + "freeze_dispatch_for_compile", ] diff --git a/vllm_fl/compilation/dispatch.py b/vllm_fl/compilation/dispatch.py new file mode 100644 index 000000000..59ce853d7 --- /dev/null +++ b/vllm_fl/compilation/dispatch.py @@ -0,0 +1,114 @@ +# Copyright (c) 2026 BAAI. All rights reserved. + +"""Bridge the dynamic FL dispatch control plane to a compiled runner.""" + +from __future__ import annotations + +import logging +from typing import Any + +from vllm_fl.dispatch import FrozenDispatchManifest, freeze_dispatch +from vllm_fl.dispatch.logger_manager import get_logger + +logger = get_logger(__name__) + +_HASH_KEY = "vllm_fl_dispatch_fingerprint" + + +def _install_logger_method_compile_guard() -> None: + """Permit Dynamo to no-op logging.Logger calls in FlagGems wrappers. + + FlagGems' public Python wrappers (``flag_gems.ops.*`` / + ``flag_gems.fused.*``) contain plain ``logger.debug(...)`` calls on their + entry paths — e.g. ``flag_gems/ops/rms_norm.py::rms_norm_forward``. + Torch Dynamo cannot trace ``logging.Logger`` methods and raises + ``torch._dynamo.exc.Unsupported`` the moment any of those wrappers is + reached from a compiled region. Registering the four unbound level + methods in ``torch._dynamo.config.ignore_logging_functions`` makes Dynamo + treat such calls as no-ops *inside compiled graphs only*; eager logging + behavior is unchanged. Idempotent. + """ + + try: + import torch._dynamo.config as _dynamo_config + except Exception: # pragma: no cover - torch always present in practice + return + ignore = getattr(_dynamo_config, "ignore_logging_functions", None) + if ignore is None or not hasattr(ignore, "add"): + return + for method_name in ("debug", "info", "warning", "error"): + ignore.add(getattr(logging.Logger, method_name)) + + +def _is_compiled_execution(vllm_config: Any) -> bool: + compilation_config = getattr(vllm_config, "compilation_config", None) + if compilation_config is None: + return False + if getattr(compilation_config, "backend", None) == "eager": + return False + + mode = getattr(compilation_config, "mode", None) + mode_name = getattr(mode, "name", str(mode)).upper() + return mode is not None and mode_name not in {"NONE", "COMPILATIONMODE.NONE"} + + +def _add_cache_fingerprint(vllm_config: Any, manifest: FrozenDispatchManifest) -> None: + """Make backend choices part of vLLM's compilation cache identity.""" + + additional_config = getattr(vllm_config, "additional_config", None) + if isinstance(additional_config, dict): + previous = additional_config.get(_HASH_KEY) + if previous is not None and previous != manifest.fingerprint: + raise RuntimeError( + "vLLM-FL dispatch selection changed after the vLLM config was " + "prepared. Rebuild the model runner before compiling." + ) + additional_config[_HASH_KEY] = manifest.fingerprint + else: + logger.warning( + "VllmConfig.additional_config is not a dict; the frozen FL dispatch " + "fingerprint could not be added to the outer vLLM cache key" + ) + + compilation_config = getattr(vllm_config, "compilation_config", None) + if compilation_config is not None: + # These fields are lazily recomputed by vLLM. Clear any path that may + # have been derived before the dispatch fingerprint was attached. + for attribute in ("cache_dir", "local_cache_dir"): + if hasattr(compilation_config, attribute): + setattr(compilation_config, attribute, "") + + +def freeze_dispatch_for_compile( + vllm_config: Any, +) -> FrozenDispatchManifest | None: + """Freeze all imported ``CachedOp`` sites before the first Dynamo trace. + + Optional operator modules can be imported by model registries without ever + being executed, so unresolved entries are recorded in the manifest rather + than failing model load. Calling one of them after freeze still raises a + deterministic error before entering Dynamo. + """ + + if not _is_compiled_execution(vllm_config): + return None + + _install_logger_method_compile_guard() + manifest = freeze_dispatch(strict=False) + _add_cache_fingerprint(vllm_config, manifest) + + if manifest.unresolved: + logger.warning( + "Frozen vLLM-FL dispatch with %d unresolved optional op(s): %s", + len(manifest.unresolved), + ", ".join(op_name for op_name, _ in manifest.unresolved), + ) + logger.info( + "Frozen vLLM-FL dispatch for torch.compile: %d op(s), fingerprint=%s", + len(manifest.selections), + manifest.fingerprint[:12], + ) + return manifest + + +__all__ = ["freeze_dispatch_for_compile"] diff --git a/vllm_fl/dispatch/__init__.py b/vllm_fl/dispatch/__init__.py index c3fb2870c..22c481a0a 100644 --- a/vllm_fl/dispatch/__init__.py +++ b/vllm_fl/dispatch/__init__.py @@ -76,7 +76,13 @@ - reference """ +import hashlib +import json import os +import threading +import weakref +from dataclasses import dataclass +from typing import Any, Iterable, Optional from .types import OpImpl, BackendImplKind, BackendPriority, match_token from .registry import OpRegistry, OpRegistrySnapshot @@ -148,6 +154,208 @@ def resolve_op(op_name: str): _OP_FAST_PATH_ENABLED = os.environ.get("VLLM_FL_OP_FAST_PATH", "1") == "1" +# --------------------------------------------------------------------------- +# Dispatch freeze for torch.compile +# --------------------------------------------------------------------------- +# When torch.compile traces a forward path that reaches a CachedOp, the +# dispatch control plane (manager lookup, policy epoch checks, IO dump, +# exception-driven fallback) graph-breaks Dynamo. The freeze mechanism +# resolves every imported CachedOp site ONCE before the first trace; after +# freeze a call is a plain ``frozen_impl.fn(*args)`` with no control-plane +# state. Selection changes => cache key changes (fingerprint is bound into +# the compile cache identity by the caller). + +@dataclass(frozen=True) +class FrozenOpSelection: + """One logical operator selected for a frozen execution phase.""" + + op_name: str + impl_id: str + kind: str + vendor: Optional[str] + callable_name: str + + def to_dict(self) -> dict[str, Any]: + return { + "op_name": self.op_name, + "impl_id": self.impl_id, + "kind": self.kind, + "vendor": self.vendor, + "callable": self.callable_name, + } + + +@dataclass(frozen=True) +class FrozenDispatchManifest: + """Stable dispatch choices used by one compiled execution phase.""" + + policy_fingerprint: str + selections: tuple[FrozenOpSelection, ...] + unresolved: tuple[tuple[str, str], ...] + fingerprint: str + + def to_dict(self) -> dict[str, Any]: + return { + "policy_fingerprint": self.policy_fingerprint, + "selections": [selection.to_dict() for selection in self.selections], + "unresolved": [ + {"op_name": op_name, "error": error} + for op_name, error in self.unresolved + ], + "fingerprint": self.fingerprint, + } + + +_CACHED_OPS: "weakref.WeakSet[CachedOp]" = weakref.WeakSet() +_FREEZE_LOCK = threading.RLock() +_DISPATCH_FROZEN = False +_FROZEN_MANIFEST: Optional[FrozenDispatchManifest] = None + + +def _callable_name(fn: Any) -> str: + module = getattr(fn, "__module__", type(fn).__module__) + qualname = getattr(fn, "__qualname__", type(fn).__qualname__) + return f"{module}.{qualname}" + + +def _is_torch_compiling() -> bool: + """Query Dynamo lazily so importing dispatch does not require torch.""" + + try: + import torch + + return bool(torch.compiler.is_compiling()) + except (AttributeError, ImportError): + return False + + +def is_dispatch_frozen() -> bool: + """Return whether runtime policy selection has been frozen.""" + + return _DISPATCH_FROZEN + + +def get_frozen_dispatch_manifest() -> Optional[FrozenDispatchManifest]: + """Return the manifest for the current frozen execution phase.""" + + return _FROZEN_MANIFEST + + +def _make_frozen_manifest( + policy_fingerprint: str, + resolved: dict[str, OpImpl], + unresolved: dict[str, str], +) -> FrozenDispatchManifest: + selections = tuple( + FrozenOpSelection( + op_name=op_name, + impl_id=impl.impl_id, + kind=impl.kind.value, + vendor=impl.vendor, + callable_name=_callable_name(impl.fn), + ) + for op_name, impl in sorted(resolved.items()) + ) + unresolved_items = tuple(sorted(unresolved.items())) + payload = { + "policy_fingerprint": policy_fingerprint, + "selections": [selection.to_dict() for selection in selections], + "unresolved": list(unresolved_items), + } + fingerprint = hashlib.sha256( + json.dumps(payload, sort_keys=True, separators=(",", ":")).encode("utf-8") + ).hexdigest() + return FrozenDispatchManifest( + policy_fingerprint=policy_fingerprint, + selections=selections, + unresolved=unresolved_items, + fingerprint=fingerprint, + ) + + +def freeze_dispatch( + op_names: Optional[Iterable[str]] = None, + *, + strict: bool = True, +) -> FrozenDispatchManifest: + """Resolve ``CachedOp`` instances before Dynamo starts tracing. + + Frozen calls contain no policy lookup, manager access, lock, IO dump, or + exception-driven backend fallback. ``strict=False`` is useful for a model + process that imported optional operator modules it will never execute; an + unresolved operator still fails immediately if it is called later. + """ + + global _DISPATCH_FROZEN, _FROZEN_MANIFEST + + requested = set(op_names) if op_names is not None else None + with _FREEZE_LOCK: + mgr = get_default_manager() + mgr.ensure_initialized() + + instances = [ + cached_op + for cached_op in list(_CACHED_OPS) + if requested is None or cached_op.op_name in requested + ] + names = sorted({cached_op.op_name for cached_op in instances}) + + resolved: dict[str, OpImpl] = {} + unresolved: dict[str, str] = {} + for op_name in names: + try: + impl = mgr._resolve_impl(op_name) + mgr._record_first_use(op_name, impl) + resolved[op_name] = impl + except Exception as exc: + unresolved[op_name] = f"{type(exc).__name__}: {exc}" + + if strict and unresolved: + details = "; ".join( + f"{op_name}: {error}" for op_name, error in unresolved.items() + ) + raise RuntimeError(f"Unable to freeze vLLM-FL dispatch: {details}") + + for cached_op in instances: + cached_op._frozen_impl = resolved.get(cached_op.op_name) + cached_op._freeze_error = unresolved.get(cached_op.op_name) + + manifest = _make_frozen_manifest( + get_policy().fingerprint(), resolved, unresolved + ) + _FROZEN_MANIFEST = manifest + _DISPATCH_FROZEN = True + return manifest + + +def thaw_dispatch() -> None: + """Leave the frozen phase. + + Existing compiled graphs must be discarded before this is used in a model + process. The function primarily exists for post-fork reset and tests. + """ + + global _DISPATCH_FROZEN, _FROZEN_MANIFEST + + with _FREEZE_LOCK: + _DISPATCH_FROZEN = False + _FROZEN_MANIFEST = None + for cached_op in list(_CACHED_OPS): + cached_op._clear_all_caches() + + +def _reset_frozen_dispatch_after_fork() -> None: + """Reset without acquiring a lock that may be owned by a vanished thread.""" + + global _FREEZE_LOCK, _DISPATCH_FROZEN, _FROZEN_MANIFEST + + _FREEZE_LOCK = threading.RLock() + _DISPATCH_FROZEN = False + _FROZEN_MANIFEST = None + for cached_op in list(_CACHED_OPS): + cached_op._clear_all_caches() + + class CachedOp: """Resolve an op once at the call site and refresh on policy changes. @@ -166,8 +374,11 @@ class CachedOp: """ __slots__ = ( + "__weakref__", "_op_name", "_impl", + "_frozen_impl", + "_freeze_error", "_use_manager_call", "_manager_id", "_manager_epoch", @@ -177,12 +388,54 @@ class CachedOp: def __init__(self, op_name: str) -> None: self._op_name = op_name self._impl = None + self._frozen_impl = None + self._freeze_error = None + self._use_manager_call = False + self._manager_id = -1 + self._manager_epoch = -1 + self._policy_epoch = -1 + _CACHED_OPS.add(self) + + @property + def op_name(self) -> str: + return self._op_name + + @property + def frozen_impl_id(self) -> Optional[str]: + impl = self._frozen_impl + return None if impl is None else impl.impl_id + + def _clear_all_caches(self) -> None: + self._impl = None + self._frozen_impl = None + self._freeze_error = None self._use_manager_call = False self._manager_id = -1 self._manager_epoch = -1 self._policy_epoch = -1 def __call__(self, *args, **kwargs): + # This is the only path used by a frozen compiled runner. Keep it + # before every manager/policy/debug check so Dynamo sees a stable + # callable and no dispatch control-plane state. + frozen_impl = self._frozen_impl + if frozen_impl is not None: + return frozen_impl.fn(*args, **kwargs) + + if _DISPATCH_FROZEN: + detail = f" ({self._freeze_error})" if self._freeze_error else "" + raise RuntimeError( + f"CachedOp '{self._op_name}' was not bound before the dispatch " + f"phase was frozen{detail}. Rebuild the compiled runner after " + "registering the operator." + ) + + if _is_torch_compiling(): + raise RuntimeError( + f"CachedOp '{self._op_name}' entered torch.compile before " + "vllm_fl.dispatch.freeze_dispatch() was called" + ) + mgr = get_default_manager() if not _OP_FAST_PATH_ENABLED: @@ -268,6 +521,14 @@ def __call__(self, *args, **kwargs): "clear_discovered_plugins", "PLUGIN_GROUP", "PLUGIN_MODULES_ENV", + # Dispatch freeze for torch.compile + "CachedOp", + "FrozenOpSelection", + "FrozenDispatchManifest", + "freeze_dispatch", + "thaw_dispatch", + "is_dispatch_frozen", + "get_frozen_dispatch_manifest", # Logging "get_logger", "set_log_level", @@ -281,5 +542,12 @@ def __call__(self, *args, **kwargs): # Convenience functions "call_op", "resolve_op", - "CachedOp", ] + + +# A manager and its registrations are process-local. Never inherit frozen +# callables across a fork; each worker binds again after loading its model. +try: + os.register_at_fork(after_in_child=_reset_frozen_dispatch_after_fork) +except AttributeError: + pass diff --git a/vllm_fl/dispatch/backends/flaggems/impl/activation.py b/vllm_fl/dispatch/backends/flaggems/impl/activation.py index d8e149275..8fb95a2ad 100644 --- a/vllm_fl/dispatch/backends/flaggems/impl/activation.py +++ b/vllm_fl/dispatch/backends/flaggems/impl/activation.py @@ -6,7 +6,41 @@ from __future__ import annotations +import flag_gems import torch +from flag_gems.fused import gelu_and_mul as _eager_gelu_and_mul +from flag_gems.modules.activation import gems_silu_and_mul as _eager_silu_and_mul + +# Resolve and materialize compiler-visible FlagGems families during backend +# registration, before Dynamo traces a frozen CachedOp. No alternate +# mathematical kernel is introduced: eager and compiled execution retain the +# original generated PointwiseDynamic JITFunctions. +from flag_gems.pt2.pointwise_dynamic import ( + ACTIVATION_POINTWISE_FAMILIES, + gelu_and_mul_pointwise, + materialize_pointwise_family_plans, + silu_and_mul_pointwise, +) + +# The four common sources captured by the transparent adapter are the NVIDIA +# PointwiseDynamic families. Other FlagGems vendors may replace these exports +# with architecture-specific implementations, so those vendors keep the +# non-PT2 path until they supply equivalent PT2 family specs. +_USE_TRANSPARENT_POINTWISE = flag_gems.vendor_name == "nvidia" + +if _USE_TRANSPARENT_POINTWISE: + # Freeze structural plans while the FlagGems backend is registered, before + # vLLM's profile run / Dynamo capture. Four generated families (SiLU, + # GELU-none, GELU-tanh, clamped SiLU) each expose six dtype/layout plans. + # The binary families retain the native contiguous->rank-1 and split->rank-2 + # materializations; clamped SiLU retains rank 2 because of its scalar + # broadcast. No Tensor or token-count value is retained. + materialize_pointwise_family_plans( + ACTIVATION_POINTWISE_FAMILIES, + ranks=(2,), + dtypes=(torch.float16, torch.bfloat16, torch.float32), + layout_classes=("contiguous_c", "split_last_dim_c"), + ) def silu_and_mul_flaggems(obj, x: torch.Tensor) -> torch.Tensor: @@ -20,11 +54,11 @@ def silu_and_mul_flaggems(obj, x: torch.Tensor) -> torch.Tensor: Returns: Output tensor of shape [..., d] """ - from flag_gems.modules.activation import gems_silu_and_mul - d = x.shape[-1] // 2 x1, x2 = x[..., :d], x[..., d:] - return gems_silu_and_mul(x1, x2) + if _USE_TRANSPARENT_POINTWISE: + return silu_and_mul_pointwise(x1, x2) + return _eager_silu_and_mul(x1, x2) def gelu_and_mul_flaggems(obj, x: torch.Tensor) -> torch.Tensor: @@ -38,9 +72,9 @@ def gelu_and_mul_flaggems(obj, x: torch.Tensor) -> torch.Tensor: Returns: Output tensor of shape [..., d] """ - from flag_gems.fused import gelu_and_mul - approximate = getattr(obj, "approximate", "none") if obj is not None else "none" d = x.shape[-1] // 2 x1, x2 = x[..., :d], x[..., d:] - return gelu_and_mul(x1, x2, approximate) + if _USE_TRANSPARENT_POINTWISE: + return gelu_and_mul_pointwise(x1, x2, approximate) + return _eager_gelu_and_mul(x1, x2, approximate) diff --git a/vllm_fl/dispatch/backends/flaggems/impl/fused_moe.py b/vllm_fl/dispatch/backends/flaggems/impl/fused_moe.py index 65bb8cea6..01925f436 100644 --- a/vllm_fl/dispatch/backends/flaggems/impl/fused_moe.py +++ b/vllm_fl/dispatch/backends/flaggems/impl/fused_moe.py @@ -4,23 +4,49 @@ FlagGems fused moe operator implementations. """ -from typing import Optional - +import flag_gems import torch +from flag_gems import ( + grouped_topk as _grouped_topk, + invoke_fused_moe_triton_kernel as _invoke_fused_moe_triton_kernel, + moe_align_block_size_triton as _moe_align_block_size_triton, + moe_sum as _moe_sum, + topk_softmax as _topk_softmax, + topk_softplus_sqrt as _topk_softplus_sqrt, +) +from flag_gems.pt2.fused_moe import ( + moe_sum as _pt2_moe_sum, + topk_softmax as _pt2_topk_softmax, +) +from flag_gems.pt2.moe_routing import ( + grouped_topk as _pt2_grouped_topk, + uses_common_moe_routing_kernels as _uses_common_moe_routing_kernels, +) + from vllm.triton_utils import triton from vllm.utils.math_utils import round_up +# These contracts capture the common/NVIDIA kernel objects. Other vendors +# may replace either public FlagGems export, so retain their original path +# until an equivalent vendor-specific PT2 contract is validated. +_USE_TRANSPARENT_MOE_PRIMITIVES = flag_gems.vendor_name == "nvidia" +_USE_TRANSPARENT_MOE_ROUTING = ( + flag_gems.vendor_name == "nvidia" + and _uses_common_moe_routing_kernels( + _grouped_topk, + _topk_softplus_sqrt, + ) +) + def moe_align_block_size_flaggems( topk_ids: torch.Tensor, block_size: int, num_experts: int, - expert_map: Optional[torch.Tensor] = None, + expert_map: torch.Tensor | None = None, pad_sorted_ids: bool = False, ignore_invalid_experts: bool = False, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - from flag_gems import moe_align_block_size_triton - max_num_tokens_padded = topk_ids.numel() + num_experts * (block_size - 1) if pad_sorted_ids: max_num_tokens_padded = round_up(max_num_tokens_padded, block_size) @@ -39,7 +65,7 @@ def moe_align_block_size_flaggems( # TODO(lms): ignore_invalid_experts not effective now # moe_align_block_size has optimize version to filtered out # all invalid experts directly when counting the number of experts - moe_align_block_size_triton( + _moe_align_block_size_triton( topk_ids, num_experts, block_size, @@ -56,20 +82,26 @@ def moe_align_block_size_flaggems( def topk_softmax_flaggems( topk_weights, topk_indices, token_expert_indices, gating_output, renormalize=False ): - from flag_gems import topk_softmax - - try: - topk_softmax( + # The FlagGems API accepts ``renormalize`` directly; a frozen compiled + # path cannot switch implementations after an exception, and Dynamo cannot + # soundly trace try/retry control flow. Both eager and compile call the + # same kernel with the same arguments. + if _USE_TRANSPARENT_MOE_PRIMITIVES: + _pt2_topk_softmax( + topk_weights, + topk_indices, + token_expert_indices, + gating_output, + renormalize, + ) + else: + _topk_softmax( topk_weights, topk_indices, token_expert_indices, gating_output, renormalize, ) - except: - topk_softmax(topk_weights, topk_indices, token_expert_indices, gating_output) - if renormalize: - topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True) return topk_weights, topk_indices @@ -95,9 +127,7 @@ def invoke_fused_moe_triton_kernel_flaggems( block_shape=None, B_bias=None, ): - from flag_gems import invoke_fused_moe_triton_kernel - - invoke_fused_moe_triton_kernel( + _invoke_fused_moe_triton_kernel( A, B, C, @@ -131,9 +161,10 @@ def grouped_topk_flaggems( bias, scoring_func=0, ): - from flag_gems import grouped_topk - - return grouped_topk( + grouped_topk_impl = ( + _pt2_grouped_topk if _USE_TRANSPARENT_MOE_ROUTING else _grouped_topk + ) + return grouped_topk_impl( scores, n_group, topk_group, @@ -146,6 +177,7 @@ def grouped_topk_flaggems( def moe_sum_flaggems(inp, out): - from flag_gems import moe_sum - - moe_sum(inp, out) + if _USE_TRANSPARENT_MOE_PRIMITIVES: + _pt2_moe_sum(inp, out) + else: + _moe_sum(inp, out) diff --git a/vllm_fl/dispatch/backends/flaggems/impl/normalization.py b/vllm_fl/dispatch/backends/flaggems/impl/normalization.py index c2e690500..335cb6036 100644 --- a/vllm_fl/dispatch/backends/flaggems/impl/normalization.py +++ b/vllm_fl/dispatch/backends/flaggems/impl/normalization.py @@ -6,16 +6,23 @@ from __future__ import annotations -from typing import Optional, Union - +import flag_gems import torch +from flag_gems.config import use_c_extension +from flag_gems.modules.normalization import gems_rms_forward +from flag_gems.pt2.rms_norm import rms_norm as _pt2_rms_norm + +# The Python adapter captures the common/NVIDIA Triton objects. Other vendors +# may replace RMSNorm, while a C++ installation deliberately selects its +# torch.ops implementation. Preserve both choices instead of bypassing them. +_USE_TRANSPARENT_RMS_NORM = flag_gems.vendor_name == "nvidia" and not use_c_extension def rms_norm_flaggems( obj, x: torch.Tensor, - residual: Optional[torch.Tensor] = None, -) -> Union[torch.Tensor, tuple[torch.Tensor, torch.Tensor]]: + residual: torch.Tensor | None = None, +) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]: """ RMS normalization using FlagGems. @@ -27,10 +34,17 @@ def rms_norm_flaggems( Returns: Normalized tensor, or tuple of (normalized, residual) if residual is provided """ - from flag_gems.modules.normalization import gems_rms_forward - # Get weight and epsilon from obj weight = obj.weight epsilon = obj.variance_epsilon + if _USE_TRANSPARENT_RMS_NORM: + variance_size_override = getattr(obj, "variance_size_override", None) + if variance_size_override is not None: + raise RuntimeError( + "FlagGems transparent RMSNorm does not support vLLM " + "variance_size_override; select a compatible backend" + ) + return _pt2_rms_norm(x, residual, weight, epsilon) + return gems_rms_forward(x, residual, weight, epsilon) diff --git a/vllm_fl/dispatch/backends/flaggems/impl/rotary.py b/vllm_fl/dispatch/backends/flaggems/impl/rotary.py index b4cb5c301..e762a4381 100644 --- a/vllm_fl/dispatch/backends/flaggems/impl/rotary.py +++ b/vllm_fl/dispatch/backends/flaggems/impl/rotary.py @@ -7,6 +7,7 @@ from __future__ import annotations import torch +from flag_gems.pt2 import rotary_embedding_inplace def rotary_embedding_flaggems( @@ -35,6 +36,23 @@ def rotary_embedding_flaggems( Returns: Tuple of (embedded_query, embedded_key) """ + if inplace: + if position_ids is None: + raise ValueError( + "The compiler-integrated FlagGems RoPE path requires position_ids" + ) + rotary_embedding_inplace( + query, + key, + cos, + sin, + position_ids, + rotary_interleaved, + ) + return query, key + + # vLLM uses the in-place path above. The generic out-of-place API has + # no manifest entry yet, so it stays on FlagGems' original implementation. from flag_gems.modules.rotary_embedding import gems_rope_forward return gems_rope_forward( @@ -44,5 +62,5 @@ def rotary_embedding_flaggems( sin, position_ids=position_ids, rotary_interleaved=rotary_interleaved, - inplace=inplace, + inplace=False, ) diff --git a/vllm_fl/dispatch/manager.py b/vllm_fl/dispatch/manager.py index 9bc172fc6..e3bab632d 100644 --- a/vllm_fl/dispatch/manager.py +++ b/vllm_fl/dispatch/manager.py @@ -8,6 +8,7 @@ import logging import os +import sys import threading from dataclasses import dataclass from typing import Callable, Dict, Optional, Set, Tuple @@ -634,4 +635,8 @@ def reset_default_manager() -> None: global _default_manager with _manager_lock: + dispatch_module = sys.modules.get("vllm_fl.dispatch") + thaw = getattr(dispatch_module, "thaw_dispatch", None) + if callable(thaw): + thaw() _default_manager = None diff --git a/vllm_fl/dispatch/policy.py b/vllm_fl/dispatch/policy.py index e7f06139f..76f7734aa 100644 --- a/vllm_fl/dispatch/policy.py +++ b/vllm_fl/dispatch/policy.py @@ -9,6 +9,7 @@ import contextvars import logging import os +import sys import threading import yaml from dataclasses import dataclass, field @@ -21,6 +22,19 @@ logger = logging.getLogger(__name__) +def _ensure_dispatch_mutable(action: str) -> None: + """Reject policy mutations after operator implementations are bound.""" + + dispatch_module = sys.modules.get("vllm_fl.dispatch") + is_frozen = getattr(dispatch_module, "is_dispatch_frozen", None) + if callable(is_frozen) and is_frozen(): + raise RuntimeError( + f"Cannot {action} while vLLM-FL dispatch is frozen. " + "Discard the compiled runner and call thaw_dispatch() first." + ) + + + # Valid preference values for VLLM_FL_PREFER PREFER_DEFAULT = "flagos" PREFER_VENDOR = "vendor" @@ -201,6 +215,7 @@ def get_policy(self) -> SelectionPolicy: def set_global_policy(self, policy: SelectionPolicy) -> SelectionPolicy: """Set the global policy and return the old policy.""" + _ensure_dispatch_mutable("set the global policy") with self._global_policy_lock: old_policy = self._global_policy self._global_policy = policy @@ -209,6 +224,7 @@ def set_global_policy(self, policy: SelectionPolicy) -> SelectionPolicy: def reset_global_policy(self) -> None: """Reset the global policy to environment defaults.""" + _ensure_dispatch_mutable("reset the global policy") with self._global_policy_lock: self._global_policy = None self.bump_policy_epoch() @@ -479,6 +495,7 @@ def __init__(self, manager: PolicyManager, policy: SelectionPolicy): self._token: Optional[contextvars.Token] = None def __enter__(self) -> "_PolicyContext": + _ensure_dispatch_mutable("enter a policy context") policy_var = self._manager._get_policy_var() self._token = policy_var.set(self._policy) self._manager.bump_policy_epoch() diff --git a/vllm_fl/worker/model_runner.py b/vllm_fl/worker/model_runner.py index 343dfdd45..a3ea3aa7c 100644 --- a/vllm_fl/worker/model_runner.py +++ b/vllm_fl/worker/model_runner.py @@ -5361,6 +5361,15 @@ def load_model(self, load_dummy_weights: bool = False) -> None: ): self.eplb_state.start_async_loop() + # Freeze the FL dispatch control plane before any Dynamo trace can + # reach a CachedOp site: after all model/operator modules have been + # loaded (so every CachedOp is registered) and before any compile call + # can trace the forward path. Frozen CachedOp sites call their + # selected implementation directly. + from vllm_fl.compilation import freeze_dispatch_for_compile + + freeze_dispatch_for_compile(self.vllm_config) + if ( self.vllm_config.compilation_config.mode == CompilationMode.STOCK_TORCH_COMPILE From 5ad7cf5d196fe3ed7e23276f574cb13024612510 Mon Sep 17 00:00:00 2001 From: JiaryCoder Date: Mon, 7 Sep 2026 04:19:47 +0000 Subject: [PATCH 2/2] fix(dispatch): import FlagGems impls at backend registration and bind ops directly Register the FlagGems implementations of the nine common ops (silu_and_mul, gelu_and_mul, rms_norm, rotary_embedding, moe_align_block_size, moe_sum, topk_softmax, invoke_fused_moe_triton_kernel, grouped_topk) by importing the impl functions inside register_builtins() and binding them directly, instead of binding FlagGemsBackend bound methods that import lazily on first call. With lazy binding, the impl modules -- and with them the flag_gems.pt2 triton_op registrations and the pointwise plan materialisation done at impl.activation import -- were only executed when Dynamo first inlined the backend method, i.e. inside the traced region. Doing that work during backend registration keeps torch.library registration and codegen out of the compiled graph, so freeze_dispatch() resolves every CachedOp to the final function object and the first trace imports nothing. Verified in a fresh process (no warm-up): after registration all impl modules, flag_gems.pt2.* and 24 pointwise plans are present and the six flag_gems_pt2 ops are registered; freeze imports nothing; the first-ever call of each op inside a fullgraph trace compiles and equals eager. In a real engine the freeze-time module state is impl=7 pt2=8 plans=24 ops=6/6. Signed-off-by: JiaryCoder --- .../backends/flaggems/register_ops.py | 35 ++++++++++++++----- 1 file changed, 26 insertions(+), 9 deletions(-) diff --git a/vllm_fl/dispatch/backends/flaggems/register_ops.py b/vllm_fl/dispatch/backends/flaggems/register_ops.py index 24da11fbd..33feea362 100644 --- a/vllm_fl/dispatch/backends/flaggems/register_ops.py +++ b/vllm_fl/dispatch/backends/flaggems/register_ops.py @@ -35,6 +35,23 @@ def register_builtins(registry) -> None: """ from .flaggems import FlagGemsBackend + # Import compiler-visible implementations while dispatch is still in its + # initialization phase. Registering a torch.library op during Dynamo + # tracing would itself be Python control flow and is intentionally banned. + from .impl.activation import ( + gelu_and_mul_flaggems, + silu_and_mul_flaggems, + ) + from .impl.fused_moe import ( + grouped_topk_flaggems, + invoke_fused_moe_triton_kernel_flaggems, + moe_align_block_size_flaggems, + moe_sum_flaggems, + topk_softmax_flaggems, + ) + from .impl.normalization import rms_norm_flaggems + from .impl.rotary import rotary_embedding_flaggems + backend = FlagGemsBackend() is_avail = backend.is_available @@ -67,7 +84,7 @@ def register_builtins(registry) -> None: op_name="silu_and_mul", impl_id="default.flagos", kind=BackendImplKind.DEFAULT, - fn=_bind_is_available(backend.silu_and_mul, is_avail), + fn=_bind_is_available(silu_and_mul_flaggems, is_avail), vendor=None, priority=BackendPriority.DEFAULT, ), @@ -75,7 +92,7 @@ def register_builtins(registry) -> None: op_name="gelu_and_mul", impl_id="default.flagos", kind=BackendImplKind.DEFAULT, - fn=_bind_is_available(backend.gelu_and_mul, is_avail), + fn=_bind_is_available(gelu_and_mul_flaggems, is_avail), vendor=None, priority=BackendPriority.DEFAULT, ), @@ -84,7 +101,7 @@ def register_builtins(registry) -> None: op_name="rms_norm", impl_id="default.flagos", kind=BackendImplKind.DEFAULT, - fn=_bind_is_available(backend.rms_norm, is_avail), + fn=_bind_is_available(rms_norm_flaggems, is_avail), vendor=None, priority=BackendPriority.DEFAULT, ), @@ -93,7 +110,7 @@ def register_builtins(registry) -> None: op_name="rotary_embedding", impl_id="default.flagos", kind=BackendImplKind.DEFAULT, - fn=_bind_is_available(backend.rotary_embedding, is_avail), + fn=_bind_is_available(rotary_embedding_flaggems, is_avail), vendor=None, priority=BackendPriority.DEFAULT, ), @@ -111,7 +128,7 @@ def register_builtins(registry) -> None: op_name="moe_align_block_size", impl_id="default.flagos", kind=BackendImplKind.DEFAULT, - fn=_bind_is_available(backend.moe_align_block_size, is_avail), + fn=_bind_is_available(moe_align_block_size_flaggems, is_avail), vendor=None, priority=BackendPriority.DEFAULT, ), @@ -120,7 +137,7 @@ def register_builtins(registry) -> None: op_name="moe_sum", impl_id="default.flagos", kind=BackendImplKind.DEFAULT, - fn=_bind_is_available(backend.moe_sum, is_avail), + fn=_bind_is_available(moe_sum_flaggems, is_avail), vendor=None, priority=BackendPriority.DEFAULT, ), @@ -129,7 +146,7 @@ def register_builtins(registry) -> None: op_name="topk_softmax", impl_id="default.flagos", kind=BackendImplKind.DEFAULT, - fn=_bind_is_available(backend.topk_softmax, is_avail), + fn=_bind_is_available(topk_softmax_flaggems, is_avail), vendor=None, priority=BackendPriority.DEFAULT, ), @@ -138,7 +155,7 @@ def register_builtins(registry) -> None: op_name="invoke_fused_moe_triton_kernel", impl_id="default.flagos", kind=BackendImplKind.DEFAULT, - fn=_bind_is_available(backend.invoke_fused_moe_triton_kernel, is_avail), + fn=_bind_is_available(invoke_fused_moe_triton_kernel_flaggems, is_avail), vendor=None, priority=BackendPriority.DEFAULT, ), @@ -147,7 +164,7 @@ def register_builtins(registry) -> None: op_name="grouped_topk", impl_id="default.flagos", kind=BackendImplKind.DEFAULT, - fn=_bind_is_available(backend.grouped_topk, is_avail), + fn=_bind_is_available(grouped_topk_flaggems, is_avail), vendor=None, priority=BackendPriority.DEFAULT, ),