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 +18 -0
- gridbook/codec.py +145 -0
- gridbook/config.py +425 -0
- gridbook/csrc/cb_fused_gemm.cu +734 -0
- gridbook/csrc/cb_gemv.cu +1905 -0
- gridbook/csrc/cb_persistent_prefill.cu +201 -0
- gridbook/csrc/cb_persistent_tc.cu +353 -0
- gridbook/csrc/cutlass_fork/sm120_cb_fused_mma.hpp +670 -0
- gridbook/csrc/cutlass_fork/sm120_cb_mma_tma.hpp +613 -0
- gridbook/csrc/cutlass_fork/sm120_cb_persistent_mma.hpp +197 -0
- gridbook/csrc/cutlass_fork/sm120_expert_row_broadcast.hpp +344 -0
- gridbook/csrc/cutlass_fork/sm120_mma_tma_orig.hpp +587 -0
- gridbook/csrc/sm120_fp8_gemm.cu +131 -0
- gridbook/csrc/smem_probe_tilem.cu +163 -0
- gridbook/csrc/toolchain_probe.cu +30 -0
- gridbook/cuda_ext.py +250 -0
- gridbook/expand.py +402 -0
- gridbook/kernels.py +230 -0
- gridbook/linear.py +492 -0
- gridbook/moe.py +1988 -0
- gridbook/moe_autotune.py +157 -0
- gridbook/moe_l2.py +211 -0
- gridbook/moe_routing.py +78 -0
- gridbook/moe_toplevel_loader.py +598 -0
- gridbook/ops.py +303 -0
- gridbook/plugin.py +143 -0
- gridbook-0.1.0.dist-info/METADATA +353 -0
- gridbook-0.1.0.dist-info/RECORD +32 -0
- gridbook-0.1.0.dist-info/WHEEL +5 -0
- gridbook-0.1.0.dist-info/entry_points.txt +2 -0
- gridbook-0.1.0.dist-info/licenses/LICENSE +201 -0
- gridbook-0.1.0.dist-info/top_level.txt +1 -0
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)
|