GLM-5.3-Flash-EXL3-TR3-2.0bpw / glm53_exl3_tp4.py
0xSero's picture
Add files using upload-large-folder tool
2333577 verified
Raw
History Blame
9.84 kB
"""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()