Download provenance/build_native_checkpoint.py from Lottolabs/gemma-4-12B-it-TT-BFP8-P150: direct link, hf CLI and curl.
- Browser
- Download file 2.08 kB
-
https://huggingface.co/Lottolabs/gemma-4-12B-it-TT-BFP8-P150/resolve/main/provenance/build_native_checkpoint.py
- Command line
-
hf download hf://Lottolabs/gemma-4-12B-it-TT-BFP8-P150/provenance/build_native_checkpoint.py
-
curl -L -o build_native_checkpoint.py https://huggingface.co/Lottolabs/gemma-4-12B-it-TT-BFP8-P150/resolve/main/provenance/build_native_checkpoint.py
2.08 kB
| #!/usr/bin/env python3 | |
| """Build a Gemma-4-12B TT-native checkpoint from the ORIGINAL HF weights. | |
| build_native_checkpoint.py --source <hf model dir named gemma-4-12B-it> --plan <plan.json> | |
| --fetch-verification <fetch-verification.json> --out <root>/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 <out>/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() | |