"""Runtime primitive for Brandon-style GLM-5.3 K2/K3/K4 EXL3 hybrids.""" from __future__ import annotations import json from pathlib import Path from typing import Any import torch import torch.distributed as dist from torch import nn import torch.nn.functional as F from safetensors import safe_open from exllamav3.modules.quant import LinearEXL3 HIDDEN = 4096 LOCAL_INTERMEDIATE = 512 EXPERTS = 288 TP = 4 SWIGLU_LIMIT = 10.0 def _packed_prefix(layer: int, expert: int, projection: str, rank: int) -> str: return ( f"model.language_model.layers.{layer}.mlp.experts.{expert}." f"{projection}.rank{rank}" ) def _load_linear( handle: Any, prefix: str, in_features: int, out_features: int, out_dtype: torch.dtype, ) -> LinearEXL3: tensors = { suffix: handle.get_tensor(f"{prefix}.{suffix}") for suffix in ("suh", "svh", "trellis", "mcg") } return LinearEXL3( config=None, in_features=in_features, out_features=out_features, suh=tensors["suh"], svh=tensors["svh"], trellis=tensors["trellis"], mcg=tensors["mcg"], out_dtype=out_dtype, key=prefix, ) class TPRankEXL3Experts(nn.Module): """One TP rank of a manifest-bound 224-tail/64-K4 routed-expert union.""" def __init__( self, artifact_root: str | Path, layer: int, tp_rank: int, device: torch.device | str, process_group: Any | None = None, out_dtype: torch.dtype = torch.float16, swiglu_limit: float = SWIGLU_LIMIT, ) -> None: super().__init__() if not 3 <= layer <= 44: raise ValueError(f"layer outside routed scope: {layer}") if not 0 <= tp_rank < TP: raise ValueError(f"TP rank outside [0, {TP}): {tp_rank}") self.artifact_root = Path(artifact_root) self.layer = layer self.tp_rank = tp_rank self.device = torch.device(device) self.process_group = process_group self.out_dtype = out_dtype self.swiglu_limit = float(swiglu_limit) self.num_experts = EXPERTS self.gate: list[LinearEXL3 | None] = [None] * EXPERTS self.up: list[LinearEXL3 | None] = [None] * EXPERTS self.down: list[LinearEXL3 | None] = [None] * EXPERTS self.expert_integer_k: list[int | None] = [None] * EXPERTS self._load() def _expected_integer_k(self) -> dict[int, int]: bitmap_path = self.artifact_root / "tier_bitmap.json" bitmap = json.loads(bitmap_path.read_text()) if bitmap.get("state") != "PASS": raise RuntimeError("hybrid tier bitmap is not PASS") layer = bitmap.get("layers", {}).get(str(self.layer), {}) declared = layer.get("expert_k") if declared is not None: expected = {int(expert): int(k) for expert, k in declared.items()} else: expected = {} for key, integer_k in ( ("tail_k2", 2), ("tail_k3", 3), ("keep_k4", 4), ): for expert in layer.get(key, []): expert = int(expert) if expert in expected: raise RuntimeError( f"duplicate tier-bitmap expert at layer {self.layer}: {expert}" ) expected[expert] = integer_k if set(expected) != set(range(EXPERTS)) or not set(expected.values()).issubset( {2, 3, 4} ): raise RuntimeError( f"invalid tier-bitmap coverage at layer {self.layer}: {len(expected)}/{EXPERTS}" ) if sum(k == 4 for k in expected.values()) != 64: raise RuntimeError( f"tier bitmap does not retain exactly 64 K4 experts at layer {self.layer}" ) return expected def _load(self) -> None: expected_integer_k = self._expected_integer_k() observed = set() sidecars = sorted( (self.artifact_root / "layers").glob(f"layer-{self.layer:02d}-part-*.json") ) if not sidecars: raise RuntimeError(f"no hybrid packed sidecars for layer {self.layer}") for sidecar_path in sidecars: sidecar = json.loads(sidecar_path.read_text()) experts = [int(value) for value in sidecar.get("experts", [])] duplicate = observed.intersection(experts) if duplicate: raise RuntimeError( f"duplicate hybrid expert assignment at layer {self.layer}: {sorted(duplicate)}" ) observed.update(experts) path = sidecar_path.with_suffix(".safetensors") with safe_open(path, framework="pt", device=str(self.device)) as handle: for expert in experts: gate_prefix = _packed_prefix( self.layer, expert, "gate_proj", self.tp_rank ) trellis_shape = handle.get_slice(f"{gate_prefix}.trellis").get_shape() integer_k = int(trellis_shape[-1]) // 16 if integer_k not in (2, 3, 4): raise RuntimeError( f"unexpected hybrid K={integer_k} at layer={self.layer} expert={expert}" ) if integer_k != expected_integer_k[expert]: raise RuntimeError( f"hybrid K disagrees with tier bitmap at layer={self.layer} " f"expert={expert}: packed={integer_k} expected={expected_integer_k[expert]}" ) self.expert_integer_k[expert] = integer_k self.gate[expert] = _load_linear( handle, gate_prefix, HIDDEN, LOCAL_INTERMEDIATE, self.out_dtype, ) self.up[expert] = _load_linear( handle, _packed_prefix(self.layer, expert, "up_proj", self.tp_rank), HIDDEN, LOCAL_INTERMEDIATE, self.out_dtype, ) self.down[expert] = _load_linear( handle, _packed_prefix(self.layer, expert, "down_proj", self.tp_rank), LOCAL_INTERMEDIATE, HIDDEN, self.out_dtype, ) if observed != set(range(EXPERTS)): raise RuntimeError( f"incomplete hybrid expert union at layer {self.layer}: {len(observed)}/{EXPERTS}" ) if any( module is None for modules in (self.gate, self.up, self.down) for module in modules ): raise RuntimeError( f"incomplete hybrid EXL3 module coverage for layer {self.layer} rank {self.tp_rank}" ) observed_counts = { k: self.expert_integer_k.count(k) for k in (2, 3, 4) } expected_counts = { k: sum(value == k for value in expected_integer_k.values()) for k in (2, 3, 4) } if observed_counts != expected_counts: raise RuntimeError( f"hybrid tier cardinality mismatch at layer {self.layer}: " f"observed={observed_counts} expected={expected_counts}" ) @torch.inference_mode() def forward( self, hidden_states: torch.Tensor, top_k_index: torch.Tensor, top_k_weights: torch.Tensor, params: dict | None = None, ) -> torch.Tensor: if hidden_states.ndim != 2 or hidden_states.shape[-1] != HIDDEN: raise ValueError( f"expected [tokens, {HIDDEN}] hidden states, got {tuple(hidden_states.shape)}" ) if ( top_k_index.shape != top_k_weights.shape or top_k_index.shape[0] != hidden_states.shape[0] ): raise ValueError("top-k routing tensors do not match hidden states") params = {} if params is None else params final = torch.zeros_like(hidden_states, dtype=self.out_dtype) mask = F.one_hot(top_k_index, num_classes=EXPERTS).permute(2, 1, 0) hit = torch.greater(mask.sum(dim=(-1, -2)), 0).nonzero().flatten().tolist() for expert in hit: top_k_pos, token_idx = torch.where(mask[expert]) x = hidden_states[token_idx].to(self.out_dtype).contiguous() gate = self.gate[expert].forward(x, params).clamp(max=self.swiglu_limit) up = self.up[expert].forward(x, params).clamp( min=-self.swiglu_limit, max=self.swiglu_limit ) activated = (F.silu(gate) * up).contiguous() current = self.down[expert].forward(activated, params) current = current * top_k_weights[token_idx, top_k_pos, None].to( current.dtype ) final.index_add_(0, token_idx, current) if self.process_group is not None: if not dist.is_initialized(): raise RuntimeError( "a process group was provided but torch.distributed is not initialized" ) dist.all_reduce(final, op=dist.ReduceOp.SUM, group=self.process_group) return final def unload(self) -> None: for modules in (self.gate, self.up, self.down): for module in modules: if module is not None: module.unload() modules.clear()