kmoss commited on
Commit
18537a3
·
verified ·
1 Parent(s): 342466c

Add Gemma4 assistant CQ artifact assistant-qdq.py

Browse files
Files changed (1) hide show
  1. assistant-qdq.py +154 -0
assistant-qdq.py ADDED
@@ -0,0 +1,154 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """QDQ a packaged Gemma4 assistant cactus weights directory back to HF safetensors."""
3
+ from __future__ import annotations
4
+
5
+ import argparse
6
+ import json
7
+ import shutil
8
+ import sys
9
+ from pathlib import Path
10
+
11
+ import torch
12
+ from safetensors.torch import save_file
13
+
14
+ sys.path.insert(0, "/workspace/turboquant_sanitized/scripts/export")
15
+ from cactus_packed_to_qdq_fp16 import ( # noqa: E402
16
+ CONFIG_FILES,
17
+ PRECISION_CQ,
18
+ PRECISION_FP16,
19
+ PRECISION_FP32,
20
+ PRECISION_INT8,
21
+ dequantize_cq_file,
22
+ dequantize_fp_file,
23
+ dequantize_int8_file,
24
+ read_header,
25
+ )
26
+
27
+
28
+ DIRECT = {
29
+ "token_embeddings": "model.embed_tokens.weight",
30
+ "output_weight": "lm_head.weight",
31
+ "output_norm": "model.norm.weight",
32
+ "pre_projection": "pre_projection.weight",
33
+ "post_projection": "post_projection.weight",
34
+ "masked_embedding_centroids": "masked_embedding.centroids.weight",
35
+ }
36
+
37
+ LAYER_SUFFIXES = {
38
+ "attn_q": "self_attn.q_proj.weight",
39
+ "attn_output": "self_attn.o_proj.weight",
40
+ "ffn_gate": "mlp.gate_proj.weight",
41
+ "ffn_up": "mlp.up_proj.weight",
42
+ "ffn_down": "mlp.down_proj.weight",
43
+ "input_norm": "input_layernorm.weight",
44
+ "attn_q_norm": "self_attn.q_norm.weight",
45
+ "post_attn_norm": "post_attention_layernorm.weight",
46
+ "pre_ffn_norm": "pre_feedforward_layernorm.weight",
47
+ "post_ffn_norm": "post_feedforward_layernorm.weight",
48
+ "layer_scalar": "layer_scalar",
49
+ }
50
+
51
+
52
+ def hf_key_for_file(path: Path) -> str | None:
53
+ stem = path.name.removesuffix(".weights")
54
+ if stem in DIRECT:
55
+ return DIRECT[stem]
56
+ parts = stem.split("_", 2)
57
+ if len(parts) == 3 and parts[0] == "layer" and parts[1].isdigit():
58
+ suffix = LAYER_SUFFIXES.get(parts[2])
59
+ if suffix:
60
+ return f"model.layers.{parts[1]}.{suffix}"
61
+ return None
62
+
63
+
64
+ def copy_runtime_files(src: Path, out: Path) -> None:
65
+ for path in src.iterdir():
66
+ if path.is_file() and path.name in CONFIG_FILES:
67
+ shutil.copy2(path, out / path.name)
68
+
69
+
70
+ def load_weight(path: Path, dtype: torch.dtype, row_batch_size: int) -> torch.Tensor:
71
+ header = read_header(path)
72
+ if header.precision in PRECISION_CQ:
73
+ return dequantize_cq_file(path, header, dtype, row_batch_size)
74
+ if header.precision in {PRECISION_FP16, PRECISION_FP32}:
75
+ return dequantize_fp_file(path, header, dtype)
76
+ if header.precision == PRECISION_INT8:
77
+ return dequantize_int8_file(path, header, dtype)
78
+ raise ValueError(f"{path.name}: unsupported precision={header.precision}")
79
+
80
+
81
+ def load_token_ordering(src: Path) -> torch.Tensor | None:
82
+ sidecar = src / "masked_embedding_token_ordering.json"
83
+ if not sidecar.exists():
84
+ return None
85
+ data = json.loads(sidecar.read_text(encoding="utf-8"))
86
+ return torch.tensor(data["values"], dtype=torch.long).reshape(tuple(data["shape"]))
87
+
88
+
89
+ def write_index_if_needed(out: Path, tensors: dict[str, torch.Tensor], shard: str = "model.safetensors") -> None:
90
+ total = sum(t.numel() * t.element_size() for t in tensors.values())
91
+ index = {
92
+ "metadata": {"total_size": str(total)},
93
+ "weight_map": {key: shard for key in sorted(tensors)},
94
+ }
95
+ (out / "model.safetensors.index.json").write_text(json.dumps(index, indent=2) + "\n", encoding="utf-8")
96
+
97
+
98
+ def main() -> None:
99
+ parser = argparse.ArgumentParser()
100
+ parser.add_argument("--cactus", default="/workspace/gemma4_assistant_quant/gemma4_e2b_it_assistant_cactus_cq4_smoke")
101
+ parser.add_argument("--out", default="/workspace/gemma4_assistant_quant/gemma4_e2b_it_assistant_cactus_qdq_smoke")
102
+ parser.add_argument("--dtype", choices=["float16", "bfloat16", "float32"], default="bfloat16")
103
+ parser.add_argument("--row-batch-size", type=int, default=512)
104
+ parser.add_argument("--force", action="store_true")
105
+ args = parser.parse_args()
106
+
107
+ dtype = {"float16": torch.float16, "bfloat16": torch.bfloat16, "float32": torch.float32}[args.dtype]
108
+ src = Path(args.cactus)
109
+ out = Path(args.out)
110
+ if out.exists():
111
+ if not args.force:
112
+ raise SystemExit(f"{out} exists; pass --force")
113
+ shutil.rmtree(out)
114
+ out.mkdir(parents=True, exist_ok=True)
115
+ copy_runtime_files(src, out)
116
+
117
+ tensors: dict[str, torch.Tensor] = {}
118
+ manifest = []
119
+ for path in sorted(src.glob("*.weights")):
120
+ key = hf_key_for_file(path)
121
+ if key is None:
122
+ raise SystemExit(f"no HF key mapping for {path.name}")
123
+ tensor = load_weight(path, dtype, args.row_batch_size)
124
+ if key in tensors:
125
+ raise SystemExit(f"duplicate HF key {key}")
126
+ tensors[key] = tensor
127
+ manifest.append({"file": path.name, "hf_key": key, "shape": list(tensor.shape), "dtype": str(tensor.dtype)})
128
+
129
+ ordering = load_token_ordering(src)
130
+ if ordering is not None:
131
+ tensors["masked_embedding.token_ordering"] = ordering
132
+ manifest.append({
133
+ "file": "masked_embedding_token_ordering.json",
134
+ "hf_key": "masked_embedding.token_ordering",
135
+ "shape": list(ordering.shape),
136
+ "dtype": str(ordering.dtype),
137
+ })
138
+
139
+ save_file(tensors, out / "model.safetensors")
140
+ write_index_if_needed(out, tensors)
141
+ (out / "qdq_manifest.json").write_text(json.dumps(manifest, indent=2) + "\n", encoding="utf-8")
142
+ summary = {
143
+ "source": str(src),
144
+ "out": str(out),
145
+ "tensor_count": len(tensors),
146
+ "dtype": args.dtype,
147
+ "bytes": sum(t.numel() * t.element_size() for t in tensors.values()),
148
+ }
149
+ (out / "qdq_summary.json").write_text(json.dumps(summary, indent=2) + "\n", encoding="utf-8")
150
+ print(json.dumps(summary, indent=2))
151
+
152
+
153
+ if __name__ == "__main__":
154
+ main()