| |
| """Convert BF16 attention / shared_experts / dense-MLP linears to FP8 e4m3. |
| |
| Targets (layers 0..77 only; layer 78 = MTP is copied verbatim): |
| - model.layers.N.self_attn.{q_a_proj,q_b_proj,kv_a_proj_with_mqa, |
| kv_b_proj,o_proj}.weight (discovered from the index; indexer.* and |
| *layernorm* are never touched) |
| - model.layers.N.mlp.shared_experts.{gate_proj,up_proj,down_proj}.weight |
| - model.layers.{0,1,2}.mlp.{gate_proj,up_proj,down_proj}.weight |
| |
| Scale convention: PER-TENSOR (one fp32 scalar per module, stored as |
| <prefix>.weight_scale shape [1]). vLLM's Fp8LinearMethod (non-block path) |
| registers PerTensorScaleParameter and requantizes fused shards with the max |
| scale; per-output-channel scales are NOT loadable there, so we fall back to |
| per-tensor (correctness over granularity). |
| |
| Modules that vLLM fuses (q_a_proj + kv_a_proj_with_mqa -> fused_qkv_a_proj, |
| gate_proj + up_proj -> gate_up_proj) share one scale (the group max) so the |
| load-time requantize_with_max_scale is a bit-exact identity instead of a |
| second lossy rounding. |
| |
| Never modifies SRC; writes DST fresh. |
| |
| Usage: convert_fp8_attn.py SRC DST |
| """ |
| import json |
| import os |
| import re |
| import shutil |
| import sys |
|
|
| import torch |
| from safetensors import safe_open |
| from safetensors.torch import save_file |
|
|
| SRC = sys.argv[1].rstrip("/") |
| DST = sys.argv[2].rstrip("/") |
| SHARD_BYTES = 4 << 30 |
| FP8_MAX = 448.0 |
| NUM_DECODER_LAYERS = 78 |
|
|
| assert os.path.abspath(DST) != os.path.abspath(SRC), "refuse overwrite" |
| assert not os.path.exists(os.path.join(DST, "model.safetensors.index.json")), \ |
| "DST already has a checkpoint" |
|
|
| wm = json.load(open(f"{SRC}/model.safetensors.index.json"))["weight_map"] |
|
|
| |
| attn_projs = set() |
| for k in wm: |
| m = re.match(r"model\.layers\.(\d+)\.self_attn\.([\w.]+)\.weight$", k) |
| if not m or int(m.group(1)) >= NUM_DECODER_LAYERS: |
| continue |
| sub = m.group(2) |
| if sub.startswith("indexer") or "norm" in sub: |
| continue |
| attn_projs.add(sub) |
| |
| |
| |
| if len(sys.argv) > 3: |
| want = set() if sys.argv[3] == "none" else set(sys.argv[3].split(",")) |
| attn_projs &= want |
| print("attention projections found:", sorted(attn_projs)) |
|
|
| MLP_PROJS = ("gate_proj", "up_proj", "down_proj") |
|
|
|
|
| def target_prefix(name: str) -> str | None: |
| """Return module prefix (name minus .weight) if this tensor is a target.""" |
| m = re.match(r"model\.layers\.(\d+)\.(.+)\.weight$", name) |
| if not m or int(m.group(1)) >= NUM_DECODER_LAYERS: |
| return None |
| li, rest = int(m.group(1)), m.group(2) |
| if rest.startswith("self_attn.") and rest[len("self_attn."):] in attn_projs: |
| return name[: -len(".weight")] |
| if rest.startswith("mlp.shared_experts.") and \ |
| rest[len("mlp.shared_experts."):] in MLP_PROJS: |
| return name[: -len(".weight")] |
| if li < 3 and rest.startswith("mlp.") and rest[len("mlp."):] in MLP_PROJS: |
| return name[: -len(".weight")] |
| return None |
|
|
|
|
| class Writer: |
| def __init__(self): |
| os.makedirs(DST, exist_ok=True) |
| self.cur, self.cur_bytes, self.n, self.total = {}, 0, 0, 0 |
| self.weight_map, self.files = {}, [] |
|
|
| def add(self, name, t): |
| nb = t.numel() * t.element_size() |
| if self.cur_bytes + nb > SHARD_BYTES and self.cur: |
| self.flush() |
| self.cur[name] = t |
| self.cur_bytes += nb |
| self.total += nb |
|
|
| def flush(self): |
| if not self.cur: |
| return |
| self.n += 1 |
| f = f"model-{self.n:05d}.safetensors" |
| save_file(self.cur, f"{DST}/{f}") |
| for k in self.cur: |
| self.weight_map[k] = f |
| self.files.append(f) |
| self.cur, self.cur_bytes = {}, 0 |
|
|
| def finalize(self): |
| self.flush() |
| out = {} |
| for i, f in enumerate(self.files, 1): |
| new = f"model-{i:05d}-of-{self.n:05d}.safetensors" |
| os.rename(f"{DST}/{f}", f"{DST}/{new}") |
| for k, v in self.weight_map.items(): |
| if v == f: |
| out[k] = new |
| json.dump({"metadata": {"total_size": self.total}, |
| "weight_map": out}, |
| open(f"{DST}/model.safetensors.index.json", "w"), indent=0) |
| print(f"index: {len(out)} tensors, {self.total/2**30:.2f} GiB") |
|
|
|
|
| |
| |
| _PARTNER = {"q_a_proj": "kv_a_proj_with_mqa", "kv_a_proj_with_mqa": "q_a_proj", |
| "gate_proj": "up_proj", "up_proj": "gate_proj"} |
|
|
|
|
| def partner_of(pfx: str) -> str | None: |
| parent, _, leaf = pfx.rpartition(".") |
| other = _PARTNER.get(leaf) |
| return f"{parent}.{other}" if other else None |
|
|
|
|
| |
| amax = {} |
| for shard in sorted(set(wm.values())): |
| with safe_open(f"{SRC}/{shard}", framework="pt") as f: |
| for name in f.keys(): |
| pfx = target_prefix(name) |
| if pfx is not None: |
| t = f.get_tensor(name) |
| assert t.dtype == torch.bfloat16 and t.dim() == 2, \ |
| (name, t.dtype, t.shape) |
| amax[pfx] = t.abs().amax().float() |
| scales = {} |
| for pfx, a in amax.items(): |
| p = partner_of(pfx) |
| if p is not None: |
| a = torch.maximum(a, amax[p]) |
| scales[pfx] = (a / FP8_MAX).clamp(min=1e-12) |
| print(f"pass 1: {len(scales)} target modules") |
|
|
|
|
| def quantize(name: str, w: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: |
| """Per-tensor (fusion-group-shared) fp8 e4m3: (w_fp8, scale fp32 [1]).""" |
| s = scales[name[: -len(".weight")]] |
| q = (w.float() / s).clamp(-FP8_MAX, FP8_MAX).to(torch.float8_e4m3fn) |
| return q, s.reshape(1) |
|
|
|
|
| w = Writer() |
| fp8_modules, bytes_before, bytes_after = [], 0, 0 |
| for shard in sorted(set(wm.values())): |
| with safe_open(f"{SRC}/{shard}", framework="pt") as f: |
| for name in sorted(f.keys()): |
| t = f.get_tensor(name) |
| pfx = target_prefix(name) |
| if pfx is None: |
| w.add(name, t) |
| continue |
| q, s = quantize(name, t) |
| w.add(name, q) |
| w.add(f"{pfx}.weight_scale", s) |
| fp8_modules.append(pfx) |
| bytes_before += t.numel() * t.element_size() |
| bytes_after += q.numel() + s.numel() * 4 |
| print(f"{shard} done", flush=True) |
| w.finalize() |
|
|
| fp8_modules.sort() |
| assert len(fp8_modules) == len(set(fp8_modules)) |
| cfg = json.load(open(f"{SRC}/config.json")) |
| cfg["quantization_config"]["fp8_modules"] = fp8_modules |
| json.dump(cfg, open(f"{DST}/config.json", "w"), indent=2) |
|
|
| for f in os.listdir(SRC): |
| if (f.endswith(".json") |
| and f not in ("config.json", "model.safetensors.index.json") |
| or f.endswith((".txt", ".jinja", ".py", ".md", ".sh"))): |
| shutil.copy2(f"{SRC}/{f}", f"{DST}/{f}") |
|
|
| print(f"converted modules: {len(fp8_modules)}") |
| by_kind = {} |
| for p in fp8_modules: |
| by_kind[p.split(".")[-1]] = by_kind.get(p.split(".")[-1], 0) + 1 |
| print("by kind:", by_kind) |
| print(f"bytes converted: {bytes_before/2**30:.2f} GiB -> " |
| f"{bytes_after/2**30:.2f} GiB") |
| print("DONE:", DST) |
|
|