nkkbr/Mini-K3-1H-kda-kernel-64-v2
01.9k
1#!/usr/bin/env python32"""Create the full module graph on meta, or explicitly allocate random weights."""3 4from __future__ import annotations5 6import argparse7from pathlib import Path8 9import torch10 11from configuration_mini_k3 import MiniK3Config12from modeling_mini_k3 import MiniK3ForCausalLM, count_logical_parameters13 14 15def main() -> None:16 parser = argparse.ArgumentParser()17 parser.add_argument("--config", type=Path, default=Path(__file__).with_name("config.json"))18 parser.add_argument(19 "--materialize",20 action="store_true",21 help="Actually allocate and randomly initialize the full model on CPU.",22 )23 parser.add_argument("--seed", type=int, default=1234, help="Random seed for initialization")24 parser.add_argument("--output", type=Path, help="Optional torch.save state-dict path")25 args = parser.parse_args()26 27 config = MiniK3Config.from_json(args.config)28 model = MiniK3ForCausalLM(config, device="meta")29 print(f"meta model defined: {count_logical_parameters(model):,} logical parameters")30 if not args.materialize:31 print("No storage allocated. Pass --materialize only on a host with enough CPU RAM.")32 return33 34 torch.manual_seed(args.seed)35 model.materialize_and_initialize("cpu")36 print(f"Full model randomly initialized on CPU with seed {args.seed}")37 if args.output is not None:38 torch.save({"config": config.to_dict(), "model": model.state_dict()}, args.output)39 print(f"saved: {args.output}")40 41 42if __name__ == "__main__":43 main()44 