#!/usr/bin/env python3 """QDQ a packaged Gemma4 assistant cactus weights directory back to HF safetensors.""" from __future__ import annotations import argparse import json import shutil import sys from pathlib import Path import torch from safetensors.torch import save_file sys.path.insert(0, "/workspace/turboquant_sanitized/scripts/export") from cactus_packed_to_qdq_fp16 import ( # noqa: E402 CONFIG_FILES, PRECISION_CQ, PRECISION_FP16, PRECISION_FP32, PRECISION_INT8, dequantize_cq_file, dequantize_fp_file, dequantize_int8_file, read_header, ) DIRECT = { "token_embeddings": "model.embed_tokens.weight", "output_weight": "lm_head.weight", "output_norm": "model.norm.weight", "pre_projection": "pre_projection.weight", "post_projection": "post_projection.weight", "masked_embedding_centroids": "masked_embedding.centroids.weight", } LAYER_SUFFIXES = { "attn_q": "self_attn.q_proj.weight", "attn_output": "self_attn.o_proj.weight", "ffn_gate": "mlp.gate_proj.weight", "ffn_up": "mlp.up_proj.weight", "ffn_down": "mlp.down_proj.weight", "input_norm": "input_layernorm.weight", "attn_q_norm": "self_attn.q_norm.weight", "post_attn_norm": "post_attention_layernorm.weight", "pre_ffn_norm": "pre_feedforward_layernorm.weight", "post_ffn_norm": "post_feedforward_layernorm.weight", "layer_scalar": "layer_scalar", } def hf_key_for_file(path: Path) -> str | None: stem = path.name.removesuffix(".weights") if stem in DIRECT: return DIRECT[stem] parts = stem.split("_", 2) if len(parts) == 3 and parts[0] == "layer" and parts[1].isdigit(): suffix = LAYER_SUFFIXES.get(parts[2]) if suffix: return f"model.layers.{parts[1]}.{suffix}" return None def copy_runtime_files(src: Path, out: Path) -> None: for path in src.iterdir(): if path.is_file() and path.name in CONFIG_FILES: shutil.copy2(path, out / path.name) def load_weight(path: Path, dtype: torch.dtype, row_batch_size: int) -> torch.Tensor: header = read_header(path) if header.precision in PRECISION_CQ: return dequantize_cq_file(path, header, dtype, row_batch_size) if header.precision in {PRECISION_FP16, PRECISION_FP32}: return dequantize_fp_file(path, header, dtype) if header.precision == PRECISION_INT8: return dequantize_int8_file(path, header, dtype) raise ValueError(f"{path.name}: unsupported precision={header.precision}") def load_token_ordering(src: Path) -> torch.Tensor | None: sidecar = src / "masked_embedding_token_ordering.json" if not sidecar.exists(): return None data = json.loads(sidecar.read_text(encoding="utf-8")) return torch.tensor(data["values"], dtype=torch.long).reshape(tuple(data["shape"])) def write_index_if_needed(out: Path, tensors: dict[str, torch.Tensor], shard: str = "model.safetensors") -> None: total = sum(t.numel() * t.element_size() for t in tensors.values()) index = { "metadata": {"total_size": str(total)}, "weight_map": {key: shard for key in sorted(tensors)}, } (out / "model.safetensors.index.json").write_text(json.dumps(index, indent=2) + "\n", encoding="utf-8") def main() -> None: parser = argparse.ArgumentParser() parser.add_argument("--cactus", default="/workspace/gemma4_assistant_quant/gemma4_e2b_it_assistant_cactus_cq4_smoke") parser.add_argument("--out", default="/workspace/gemma4_assistant_quant/gemma4_e2b_it_assistant_cactus_qdq_smoke") parser.add_argument("--dtype", choices=["float16", "bfloat16", "float32"], default="bfloat16") parser.add_argument("--row-batch-size", type=int, default=512) parser.add_argument("--force", action="store_true") args = parser.parse_args() dtype = {"float16": torch.float16, "bfloat16": torch.bfloat16, "float32": torch.float32}[args.dtype] src = Path(args.cactus) out = Path(args.out) if out.exists(): if not args.force: raise SystemExit(f"{out} exists; pass --force") shutil.rmtree(out) out.mkdir(parents=True, exist_ok=True) copy_runtime_files(src, out) tensors: dict[str, torch.Tensor] = {} manifest = [] for path in sorted(src.glob("*.weights")): key = hf_key_for_file(path) if key is None: raise SystemExit(f"no HF key mapping for {path.name}") tensor = load_weight(path, dtype, args.row_batch_size) if key in tensors: raise SystemExit(f"duplicate HF key {key}") tensors[key] = tensor manifest.append({"file": path.name, "hf_key": key, "shape": list(tensor.shape), "dtype": str(tensor.dtype)}) ordering = load_token_ordering(src) if ordering is not None: tensors["masked_embedding.token_ordering"] = ordering manifest.append({ "file": "masked_embedding_token_ordering.json", "hf_key": "masked_embedding.token_ordering", "shape": list(ordering.shape), "dtype": str(ordering.dtype), }) save_file(tensors, out / "model.safetensors") write_index_if_needed(out, tensors) (out / "qdq_manifest.json").write_text(json.dumps(manifest, indent=2) + "\n", encoding="utf-8") summary = { "source": str(src), "out": str(out), "tensor_count": len(tensors), "dtype": args.dtype, "bytes": sum(t.numel() * t.element_size() for t in tensors.values()), } (out / "qdq_summary.json").write_text(json.dumps(summary, indent=2) + "\n", encoding="utf-8") print(json.dumps(summary, indent=2)) if __name__ == "__main__": main()