0xSero's picture
Add files using upload-large-folder tool
5f827be verified
Raw
History Blame
33.9 kB
"""Laguna hybrid NVFP4 + calibrated EXL3 tail runtime for TP1 or TP2.
Enabled only when LAGUNA_HYBRID_TIER is TR2 or TR3. The stock compressed-
tensors NVFP4 MoE method is replaced with a fail-closed two-tier method:
* saliency-selected hot experts are loaded byte-for-byte from the source
checkpoint into a compact native vLLM CUTLASS NVFP4 kernel;
* the remaining experts are loaded from the calibrated TP2 rank-sliced EXL3
tail; TP1 reconstructs the full expert by executing both stored slices into
the kernel's additive fp32 output buffer;
* both tiers consume the same router weights and their local outputs are added
before vLLM performs its normal TP reduction.
The trellis path uses preallocated scratch and fixed chunking and is therefore
compatible with normal CUDA graph capture. This module does not change vLLM's
execution mode or graph configuration.
"""
from __future__ import annotations
import dataclasses
import json
import os
import re
import threading
from pathlib import Path
from typing import Any
TIER = os.environ.get("LAGUNA_HYBRID_TIER", "").upper()
if TIER in ("TR2", "TR3"):
import torch
import torch.nn as nn
from safetensors import safe_open
from vllm.distributed import (
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
)
from vllm.model_executor.layers.fused_moe import SharedExperts
from vllm.model_executor.layers.fused_moe.activation import MoEActivation
from vllm.model_executor.layers.fused_moe.config import FusedMoEParallelConfig
from vllm.model_executor.layers.fused_moe.oracle.nvfp4 import (
convert_to_nvfp4_moe_kernel_format,
make_nvfp4_moe_kernel,
make_nvfp4_moe_quant_config,
select_nvfp4_moe_backend,
)
from vllm.model_executor.layers.quantization.compressed_tensors.compressed_tensors_moe.compressed_tensors_moe_w4a4_nvfp4 import (
CompressedTensorsW4A4Nvfp4MoEMethod as _StockNvfp4MoE,
)
from vllm.model_executor.layers.quantization.utils.quant_utils import (
kNvfp4Dynamic,
kNvfp4Static,
)
from vllm.model_executor.utils import set_weight_attrs
from exllamav3.ext import exllamav3_ext as _exl3
BITS = 2 if TIER == "TR2" else 3
ARTIFACT_TP = 2
EXPERTS = 256
HIDDEN = 3072
INTERMEDIATE = 1024
TOP_K = 10
MCG = 0xCBAC1FED
CHUNK = int(os.environ.get("LAGUNA_HYBRID_TRELLIS_CHUNK", "128"))
if CHUNK <= 0:
raise RuntimeError("LAGUNA_HYBRID_TRELLIS_CHUNK must be positive")
MODEL_DIR = Path(os.environ.get("LAGUNA_HYBRID_MODEL_DIR", "/model"))
TAIL_DIR = Path(os.environ.get("LAGUNA_HYBRID_TAIL_DIR", "/tail"))
MAP_PATH = Path(
os.environ.get(
"LAGUNA_HYBRID_TIER_MAP",
str(TAIL_DIR / f"laguna-{TIER.lower()}-tier-map.json"),
)
)
if not MAP_PATH.is_file():
raise FileNotFoundError(f"Laguna hybrid tier map is missing: {MAP_PATH}")
_MAP_PAYLOAD = json.loads(MAP_PATH.read_text(encoding="utf-8"))
if _MAP_PAYLOAD.get("tier") != TIER:
raise RuntimeError(
f"tier map says {_MAP_PAYLOAD.get('tier')!r}, runtime requested {TIER}"
)
if not _MAP_PAYLOAD.get("coverage_gate", {}).get("passed"):
raise RuntimeError("Laguna hybrid tier map failed its coverage gate")
_EXPECTED_BUDGET = {
"TR2": {
"tail_encoding": "trellis2",
"tail_weight_bits": 2,
"hot_experts_per_sparse_layer": 96,
"tail_experts_per_sparse_layer": 160,
"average_weight_bits": 2.75,
},
"TR3": {
"tail_encoding": "trellis3",
"tail_weight_bits": 3,
"hot_experts_per_sparse_layer": 64,
"tail_experts_per_sparse_layer": 192,
"average_weight_bits": 3.25,
},
}[TIER]
_BUDGET = _MAP_PAYLOAD.get("bit_budget")
if not isinstance(_BUDGET, dict):
raise RuntimeError("Laguna hybrid tier map lacks a bit-budget contract")
_REQUIRED_BUDGET = {
"hot_encoding": "nvfp4",
"hot_weight_bits": 4,
"tail_encoding": _EXPECTED_BUDGET["tail_encoding"],
"tail_weight_bits": _EXPECTED_BUDGET["tail_weight_bits"],
"hot_experts_per_sparse_layer": _EXPECTED_BUDGET[
"hot_experts_per_sparse_layer"
],
"tail_experts_per_sparse_layer": _EXPECTED_BUDGET[
"tail_experts_per_sparse_layer"
],
}
for _field, _value in _REQUIRED_BUDGET.items():
if _BUDGET.get(_field) != _value:
raise RuntimeError(
f"Laguna {TIER} bit budget {_field}="
f"{_BUDGET.get(_field)!r} != {_value!r}"
)
_average = _EXPECTED_BUDGET["average_weight_bits"]
if (
float(_BUDGET.get("target_average_weight_bits", -1.0)) != _average
or float(_BUDGET.get("nominal_average_weight_bits", -1.0)) != _average
or _BUDGET.get("budget_passed") is not True
):
raise RuntimeError(
f"Laguna {TIER} tier map failed its exact bit-budget contract"
)
_LAYER_RE = re.compile(r"(?:^|\.)layers\.(\d+)(?:\.|$)")
_RUNTIME: dict[tuple[int, int], dict[str, torch.Tensor | int]] = {}
_RUNTIME_LOCK = threading.Lock()
def _layer_index(value: str | None) -> int:
match = _LAYER_RE.search(str(value or ""))
if match is None:
raise RuntimeError(f"cannot resolve Laguna layer index from {value!r}")
layer = int(match.group(1))
if not 1 <= layer <= 47:
raise RuntimeError(f"unexpected Laguna sparse layer index: {layer}")
return layer
def _tier_layer(layer: int) -> tuple[list[int], list[int]]:
entry = _MAP_PAYLOAD["layers"].get(str(layer))
if entry is None:
raise RuntimeError(f"tier map has no sparse layer {layer}")
hot = [int(value) for value in entry["hot_experts"]]
tail = [int(value) for value in entry["tail_experts"]]
expected_hot = 96 if TIER == "TR2" else 64
expected_tail = EXPERTS - expected_hot
if len(hot) != expected_hot or len(tail) != expected_tail:
raise RuntimeError(
f"layer {layer}: hot/tail counts {len(hot)}/{len(tail)} "
f"!= {expected_hot}/{expected_tail}"
)
if set(hot) | set(tail) != set(range(EXPERTS)) or set(hot) & set(tail):
raise RuntimeError(f"layer {layer}: tier map is not a 256-expert partition")
expected_tail_encoding = "trellis2" if TIER == "TR2" else "trellis3"
if (
entry.get("hot_encoding") != "nvfp4"
or entry.get("tail_encoding") != expected_tail_encoding
):
raise RuntimeError(
f"layer {layer}: tier encodings do not match {TIER}"
)
return hot, tail
def _runtime(device: torch.device, max_rows: int) -> dict[str, Any]:
key = (device.index if device.index is not None else torch.cuda.current_device(), max_rows)
with _RUNTIME_LOCK:
existing = _RUNTIME.get(key)
if existing is not None:
return existing
concurrency = int(_exl3.exl3_moe_max_concurrency(key[0]))
value: dict[str, Any] = {
"max_rows": max_rows,
"cap": CHUNK,
"xh": torch.empty((max_rows, HIDDEN), dtype=torch.float16, device=device),
"out32": torch.empty((max_rows, HIDDEN), dtype=torch.float32, device=device),
"tg": torch.empty(
(concurrency, CHUNK, HIDDEN), dtype=torch.float16, device=device
),
"tu": torch.empty(
(concurrency, CHUNK, HIDDEN), dtype=torch.float16, device=device
),
"ig": torch.empty(
(concurrency, CHUNK, INTERMEDIATE // ARTIFACT_TP),
dtype=torch.float16,
device=device,
),
"iu": torch.empty(
(concurrency, CHUNK, INTERMEDIATE // ARTIFACT_TP),
dtype=torch.float16,
device=device,
),
"flat_token": torch.arange(
CHUNK, dtype=torch.int64, device=device
).repeat_interleave(TOP_K),
"ones": torch.ones(CHUNK * TOP_K, dtype=torch.int64, device=device),
}
_RUNTIME[key] = value
print(
f"[laguna-hybrid] shared EXL3 runtime allocated: max_rows={max_rows} "
f"chunk={CHUNK} concurrency={concurrency}",
flush=True,
)
return value
class LagunaHybridNvfp4Exl3MoE(_StockNvfp4MoE):
def __init__(self, moe, layer_name: str | None = None, use_a16: bool = False):
super().__init__(moe, layer_name, use_a16)
self.layer_name = layer_name
self.layer_index = _layer_index(layer_name)
self.hot, self.tail = _tier_layer(self.layer_index)
self.hot_pos = {expert: index for index, expert in enumerate(self.hot)}
self.tail_pos = {expert: index for index, expert in enumerate(self.tail)}
self._seen: set[tuple[int, str, str]] = set()
self.hot_kernel = None
self.hot_layer = None
self.hot_quant_config = None
self.tail_slabs: dict[int, dict[str, dict[str, torch.Tensor]]] = {}
self.tail_ptrs: dict[int, list[torch.Tensor]] = {}
self.runtime_tp = 0
self.hot_expert_map = None
self.tail_lut = None
self.runtime = None
def create_weights(
self,
layer: torch.nn.Module,
num_experts: int,
hidden_size: int,
intermediate_size_per_partition: int,
params_dtype: torch.dtype,
**extra_weight_attrs,
) -> None:
if num_experts != EXPERTS or hidden_size != HIDDEN:
raise RuntimeError(
f"unexpected Laguna MoE geometry E={num_experts} H={hidden_size}"
)
runtime_tp = get_tensor_model_parallel_world_size()
if runtime_tp not in (1, ARTIFACT_TP):
raise RuntimeError(
"Laguna hybrid runtime supports tensor parallel world size 1 or 2"
)
expected_local_i = INTERMEDIATE // runtime_tp
if intermediate_size_per_partition != expected_local_i:
raise RuntimeError(
"Laguna hybrid intermediate partition mismatch: "
f"{intermediate_size_per_partition} != {expected_local_i} "
f"for TP{runtime_tp}"
)
self.runtime_tp = runtime_tp
layer.num_experts = num_experts
layer.params_dtype = params_dtype
hot_count = len(self.hot)
rank = get_tensor_model_parallel_rank()
def weight_loader(
param,
loaded,
name_mapped=None,
*,
shard_id=None,
expert_id=None,
return_success=False,
**_kwargs,
):
expert = int(expert_id)
if expert in self.tail_pos:
return True if return_success else None
local = self.hot_pos.get(expert)
if local is None:
return False if return_success else None
name = str(name_mapped or "")
shard = str(shard_id)
if shard not in ("w1", "w2", "w3"):
raise RuntimeError(f"unexpected expert shard {shard!r}")
family = "w13" if ".w13_" in name else "w2"
if "input_global_scale" in name:
field = "input_global_scale"
elif "weight_global_scale" in name:
field = "weight_global_scale"
elif "weight_scale" in name:
field = "weight_scale"
elif "weight_packed" in name:
field = "weight_packed"
else:
raise RuntimeError(f"unrecognized compact NVFP4 parameter: {name}")
if loaded.ndim >= 2 and runtime_tp == ARTIFACT_TP:
if shard in ("w1", "w3"):
loaded = loaded.chunk(ARTIFACT_TP, 0)[rank]
else:
loaded = loaded.chunk(ARTIFACT_TP, 1)[rank]
destination = param.data[local]
if family == "w13":
if field in ("weight_packed", "weight_scale"):
half = destination.shape[0] // 2
destination = (
destination[:half] if shard == "w1" else destination[half:]
)
elif field in ("weight_global_scale", "input_global_scale"):
destination = destination[0 if shard == "w1" else 1]
destination.copy_(loaded.reshape(destination.shape).to(destination.dtype))
self._seen.add((expert, shard, field))
return True if return_success else None
def parameter(name: str, shape: tuple[int, ...], dtype: torch.dtype) -> None:
value = nn.Parameter(
torch.empty(
shape,
dtype=dtype,
device=torch.cuda.current_device(),
),
requires_grad=False,
)
set_weight_attrs(value, {**extra_weight_attrs, "weight_loader": weight_loader})
layer.register_parameter(name, value)
local_i = expected_local_i
parameter(
"w13_weight_packed",
(hot_count, 2 * local_i, HIDDEN // 2),
torch.uint8,
)
parameter(
"w2_weight_packed",
(hot_count, HIDDEN, local_i // 2),
torch.uint8,
)
parameter(
"w13_weight_scale",
(hot_count, 2 * local_i, HIDDEN // 16),
torch.float8_e4m3fn,
)
parameter(
"w2_weight_scale",
(hot_count, HIDDEN, local_i // 16),
torch.float8_e4m3fn,
)
parameter("w13_weight_global_scale", (hot_count, 2), torch.float32)
parameter("w2_weight_global_scale", (hot_count,), torch.float32)
parameter("w13_input_global_scale", (hot_count, 2), torch.float32)
parameter("w2_input_global_scale", (hot_count,), torch.float32)
print(
f"[laguna-hybrid] layer {self.layer_index}: allocated "
f"{hot_count} NVFP4 + {len(self.tail)} trellis{BITS} experts "
f"for TP{runtime_tp} rank {rank}",
flush=True,
)
def _validate_hot_load(self) -> None:
required = {
(expert, shard, field)
for expert in self.hot
for shard in ("w1", "w2", "w3")
for field in (
"weight_packed",
"weight_scale",
"weight_global_scale",
"input_global_scale",
)
}
missing = required - self._seen
if missing:
raise RuntimeError(
f"layer {self.layer_index}: missing {len(missing)} hot NVFP4 tensors; "
f"first={sorted(missing)[:5]}"
)
def _build_hot_kernel(self, layer) -> None:
backend, experts_cls = select_nvfp4_moe_backend(
config=dataclasses.replace(
self.moe,
num_experts=len(self.hot),
num_local_experts=len(self.hot),
num_logical_experts=len(self.hot),
intermediate_size=self.moe.intermediate_size_per_partition,
moe_parallel_config=FusedMoEParallelConfig.make_no_parallel(),
),
weight_key=kNvfp4Static,
activation_key=kNvfp4Dynamic,
)
kept_moe = dataclasses.replace(
self.moe,
num_experts=len(self.hot),
num_local_experts=len(self.hot),
num_logical_experts=len(self.hot),
intermediate_size=self.moe.intermediate_size_per_partition,
moe_parallel_config=FusedMoEParallelConfig.make_no_parallel(),
)
compact = nn.Module()
compact.activation = getattr(layer, "activation", MoEActivation.SILU)
compact.moe_config = kept_moe
compact.local_num_experts = len(self.hot)
compact.swiglu_limit = getattr(layer, "swiglu_limit", None)
converted = convert_to_nvfp4_moe_kernel_format(
nvfp4_backend=backend,
layer=compact,
w13=layer.w13_weight_packed,
w13_scale=layer.w13_weight_scale,
w13_scale_2=(1.0 / layer.w13_weight_global_scale[:, 0].contiguous()),
a13_scale=(1.0 / layer.w13_input_global_scale),
w2=layer.w2_weight_packed,
w2_scale=layer.w2_weight_scale,
w2_scale_2=(1.0 / layer.w2_weight_global_scale),
a2_scale=(1.0 / layer.w2_input_global_scale),
is_act_and_mul=True,
)
(
compact.w13_weight,
compact.w13_weight_scale,
compact.w13_weight_scale_2,
compact.w13_input_scale,
compact.w2_weight,
compact.w2_weight_scale,
compact.w2_weight_scale_2,
compact.w2_input_scale,
) = converted
quant_config = make_nvfp4_moe_quant_config(
backend=backend,
w13_scale=compact.w13_weight_scale,
w2_scale=compact.w2_weight_scale,
w13_scale_2=compact.w13_weight_scale_2,
w2_scale_2=compact.w2_weight_scale_2,
a13_scale=compact.w13_input_scale,
a2_scale=compact.w2_input_scale,
swiglu_limit=compact.swiglu_limit,
layer=compact,
)
kernel = make_nvfp4_moe_kernel(
moe_quant_config=quant_config,
moe_config=kept_moe,
experts_cls=experts_cls,
backend=backend,
routing_tables=None,
layer=compact,
)
kernel.fused_experts.process_weights_after_loading(compact)
self.hot_kernel = kernel
self.hot_layer = compact
self.hot_quant_config = quant_config
self.moe_kernel = kernel
self._hot_keepalive = converted
device = compact.w13_weight.device
# CutlassExpertsFp4 explicitly does not support expert_map. Its
# apply implementation accepts that argument for the modular
# kernel interface but does not forward it to run_cutlass_moe_fp4.
# Keep a lookup table here and remap router IDs before invoking the
# compact kernel instead.
hot_expert_map = torch.full(
(EXPERTS,), -1, dtype=torch.int32, device=device
)
for expert, local in self.hot_pos.items():
hot_expert_map[expert] = local
self.hot_expert_map = hot_expert_map
for name in (
"w13_weight_packed",
"w2_weight_packed",
"w13_weight_scale",
"w2_weight_scale",
"w13_weight_global_scale",
"w2_weight_global_scale",
"w13_input_global_scale",
"w2_input_global_scale",
):
if hasattr(layer, name):
delattr(layer, name)
print(
f"[laguna-hybrid] layer {self.layer_index}: native hot kernel built "
f"with backend={backend} experts={len(self.hot)}",
flush=True,
)
def _load_tail(self, layer) -> None:
runtime_tp = get_tensor_model_parallel_world_size()
rank = get_tensor_model_parallel_rank()
if runtime_tp != self.runtime_tp:
raise RuntimeError(
f"Laguna hybrid TP changed during load: {self.runtime_tp} -> {runtime_tp}"
)
artifact_ranks = (
[rank] if runtime_tp == ARTIFACT_TP else list(range(ARTIFACT_TP))
)
device = self.hot_layer.w13_weight.device
count = len(self.tail)
shapes = {
"gate_proj": {
"trellis": (
count,
HIDDEN // 16,
(INTERMEDIATE // ARTIFACT_TP) // 16,
16 * BITS,
),
"suh": (count, HIDDEN),
"svh": (count, INTERMEDIATE // ARTIFACT_TP),
},
"up_proj": {
"trellis": (
count,
HIDDEN // 16,
(INTERMEDIATE // ARTIFACT_TP) // 16,
16 * BITS,
),
"suh": (count, HIDDEN),
"svh": (count, INTERMEDIATE // ARTIFACT_TP),
},
"down_proj": {
"trellis": (
count,
(INTERMEDIATE // ARTIFACT_TP) // 16,
HIDDEN // 16,
16 * BITS,
),
"suh": (count, INTERMEDIATE // ARTIFACT_TP),
"svh": (count, HIDDEN),
},
}
slabs_by_rank: dict[int, dict[str, dict[str, torch.Tensor]]] = {}
for artifact_rank in artifact_ranks:
slabs: dict[str, dict[str, torch.Tensor]] = {}
for projection, fields in shapes.items():
slabs[projection] = {
field: torch.empty(
shape,
dtype=(
torch.int16
if field == "trellis"
else torch.float16
),
device=device,
)
for field, shape in fields.items()
}
slabs_by_rank[artifact_rank] = slabs
for local, expert in enumerate(self.tail):
path = (
TAIL_DIR
/ TIER.lower()
/ f"layer-{self.layer_index:02d}"
/ f"expert-{expert:03d}.safetensors"
)
if not path.is_file():
raise FileNotFoundError(f"missing Laguna tail expert artifact: {path}")
with safe_open(str(path), framework="pt", device="cpu") as handle:
metadata = handle.metadata() or {}
expected_meta = {
"format": "exl3-trellis",
"bits": str(BITS),
"tp": str(ARTIFACT_TP),
"layer": str(self.layer_index),
"expert": str(expert),
"mcg_multiplier": hex(MCG),
"hessian": "routed-real-activations",
}
for key, expected in expected_meta.items():
if metadata.get(key) != expected:
raise RuntimeError(
f"{path}: metadata {key}={metadata.get(key)!r} != {expected!r}"
)
for artifact_rank in artifact_ranks:
for projection in ("gate_proj", "up_proj", "down_proj"):
marker = handle.get_tensor(
f"{projection}.rank{artifact_rank}.mcg"
)
if (int(marker.item()) & 0xFFFFFFFF) != MCG:
raise RuntimeError(f"{path}: wrong MCG marker")
for field in ("trellis", "suh", "svh"):
source = handle.get_tensor(
f"{projection}.rank{artifact_rank}.{field}"
)
destination = slabs_by_rank[artifact_rank][projection][
field
][local]
if tuple(source.shape) != tuple(destination.shape):
raise RuntimeError(
f"{path}:{projection}.rank{artifact_rank}.{field} "
f"shape {tuple(source.shape)} != "
f"{tuple(destination.shape)}"
)
destination.copy_(source.to(destination.dtype))
self.tail_slabs = slabs_by_rank
pointers_by_rank: dict[int, list[torch.Tensor]] = {}
for artifact_rank, slabs in slabs_by_rank.items():
pointers: list[torch.Tensor] = []
for projection in ("gate_proj", "up_proj", "down_proj"):
for field in ("trellis", "suh", "svh"):
slab = slabs[projection][field]
step = slab.stride(0) * slab.element_size()
pointers.append(
torch.tensor(
[
slab.data_ptr() + index * step
for index in range(count)
],
dtype=torch.int64,
device=device,
)
)
pointers_by_rank[artifact_rank] = pointers
self.tail_ptrs = pointers_by_rank
tail_lut = torch.full((EXPERTS,), count, dtype=torch.int64, device=device)
for expert, local in self.tail_pos.items():
tail_lut[expert] = local
self.tail_lut = tail_lut
from vllm.config import get_current_vllm_config
max_rows = int(
get_current_vllm_config().scheduler_config.max_num_batched_tokens
)
self.runtime = _runtime(device, max_rows)
print(
f"[laguna-hybrid] layer {self.layer_index}: loaded {count} calibrated "
f"trellis{BITS} experts from artifact ranks {artifact_ranks} "
f"for TP{runtime_tp} rank {rank}",
flush=True,
)
def _validate_execution_contract(self, layer) -> None:
if not bool(getattr(self.moe, "is_act_and_mul", False)):
raise RuntimeError(
"Laguna hybrid requires SwiGLU act-and-multiply experts"
)
activation = getattr(layer, "activation", MoEActivation.SILU)
if activation != MoEActivation.SILU:
raise RuntimeError(
f"Laguna hybrid requires SiLU experts, got {activation!r}"
)
swiglu_limit = getattr(layer, "swiglu_limit", None)
if swiglu_limit not in (None, 0, 0.0):
raise RuntimeError(
"Laguna hybrid EXL3 tail does not support a nonzero "
f"SwiGLU limit, got {swiglu_limit!r}"
)
if bool(getattr(layer, "apply_router_weight_on_input", False)):
raise RuntimeError(
"Laguna hybrid requires router weights to be applied to "
"expert outputs"
)
def process_weights_after_loading(self, layer) -> None:
self._validate_execution_contract(layer)
self._validate_hot_load()
self._build_hot_kernel(layer)
self._load_tail(layer)
def get_fused_moe_quant_config(self, layer):
if self.hot_quant_config is None:
raise RuntimeError("Laguna hot quant config requested before initialization")
return self.hot_quant_config
def _apply_tail(self, x, topk_weights, topk_ids):
runtime = self.runtime
if runtime is None or self.tail_lut is None:
raise RuntimeError("Laguna trellis runtime is not initialized")
rows = int(x.shape[0])
if rows > int(runtime["max_rows"]):
raise RuntimeError(
f"Laguna trellis rows {rows} exceed planned capacity "
f"{runtime['max_rows']}"
)
xh = runtime["xh"][:rows]
xh.copy_(x)
out = runtime["out32"][:rows]
out.zero_()
local_ids = self.tail_lut[topk_ids.long()]
weights = topk_weights.to(torch.float16)
count = len(self.tail)
cap = int(runtime["cap"])
for start in range(0, rows, cap):
chunk_rows = min(cap, rows - start)
flat = local_ids[start : start + chunk_rows].reshape(-1)
order = torch.argsort(flat)
token = runtime["flat_token"][: chunk_rows * TOP_K].index_select(
0, order
).contiguous()
sorted_weights = weights[start : start + chunk_rows].reshape(-1).index_select(
0, order
).contiguous()
counts = torch.zeros(count + 1, dtype=torch.int64, device=x.device)
counts.scatter_add_(0, flat, runtime["ones"][: chunk_rows * TOP_K])
# The EXL3 fused down projection scatter-adds with atomicAdd.
# In TP1, executing both physically stored TP2 artifact slices
# into the same zeroed fp32 buffer reconstructs the exact
# tensor-parallel sum without allocating inside graph replay.
for artifact_rank in sorted(self.tail_ptrs):
_exl3.exl3_moe(
xh[start : start + chunk_rows],
out[start : start + chunk_rows],
counts,
token,
sorted_weights,
runtime["tg"],
runtime["tu"],
runtime["ig"],
runtime["iu"],
0,
BITS,
BITS,
BITS,
*self.tail_ptrs[artifact_rank],
True,
False,
True,
False,
True,
False,
0.0,
)
return out
def apply(
self,
layer,
x,
topk_weights,
topk_ids,
shared_experts: SharedExperts | None,
shared_experts_input,
):
if (
self.hot_kernel is None
or self.hot_expert_map is None
or self.hot_layer is None
):
raise RuntimeError("Laguna hot kernel is not initialized")
# The router produces original IDs in [0, 256). Remap hot routes
# into the compact [0, hot_count) space. Tail routes are sent to
# the valid placeholder expert 0 with an exact zero router weight;
# this avoids unsupported/out-of-range IDs without renormalizing
# the original top-k weights.
local_hot_ids = self.hot_expert_map[topk_ids.long()]
hot_route = local_hot_ids >= 0
local_hot_ids = torch.where(
hot_route, local_hot_ids, torch.zeros_like(local_hot_ids)
)
local_hot_weights = topk_weights * hot_route.to(topk_weights.dtype)
hot = self.hot_kernel.apply(
x,
self.hot_layer.w13_weight,
self.hot_layer.w2_weight,
local_hot_weights,
local_hot_ids,
activation=self.hot_layer.activation,
global_num_experts=len(self.hot),
expert_map=None,
apply_router_weight_on_input=layer.apply_router_weight_on_input,
shared_experts=shared_experts,
shared_experts_input=shared_experts_input,
)
tail = self._apply_tail(x, topk_weights, topk_ids)
return torch.add(hot.float(), tail).to(x.dtype)
import vllm.model_executor.layers.quantization.compressed_tensors.compressed_tensors_moe.compressed_tensors_moe_w4a4_nvfp4 as _target
_target.CompressedTensorsW4A4Nvfp4MoEMethod = LagunaHybridNvfp4Exl3MoE
print(
f"[laguna-hybrid] installed {TIER} runtime: bits={BITS} "
f"artifact_tp={ARTIFACT_TP} runtime_tp=1-or-2 "
f"map={MAP_PATH} tail={TAIL_DIR} CUDA graphs unchanged",
flush=True,
)