File size: 9,835 Bytes
2333577
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
"""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()