#!/usr/bin/env python3 """Build a Gemma-4-12B TT-native checkpoint from the ORIGINAL HF weights. build_native_checkpoint.py --source --plan --fetch-verification --out /gemma-4-12B-it The runtime's own weight path (create_tt_model -> ttnn.as_tensor) quantizes every tensor per the plan and serializes the exact host TTNN tensors into /tensors; then native_checkpoint.finalize writes the lossless small tensors, metadata and manifest. Needs the device (model construction uploads to DRAM). No proof is written here: run the proof recipe (prove_native.sh) afterwards. """ import argparse import json import os import shutil import sys from pathlib import Path sys.path.insert(0, str(Path(__file__).resolve().parent)) import native_checkpoint as nc # noqa: E402 def main(): ap = argparse.ArgumentParser() ap.add_argument("--source", required=True) ap.add_argument("--plan", required=True) ap.add_argument("--fetch-verification", required=True) ap.add_argument("--out", required=True) a = ap.parse_args() out = Path(a.out) if "12b" not in out.name.lower(): raise SystemExit("checkpoint dir basename must contain 12B (runtime policy key)") if out.exists(): raise SystemExit(f"{out} exists; refusing to overwrite") (out / "tensors").mkdir(parents=True) shutil.copyfile(a.plan, out / "precision_plan.json") os.environ["TT_CACHE_PATH"] = str(out / "tensors") os.environ["GEMMA4_PRECISION_PLAN"] = str(out / "precision_plan.json") import ttnn from models.demos.gemma4.tt.common import create_tt_model mesh = ttnn.open_mesh_device(mesh_shape=ttnn.MeshShape(1, 1)) try: create_tt_model(mesh_device=mesh, max_batch_size=1, max_seq_len=2048, model_path=a.source, create_kv_cache=False) finally: ttnn.close_mesh_device(mesh) print(json.dumps(nc.finalize(out, a.source, a.fetch_verification), indent=1)) print("BUILD_DONE") if __name__ == "__main__": main()