File size: 2,084 Bytes
8e6b0e0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
#!/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()