gridbook 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
gridbook/__init__.py ADDED
@@ -0,0 +1,18 @@
1
+ """vLLM out-of-tree plugin for the NVFP4-CB / FP8-CB product-codebook formats.
2
+
3
+ Registers a ``QuantizationConfig`` (``"gridbook"``, with ``"prismaquant"`` kept
4
+ as a legacy alias) plus the linear and fused-MoE methods that serve
5
+ codebook-quantized weights. The CUDA kernels ship as sources under
6
+ ``gridbook/csrc`` and are JIT-compiled on first use; without nvcc the plugin
7
+ falls back to a correct-but-slow Triton path.
8
+
9
+ ``register`` is lazy so ``import gridbook.codec`` / ``import gridbook.kernels``
10
+ (the format and correctness tests) work without vLLM installed.
11
+ """
12
+
13
+ __version__ = "0.1.0"
14
+
15
+
16
+ def register() -> None:
17
+ from .plugin import register as _register
18
+ _register()
gridbook/codec.py ADDED
@@ -0,0 +1,145 @@
1
+ """Self-contained (no `prismaquant` import at serve time) CB codec helpers:
2
+
3
+ * load-time preprocessing that turns the shipped layout (LAYOUT.md) into the
4
+ small resident tensors the Triton kernel consumes — the flat codebook, the
5
+ pre-decoded fp4 scale plane, and an 8-byte-padded index stream. None of these
6
+ is a dense [N,K] weight (INV-1 holds);
7
+ * activation QDQ that reproduces the emulation gate's served-activation buckets
8
+ (fp4 group-16 RTN / fp8 dynamic per-token) so served KL is comparable to the
9
+ emulated prediction.
10
+ """
11
+ from __future__ import annotations
12
+
13
+ import torch
14
+ import torch.nn.functional as F
15
+
16
+ VEC_DIM = 8
17
+ SUPERBLOCK = 256
18
+ FP4_GROUP = 16
19
+ _E4M3 = torch.float8_e4m3fn
20
+ NVFP4_GRID_MAX = 6.0
21
+ FP8_ELEMENT_MAX = 448.0
22
+
23
+ # E2M1 element grid (sorted ascending), for the fp4 activation RTN.
24
+ _E2M1 = (0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0)
25
+
26
+ # Two-tier v2 scale coding (docs/nvfp4-cb-plan/two-tier-scale-spec.md §1). Kept
27
+ # in sync with prismaquant.nvfp4_cb_formats (the reference the kernel matches).
28
+ SCALE_CODING_TWO_TIER = "two_tier"
29
+ TWO_TIER_SUPER_BIAS = 127
30
+ TWO_TIER_SUB_TABLE = (1.0, 1.125, 1.25, 1.375, 1.5, 1.625, 1.75, 1.875,
31
+ 2.0, 2.25, 2.5, 2.75, 3.0, 3.25, 3.5, 3.75)
32
+
33
+
34
+ def type_size(k: int, is_fp4: bool) -> int:
35
+ return 4 * int(k) + (16 if is_fp4 else 0)
36
+
37
+
38
+ def build_flat_codebook(sub_tables: list[torch.Tensor]) -> torch.Tensor:
39
+ """Concatenate product sub-tables (each (2^w, sub_dim)) into the flat layout
40
+ the kernel gathers from: block ``s`` = ``sub_tables[s].reshape(-1)`` (row
41
+ major, so entry (idx, local) sits at idx*sub_dim + local)."""
42
+ return torch.cat([t.reshape(-1).to(torch.bfloat16).contiguous()
43
+ for t in sub_tables]).contiguous()
44
+
45
+
46
+ def build_compose_table(sub_table) -> torch.Tensor:
47
+ """Two-tier v2 (docs/nvfp4-cb-plan/two-tier-scale-spec.md §1): the (256,16)
48
+ compose table ``T[c]·2^(E-127)`` flattened to (4096,) fp32, bit-exact to
49
+ ``nvfp4_cb_formats._two_tier_tables`` (float64 product -> fp32). The kernel
50
+ gathers ``compose[super_e*16 + code]`` per group — no resident fp32 plane
51
+ (spec §4/G4), just this 16 KiB constant table."""
52
+ T = torch.tensor(list(sub_table), dtype=torch.float64)
53
+ exps = torch.arange(256, dtype=torch.float64)
54
+ compose = (T[None, :] * torch.pow(2.0, exps[:, None] - 127.0)).to(
55
+ torch.float32) # (256, 16)
56
+ return compose.reshape(-1).contiguous() # (4096,)
57
+
58
+
59
+ def decode_fp4_scale_plane(qw: torch.Tensor, k: int) -> torch.Tensor:
60
+ """(N, n_sb*type_size) uint8 -> (N, n_sb*16) fp32 group-16 scales, decoded
61
+ from the E4M3 scale plane that follows each superblock's 4k index bytes."""
62
+ n, row_bytes = qw.shape
63
+ ts = type_size(k, is_fp4=True)
64
+ n_sb = row_bytes // ts
65
+ blk = qw.reshape(n, n_sb, ts)
66
+ plane = blk[:, :, 4 * k:4 * k + FP4_GROUP].contiguous() # (N, n_sb, 16)
67
+ return plane.view(_E4M3).to(torch.float32).reshape(n, n_sb * FP4_GROUP)
68
+
69
+
70
+ PAD_BYTES = 16
71
+
72
+
73
+ def pad_qweight(qw: torch.Tensor) -> torch.Tensor:
74
+ """Right-pad each row by ``PAD_BYTES`` so the padded buffer satisfies BOTH
75
+ invariants every consumer of a padded row depends on:
76
+
77
+ 1. **>= 8 bytes of read slack per row.** The decode/expand kernels extract a
78
+ codeword with an 8-byte window anchored at the codeword's first byte, so
79
+ the last codeword of the last superblock of the last row reads up to 7
80
+ bytes past the packed data. Without the slack that is an out-of-bounds
81
+ global read (illegal-memory-access, or silent garbage in the final
82
+ output rows).
83
+ 2. **The padded row stride stays a 16-byte multiple** whenever the UNPADDED
84
+ one was. The fp8 CUTLASS entries (``cb_fused_prefill_mm_scaled``, the
85
+ persistent-TC prefill) take the row stride explicitly and TORCH_CHECK
86
+ ``stride(0) % 16 == 0`` — TMA needs 16-byte-aligned row starts. Every fp8
87
+ rung has ``type_size = 4k`` in {112,128,144,160,176,192}, so
88
+ ``row_bytes`` is 16-aligned and ``row_bytes + 16`` still is; the old
89
+ ``+ 8`` pad was NOT, which is why the padded buffer could not be shared
90
+ with the registered ``cb_qweight`` parameter and both had to stay
91
+ resident (see ``linear.process_weights_after_loading``). Pad width is
92
+ therefore load-bearing, not a spare-bytes choice: dropping it back to 8
93
+ re-breaks the fp8 prefill kernels' alignment check.
94
+
95
+ (fp4 rungs carry an odd ``type_size`` — ``4k+16`` v1, ``4k+9`` v2 — so their
96
+ row stride is not 16-aligned either way; they only ever hit kernels that
97
+ take the stride explicitly and require ``stride(1) == 1``.)
98
+
99
+ Consumers must read the row stride from the tensor (``.stride(0)``), never
100
+ derive it as ``size(1) - PAD_BYTES``.
101
+ """
102
+ return F.pad(qw.contiguous(), (0, PAD_BYTES), value=0).contiguous()
103
+
104
+
105
+ # Signed E2M1 grid, cached per device: building it per call allocated a CPU
106
+ # tensor and H2D-copied it on EVERY fp4 activation QDQ — a hidden sync in the
107
+ # eager decode hot path, and a hard error under CUDA-graph capture (unpinned
108
+ # CPU->CUDA copy). Warmup forwards populate the cache before any capture.
109
+ _FP4_QDQ_GRID: dict = {}
110
+
111
+
112
+ def _fp4_qdq_grid(device: torch.device) -> torch.Tensor:
113
+ grid = _FP4_QDQ_GRID.get(device)
114
+ if grid is None:
115
+ grid = torch.tensor(sorted({v for a in _E2M1 for v in (a, -a)}),
116
+ dtype=torch.float32, device=device)
117
+ _FP4_QDQ_GRID[device] = grid
118
+ return grid
119
+
120
+
121
+ def fp4_group16_act_qdq(x: torch.Tensor) -> torch.Tensor:
122
+ """W4A4 activation bucket: RTN to E2M1 at group-16 amax/6 scale (mirrors
123
+ format_registry `_make_rtn('fp4_e2m1', 16)`)."""
124
+ in_f = x.shape[-1]
125
+ grid = _fp4_qdq_grid(x.device)
126
+ w = x.reshape(-1, in_f).float().reshape(-1, in_f // FP4_GROUP, FP4_GROUP)
127
+ scale = w.abs().amax(dim=-1, keepdim=True).clamp_min(1e-8) / NVFP4_GRID_MAX
128
+ xg = w / scale
129
+ idx = torch.bucketize(xg.contiguous(), grid)
130
+ lo = grid[(idx - 1).clamp_min(0)]
131
+ hi = grid[idx.clamp_max(grid.numel() - 1)]
132
+ q = torch.where((hi - xg).abs() < (xg - lo).abs(), hi, lo)
133
+ return (q * scale).reshape(x.shape).to(x.dtype)
134
+
135
+
136
+ def fp8_dynamic_act_qdq(x: torch.Tensor) -> torch.Tensor:
137
+ """W8A8 activation bucket: vLLM dynamic per-token E4M3 (mirrors
138
+ fp8_dynamic.fp8_dynamic_activation_qdq_vllm)."""
139
+ rows = x.reshape(-1, x.shape[-1]).float()
140
+ min_scale = 1.0 / (FP8_ELEMENT_MAX * 512.0)
141
+ scale = (rows.abs().amax(dim=-1, keepdim=True) / FP8_ELEMENT_MAX
142
+ ).clamp_min(min_scale)
143
+ q = (rows / scale).clamp(-FP8_ELEMENT_MAX, FP8_ELEMENT_MAX).to(_E4M3)
144
+ deq = q.to(torch.float32) * scale
145
+ return deq.reshape(x.shape).to(x.dtype)
gridbook/config.py ADDED
@@ -0,0 +1,425 @@
1
+ """``PrismaQuantConfig`` — the vLLM quantization config for the NVFP4-CB /
2
+ FP8-CB out-of-tree lane (docs/nvfp4-cb-plan/serving-kernel.md §2, LAYOUT.md §4).
3
+
4
+ vLLM auto-detects us from ``quant_method == "prismaquant"``. The exporter writes
5
+ ``config.json['quantization_config']`` as a *pointer* (``config_file`` ->
6
+ ``quant_config.json`` + ``codebook_file`` -> ``cb_codebooks.pqcb``); the full
7
+ ``config_groups`` / ``ignore`` live in ``quant_config.json``. We resolve that
8
+ sidecar **lazily** (via ``get_current_vllm_config()``, the same handle
9
+ ``get_codebooks`` uses) since ``from_config`` runs before the model dir is
10
+ plumbed. Inlined configs (``config_groups`` already present) are also accepted.
11
+
12
+ **Mixed-container dispatch (serving-kernel.md §2).** A config group with a
13
+ ``"scheme"`` key is a CB group (our nvfp4_cb/fp8_cb vocabulary) -> our
14
+ ``PrismaQuantCBLinearMethod``. A group WITHOUT it uses the exact stock
15
+ compressed-tensors vocabulary -> a real ``CompressedTensorsConfig`` we construct
16
+ and delegate to (``CompressedTensorsW4A4Nvfp4`` for NVFP4 groups, the fp8 scheme
17
+ for FP8_DYNAMIC). ``ignore`` -> ``UnquantizedLinearMethod``.
18
+ """
19
+ from __future__ import annotations
20
+
21
+ import json
22
+ import os
23
+ from typing import Any
24
+
25
+ import torch
26
+ from vllm.model_executor.layers.linear import LinearBase, UnquantizedLinearMethod
27
+ from vllm.model_executor.layers.quantization.base_config import (
28
+ QuantizationConfig,
29
+ QuantizeMethodBase,
30
+ )
31
+ from vllm.model_executor.layers.vocab_parallel_embedding import (
32
+ UnquantizedEmbeddingMethod,
33
+ VocabParallelEmbedding,
34
+ )
35
+
36
+ try:
37
+ from vllm.model_executor.layers.fused_moe import RoutedExperts
38
+ except Exception: # pragma: no cover - older vLLM
39
+ RoutedExperts = None
40
+
41
+ _MOE_LEAVES = ("gate_up_proj", "down_proj", "gate_proj", "up_proj")
42
+
43
+ # vLLM fuses these siblings into one module; packed_modules_mapping is populated
44
+ # by dispatch time, but we keep the standard mapping as a fallback.
45
+ _FUSED_FALLBACK = {
46
+ "qkv_proj": ["q_proj", "k_proj", "v_proj"],
47
+ "gate_up_proj": ["gate_proj", "up_proj"],
48
+ "in_proj_qkvz": ["in_proj_qkv", "in_proj_z"],
49
+ "in_proj_ba": ["in_proj_b", "in_proj_a"],
50
+ }
51
+
52
+
53
+ def _canonical_prefix(prefix: str) -> str:
54
+ """vLLM serving prefix -> canonical target namespace. Some model classes
55
+ wrap the LM (`language_model.model.layers.*` on Qwen3.5-class VL) while
56
+ targets are canonical `model.layers.*`; strip the wrapper when the next
57
+ component is `model.` (measured via PRISMAQUANT_DEBUG_PREFIXES,
58
+ 2026-07-22 — every LM layer resolved no-scheme without this)."""
59
+ if prefix.startswith("language_model.model."):
60
+ return prefix[len("language_model."):]
61
+ if prefix.startswith("language_model."):
62
+ return "model." + prefix[len("language_model."):]
63
+ # Pre-fix multimodal CHECKPOINT namespace (shipped 27B gridbook artifact):
64
+ # ``model.language_model.layers.*`` denotes the same Linear as canonical
65
+ # ``model.layers.*``. Normalising it here (as well as in
66
+ # ``_canonical_target``) keeps probe-side and target-side on one string.
67
+ if prefix.startswith("model.language_model."):
68
+ return "model." + prefix[len("model.language_model."):]
69
+ return prefix
70
+
71
+
72
+ def _candidate_bases(name: str) -> list[str]:
73
+ """Every namespace vintage *name* can legitimately be matched against,
74
+ **most specific first** (the string as given, then its canonical form).
75
+
76
+ THE one place that answers "which namespace am I in?". A stored target /
77
+ serving prefix reaches us in one of three vintages — the old multimodal
78
+ CHECKPOINT form (``model.language_model.*``), the canonical form
79
+ (``model.*``), and the vLLM wrapper-class SERVING form
80
+ (``language_model.model.*``) — and ``apply_vllm_mapper`` can move the
81
+ stored keys into a *fourth*, the mapper's own namespace, AFTER
82
+ ``_ensure_resolved`` canonicalised them. Anything that matches a prefix
83
+ against ``target_scheme`` / ``ignore`` must therefore try both sides, and
84
+ must do so HERE: the dense fused path grew its own single-namespace copy of
85
+ this logic and silently mis-resolved for it (issue #1). A future fifth
86
+ vintage should mean editing this function and nothing else.
87
+ """
88
+ canonical = _canonical_prefix(name)
89
+ return [name] if canonical == name else [name, canonical]
90
+
91
+
92
+ def _canonical_target(name: str) -> str:
93
+ """Stored ``config_groups[*].targets`` / ``ignore`` entry -> canonical
94
+ target namespace, so historical checkpoint-namespace artifacts resolve
95
+ against the canonicalised serving prefixes ``_canonical_prefix`` produces.
96
+
97
+ Rewrites (prefix-anchored only):
98
+ ``model.language_model.`` -> ``model.`` (old multimodal ckpt)
99
+ ``language_model.model.`` -> ``model.`` (serving wrapper form)
100
+ ``language_model.<rest>`` -> ``model.<rest>``
101
+ Everything else (``visual.*``, ``mtp.*``, plain ``model.layers.*``,
102
+ bare leaf names) passes through untouched."""
103
+ return _canonical_prefix(name)
104
+
105
+
106
+ def _resolve_model_file(model_dir: str, fname: str) -> str:
107
+ """Local path for a sidecar file next to the model. When the model was
108
+ given as a Hub repo id (``vllm serve rdtand/...``) rather than a local
109
+ directory, fetch the sidecar from the Hub — vLLM's own loader handles the
110
+ weights that way, but OUR sidecars (quant_config.json, the .pqcb codebook
111
+ blob) were opened with a plain path join, which broke every serve-by-id
112
+ until 2026-07-22."""
113
+ if os.path.isdir(model_dir):
114
+ return os.path.join(model_dir, fname)
115
+ from huggingface_hub import hf_hub_download
116
+ return hf_hub_download(repo_id=model_dir, filename=fname)
117
+
118
+
119
+ class PrismaQuantConfig(QuantizationConfig):
120
+ """Per-layer dispatch: CB decode / stock-CT delegation / unquantized."""
121
+
122
+ def __init__(self, raw_config: dict) -> None:
123
+ super().__init__()
124
+ self._raw_config = dict(raw_config or {})
125
+ self.codebook_file = self._raw_config.get("codebook_file",
126
+ "cb_codebooks.pqcb")
127
+ # Resolved lazily (the sidecar quant_config.json needs the model dir).
128
+ self._resolved = False
129
+ self._full_config: dict = {}
130
+ self.config_groups: dict = {}
131
+ self.ignore: list[str] = []
132
+ self.target_scheme: dict[str, dict] = {} # CB module -> scheme dict
133
+ self._cb_targets: set[str] = set()
134
+ self.ct_config = None # stock CompressedTensorsConfig
135
+ self._codebooks: dict[str, torch.Tensor] | None = None
136
+
137
+ # -- lazy resolution of the (possibly pointer) quant config --------------
138
+ def _ensure_resolved(self) -> None:
139
+ if self._resolved:
140
+ return
141
+ cfg = self._raw_config
142
+ if "config_groups" not in cfg:
143
+ cfg_file = cfg.get("config_file", "quant_config.json")
144
+ from vllm.config import get_current_vllm_config
145
+ model_dir = get_current_vllm_config().model_config.model
146
+ with open(_resolve_model_file(model_dir, cfg_file)) as fh:
147
+ cfg = json.load(fh)
148
+ self.codebook_file = cfg.get("codebook_file", self.codebook_file)
149
+ # Normalise stored namespaces ONCE, here, so all downstream resolution
150
+ # (ours and the delegated CT config's) sees canonical target names.
151
+ cfg = dict(cfg)
152
+ cfg["config_groups"] = {
153
+ name: {**g, "targets": [_canonical_target(t)
154
+ for t in g.get("targets", [])]}
155
+ for name, g in cfg["config_groups"].items()
156
+ }
157
+ cfg["ignore"] = [_canonical_target(i) for i in cfg.get("ignore", [])]
158
+ self._full_config = cfg
159
+ self.config_groups = cfg["config_groups"]
160
+ self.ignore = list(cfg["ignore"])
161
+ stock_groups: dict = {}
162
+ for name, g in self.config_groups.items():
163
+ if "scheme" in g: # CB group (our vocabulary)
164
+ for t in g["targets"]:
165
+ self.target_scheme[t] = g["scheme"]
166
+ self._cb_targets.add(t)
167
+ else: # stock CT vocabulary
168
+ stock_groups[name] = g
169
+ self._alias_collapsed_shared_prefixes()
170
+ self.ct_config = (self._build_ct_config(stock_groups)
171
+ if stock_groups else None)
172
+ self._resolved = True
173
+
174
+ def _alias_collapsed_shared_prefixes(self) -> None:
175
+ """HunYuan-V3-style shared-expert dispatch collapse. HYV3MoEFused builds
176
+ its shared MLP with ``prefix=f"{prefix}"`` — the ``.shared_mlp`` segment
177
+ never reaches ``get_quant_method``, which instead sees the PARENT-prefix
178
+ names ``…mlp.gate_up_proj`` / ``…mlp.down_proj``. Module paths (params,
179
+ checkpoint tensors) DO keep ``.shared_mlp.``, so only the dispatch key
180
+ collapses. Alias every ``….shared_mlp.<leaf>`` CB target and ignore
181
+ entry to its collapsed form so the CB method owns the shared expert
182
+ natively (packed decode in-kernel) instead of vLLM building plain bf16
183
+ Linears that the loader must fill by decode-at-load. Collision-safe: a
184
+ layer is either dense-MLP (real ``…mlp.<leaf>`` keys, no shared_mlp) or
185
+ MoE-with-shared (no real collapsed keys), and ``setdefault`` keeps any
186
+ real key authoritative. Archs that thread correct shared prefixes are
187
+ unaffected (their aliases match no module prefix). Runs before the
188
+ delegated-CT build so CT's ignore covers the aliases too."""
189
+ for t in [k for k in self.target_scheme if ".shared_mlp." in k]:
190
+ alias = t.replace(".shared_mlp.", ".")
191
+ self.target_scheme.setdefault(alias, self.target_scheme[t])
192
+ self._cb_targets.add(alias)
193
+ self.ignore.extend(ig.replace(".shared_mlp.", ".")
194
+ for ig in list(self.ignore) if ".shared_mlp." in ig)
195
+
196
+ def _build_ct_config(self, stock_groups: dict):
197
+ """A stock CompressedTensorsConfig over the non-CB groups. They are
198
+ already CT vocabulary; we re-key quant_method, add our CB modules to
199
+ CT's ignore (so CT never owns them), and give it a valid top-level
200
+ format (our container's is a CB marker; stock groups carry per-group
201
+ formats that CT reads under "mixed-precision")."""
202
+ from vllm.model_executor.layers.quantization.compressed_tensors.compressed_tensors import ( # noqa: E501
203
+ CompressedTensorsConfig,
204
+ )
205
+ ct_dict = dict(self._full_config)
206
+ ct_dict["quant_method"] = "compressed-tensors"
207
+ ct_dict["config_groups"] = dict(stock_groups)
208
+ ct_dict["ignore"] = list(self.ignore) + sorted(self._cb_targets)
209
+ ct_dict.pop("codebook_file", None)
210
+ ct_dict.pop("provenance", None)
211
+ raw_fmt = str(self._full_config.get("format", ""))
212
+ if raw_fmt in ("", "nvfp4_cb", "fp8_cb", "cb", "mixed-precision"):
213
+ ct_dict["format"] = "mixed-precision"
214
+ return CompressedTensorsConfig.from_config(ct_dict)
215
+
216
+ def __repr__(self) -> str:
217
+ return (f"PrismaQuantConfig(resolved={self._resolved}, "
218
+ f"cb_targets={len(self.target_scheme)}, "
219
+ f"stock_ct={'yes' if self.ct_config is not None else 'no'})")
220
+
221
+ @classmethod
222
+ def get_name(cls):
223
+ return "gridbook"
224
+
225
+ def get_supported_act_dtypes(self) -> list[torch.dtype]:
226
+ return [torch.bfloat16, torch.float16]
227
+
228
+ @classmethod
229
+ def get_min_capability(cls) -> int:
230
+ return 80
231
+
232
+ @classmethod
233
+ def get_config_filenames(cls) -> list[str]:
234
+ return []
235
+
236
+ @classmethod
237
+ def from_config(cls, config: dict[str, Any]) -> "PrismaQuantConfig":
238
+ # Defer parsing: a pointer config resolves quant_config.json lazily.
239
+ return cls(config)
240
+
241
+ @classmethod
242
+ def override_quantization_method(cls, hf_quant_cfg, user_quant, **kwargs):
243
+ # "gridbook" is the registry key going forward; "prismaquant" is the
244
+ # legacy key older local artifacts carry — both dispatch here.
245
+ if user_quant in ("gridbook", "prismaquant"):
246
+ return "gridbook"
247
+ if hf_quant_cfg is not None and \
248
+ hf_quant_cfg.get("quant_method") in ("gridbook", "prismaquant"):
249
+ return "gridbook"
250
+ return None
251
+
252
+ # -- codebook sidecar (loaded once, shared across all layers) ------------
253
+ def get_codebooks(self) -> dict[str, torch.Tensor]:
254
+ if self._codebooks is None:
255
+ from safetensors.torch import load_file
256
+ from vllm.config import get_current_vllm_config
257
+ model_dir = get_current_vllm_config().model_config.model
258
+ self._codebooks = load_file(
259
+ _resolve_model_file(model_dir, self.codebook_file))
260
+ return self._codebooks
261
+
262
+ # -- per-prefix scheme resolution (handles vLLM fused qkv/gate_up) -------
263
+ def _is_ignored(self, prefix: str) -> bool:
264
+ return any(ig in base for base in _candidate_bases(prefix)
265
+ for ig in self.ignore)
266
+
267
+ def shard_target_keys(self, prefix: str, *,
268
+ unfused_fallback: bool = False) -> list[str]:
269
+ """``target_scheme`` keys naming the CB shards of (possibly fused)
270
+ *prefix*, in shard order — ``[]`` if none resolve.
271
+
272
+ THE single owner of fused-shard resolution: ``_scheme_for_prefix``
273
+ (which format does this module decode as?) and
274
+ ``PrismaQuantCBLinearMethod._shard_roles`` (which per-role codebooks
275
+ does it concatenate?) must agree module-for-module, and before issue #1
276
+ they were two hand-rolled copies that had already drifted — the copies
277
+ built their shard keys from the CANONICAL prefix only, so once
278
+ ``apply_vllm_mapper`` moved the stored keys into the mapper's namespace
279
+ a fused GDN ``in_proj_qkvz`` resolved to nothing (silent BF16
280
+ fall-through) and every dense ``_shard_roles`` returned ``[]`` (a
281
+ load-time width assert). Namespace choice is delegated wholesale to
282
+ ``_candidate_bases``.
283
+
284
+ Bases are tried in order and the FIRST base with any hit wins **whole**:
285
+ hits are never mixed across bases, because two vintages of one key can
286
+ name two different on-disk tensors and pairing shards across them would
287
+ silently fuse the wrong weights.
288
+
289
+ ``unfused_fallback`` reproduces ``_shard_roles``' extra ``or [leaf]``
290
+ rung — a plain Linear is its own single role. ``_scheme_for_prefix``
291
+ deliberately omits it (it has already tried the exact keys itself, and
292
+ a bare-leaf retry there would only re-ask the same question).
293
+ """
294
+ pmm = getattr(self, "packed_modules_mapping", {}) or {}
295
+ for base in _candidate_bases(prefix):
296
+ leaf = base.split(".")[-1]
297
+ shard_leaves = pmm.get(leaf) or _FUSED_FALLBACK.get(leaf)
298
+ if shard_leaves is None:
299
+ if not unfused_fallback:
300
+ continue
301
+ shard_leaves = [leaf]
302
+ stem = base[: -len(leaf)]
303
+ hits = [stem + sl for sl in shard_leaves
304
+ if stem + sl in self.target_scheme]
305
+ if hits:
306
+ return hits
307
+ return []
308
+
309
+ def _scheme_for_prefix(self, prefix: str) -> dict | None:
310
+ for base in _candidate_bases(prefix):
311
+ if base in self.target_scheme:
312
+ return self.target_scheme[base]
313
+ schemes = [self.target_scheme[k]
314
+ for k in self.shard_target_keys(prefix)]
315
+ if not schemes:
316
+ return None
317
+ fmt_keys = ("grid", "mode", "k", "n_sub", "type_size")
318
+ sig = {kk: schemes[0][kk] for kk in fmt_keys}
319
+ for s in schemes[1:]:
320
+ if {kk: s[kk] for kk in fmt_keys} != sig:
321
+ raise ValueError(
322
+ f"fused module {prefix} maps to mixed CB decode "
323
+ "formats — export union-find should prevent this")
324
+ return schemes[0]
325
+
326
+ def get_quant_method(self, layer: torch.nn.Module,
327
+ prefix: str) -> "QuantizeMethodBase | None":
328
+ self._ensure_resolved()
329
+ from .linear import PrismaQuantCBLinearMethod
330
+
331
+ # Keep the delegated CT config's fused-module mapping in lockstep.
332
+ if self.ct_config is not None:
333
+ self.ct_config.packed_modules_mapping = getattr(
334
+ self, "packed_modules_mapping", {}) or {}
335
+
336
+ if isinstance(layer, LinearBase):
337
+ # 1) CB target (has a "scheme") — ours (precise, fused-aware; ahead
338
+ # of the substring ignore test).
339
+ scheme = self._scheme_for_prefix(prefix)
340
+ if os.environ.get("PRISMAQUANT_DEBUG_PREFIXES") == "1":
341
+ import sys
342
+ print(f"[pq-prefix] {prefix} -> "
343
+ f"{'CB' if scheme is not None else 'no-scheme'}",
344
+ file=sys.stderr, flush=True)
345
+ if scheme is not None:
346
+ return PrismaQuantCBLinearMethod(self, scheme, prefix)
347
+ # 2) explicitly-ignored -> BF16 passthrough.
348
+ if self._is_ignored(prefix):
349
+ return UnquantizedLinearMethod()
350
+ # 3) stock NVFP4 / FP8_DYNAMIC -> compressed-tensors delegation
351
+ # (canonical prefix — CT targets are serving-namespace names).
352
+ if self.ct_config is not None:
353
+ return self.ct_config.get_quant_method(
354
+ layer, _canonical_prefix(prefix))
355
+ return UnquantizedLinearMethod()
356
+
357
+ if isinstance(layer, VocabParallelEmbedding):
358
+ if self.ct_config is not None:
359
+ method = self.ct_config.get_quant_method(layer, prefix)
360
+ if method is not None:
361
+ return method
362
+ return UnquantizedEmbeddingMethod()
363
+
364
+ # FusedMoE expert stacks (RoutedExperts): a CB expert group -> our MoE
365
+ # method; else delegate to the stock CT MoE path.
366
+ if RoutedExperts is not None and isinstance(layer, RoutedExperts):
367
+ scheme = self._moe_scheme_for_prefix(prefix)
368
+ if scheme is not None:
369
+ from .moe import PrismaQuantCBMoEMethod
370
+ return PrismaQuantCBMoEMethod(
371
+ self, layer.moe_config, scheme, prefix)
372
+ if self.ct_config is not None:
373
+ return self.ct_config.get_quant_method(layer, prefix)
374
+ return None
375
+ return None
376
+
377
+ def _moe_scheme_for_prefix(self, prefix: str) -> dict | None:
378
+ """A CB expert stack (targets like ``…experts.gate_up_proj`` /
379
+ ``…experts.down_proj``) under this FusedMoE prefix — return its scheme
380
+ (uniform per layer, so any matching target's scheme is the layer's)."""
381
+ # Canonicalise BOTH sides, exactly as ``_scheme_for_prefix`` does for
382
+ # Linears. Without this the multimodal wrapper breaks experts ONLY:
383
+ # vLLM hands us the serving prefix ``language_model.model.layers.N.mlp.
384
+ # experts`` while the checkpoint-namespace targets read
385
+ # ``model.language_model.layers.N.mlp.experts.gate_up_proj``, so a raw
386
+ # ``startswith`` misses, no CB MoE method is created, no
387
+ # ``w13_cb_qweight``/``w2_cb_qweight`` params exist, and the arch's own
388
+ # expert mapping then derives ``experts.w2_weight.cb_qweight`` and
389
+ # AttributeErrors (35B CB serve boot). Dense Linears were unaffected
390
+ # because their lookup already canonicalised — that asymmetry WAS the bug.
391
+ #
392
+ # Structurally different from the dense lookup (the TARGET is longer
393
+ # than the prefix here, so this is a ``startswith``, not a key lookup),
394
+ # but the namespace question is the same one — so it comes from the same
395
+ # ``_candidate_bases``, on BOTH sides. Cross-vintage matches are safe
396
+ # here: ``_canonical_prefix`` only rewrites the ``language_model``
397
+ # wrapper, i.e. it renames the SAME module; it can never move a match to
398
+ # a different layer index or leaf.
399
+ bases = _candidate_bases(prefix)
400
+ for name, sch in self.target_scheme.items():
401
+ if name.split(".")[-1] not in _MOE_LEAVES:
402
+ continue
403
+ variants = _candidate_bases(name)
404
+ if any(v.startswith(b) for v in variants for b in bases):
405
+ return sch
406
+ return None
407
+
408
+ def apply_vllm_mapper(self, hf_to_vllm_mapper):
409
+ self._ensure_resolved()
410
+ # vLLM hands us the UNSTACKED mapper (get_unstacked_mapper()), so the
411
+ # q_proj->qkv_proj fusion is NOT rewritten (per-role leaf names survive
412
+ # for _scheme_for_prefix to re-fuse) — but genuine renames/prefixes ARE
413
+ # applied. For hybrid/VLM checkpoints that means the module-nesting
414
+ # prefix (e.g. Qwen3-VL: ``model.language_model.`` -> ``language_model.
415
+ # model.``) must be applied to the CB target keys too: _scheme_for_prefix
416
+ # matches serve-time prefixes EXACTLY (unlike the substring ignore test),
417
+ # so an un-remapped key silently falls through to unquantized and the
418
+ # cb_qweight load then fails ("no parameter named …cb_qweight"). Mirror
419
+ # exactly what the delegated stock-CT config does for its own targets.
420
+ self.ignore = hf_to_vllm_mapper.apply_list(self.ignore)
421
+ self.target_scheme = hf_to_vllm_mapper.apply_dict(self.target_scheme)
422
+ self._cb_targets = set(
423
+ hf_to_vllm_mapper.apply_list(sorted(self._cb_targets)))
424
+ if self.ct_config is not None:
425
+ self.ct_config.apply_vllm_mapper(hf_to_vllm_mapper)