HAKO-v1 / run_all.py
PowerMachine's picture
HAKO upload: run_all.py
8433089 verified
Raw History Blame Contribute Delete
11.5 kB
#!/usr/bin/env python3
"""HAKO end-to-end runner: phases 0-4 under a global wall-clock budget.
Phase 0 environment + math property suite (fail-fast; no training if any
theorem check fails -- proofs come BEFORE code execution).
Phase 1 Qwen2.5-0.5B: fetch -> decompose (int4 kept) -> Z -> K-Means++ ->
plateau -> GHSOM growth -> head pretraining.
Phase 2 Granite-4.0-1B: decompose -> register -> joint cooperative fusion
with MoE cross-attention, chain recursion, diffusion game loop.
Phase 3 NLP/NLG: byte-BPE (Lemma-1 verified) + conditioned decoder on
ultra-alpaca-ptbr + corpus-ptbr-v1 + EN + zh.
Phase 4 final validation, checkpoint verification, HF publish (only if
every phase succeeded and no token leak is detected).
Usage: python3 run_all.py [--phase N] [--smoke]
"""
from __future__ import annotations
import argparse
import json
import logging
import os
import sys
import time
from pathlib import Path
ROOT = Path(__file__).resolve().parent
sys.path.insert(0, str(ROOT))
# Managed caches: keep every HF artifact inside the project tree so the
# MemoryManager can sweep them (never grow the user's global cache).
os.environ.setdefault("HF_HOME", "/home/z/my-project/work/hf_home")
os.environ.setdefault("HF_DATASETS_CACHE", "/home/z/my-project/work/hf_home/datasets")
os.environ.setdefault("TMPDIR", "/home/z/my-project/work/tmp")
logging.basicConfig(level=logging.INFO,
format="%(asctime)s %(name)s %(levelname)s %(message)s")
log = logging.getLogger("hako.run")
def main() -> int:
ap = argparse.ArgumentParser()
ap.add_argument("--phase", type=int, default=0)
ap.add_argument("--start", type=int, default=1,
help="first training phase to run (1..3); earlier "
"phases are restored from checkpoints")
ap.add_argument("--smoke", action="store_true",
help="short smoke budget for CI-like validation")
args = ap.parse_args()
from hako.config import CONFIG
from hako.memory_manager import MemoryManager
from hako.telemetry import Telemetry
cfg = CONFIG
cfg.ensure_dirs()
mem = MemoryManager(cfg)
tel = Telemetry(cfg)
hw = cfg.derived()
log.info("hardware: cpus=%d ram=%.2fGiB disk_free=%.2fGiB",
hw.cpu_count, hw.ram_total / 2**30, hw.disk_free / 2**30)
smoke = args.smoke
if smoke:
cfg.budget_phase1 = 240
cfg.budget_phase2 = 180
cfg.budget_phase3 = 200
cfg.budget_phase4 = 60
failures = 0
# ---------------- phase 0: math property suite ------------------------
if args.phase <= 0:
log.info("=== PHASE 0: mathematical property verification ===")
import subprocess
r = subprocess.run([sys.executable,
str(ROOT / "tests" / "test_math_properties.py")],
capture_output=True, text=True)
print(r.stdout[-4000:])
if r.returncode != 0:
log.error("math property suite FAILED -- aborting (proofs first)")
tel.orchestration(event="phase0_failed")
return 2
tel.orchestration(event="phase0_math_verified")
# ---------------- phase 0.5: corpus cache (isolated RAM) --------------
cache_file = Path(cfg.work_dir) / "cache" / "corpus_ptbr.jsonl"
if not cache_file.exists():
import subprocess
cache_file.parent.mkdir(parents=True, exist_ok=True)
log.info("extracting corpus-ptbr-v1 cache (isolated subprocess)...")
r = subprocess.run(
[sys.executable, "-m", "hako.sources.corpus_cache",
str(cache_file), "6000"],
cwd=str(ROOT), capture_output=True, text=True, timeout=600)
if r.returncode != 0:
log.warning("corpus cache extraction failed: %s",
r.stderr[-300:])
else:
log.info("corpus cache ready (%s)", cache_file.name)
tel.orchestration(event="corpus_cache_ready")
carry: dict = {}
# ---------------- phase 1 ---------------------------------------------
if args.start <= 1 and args.phase <= 1:
from hako.train.phase1_qwen import run as run1
log.info("=== PHASE 1: Qwen source (budget %ds) ===", cfg.budget_phase1)
try:
carry = run1(cfg, tel, mem, cfg.budget_phase1)
tel.orchestration(event="phase1_ok",
steps=int(carry.get("steps", 0)))
except Exception as exc: # noqa: BLE001
log.exception("phase 1 failed")
tel.orchestration(event="phase1_failed", error=str(exc)[:300])
failures += 1
if smoke:
return 3
elif args.start == 2 and failures == 0:
# resume: rebuild the system from the phase-1 checkpoint
import numpy as _np
import torch as _torch
from hako.sources import loader as _loader
from hako.train.restore import restore_system
ck1 = Path(cfg.ckpt_dir) / "hako_phase1.npz"
zw = Path(cfg.artifacts_dir) / "phase1_Z_workspace.npz"
if ck1.exists() and zw.exists():
sys_, _ = restore_system(cfg, tel, mem, ck1)
data = _np.load(zw)
from hako.nlp.datasets import Corpus
corpus = Corpus()
docs = _loader.fetch_datasets_streaming(
cfg, mem, {"pt_alpaca": 9000, "en": 6000, "zh": 2200})
for lang, dl in docs.items():
corpus.add_docs(lang, dl)
cache_file = Path(cfg.work_dir) / "cache" / "corpus_ptbr.jsonl"
if cache_file.exists():
import json as _json
with open(cache_file, "r", encoding="utf-8") as fh:
corpus.add_docs("pt", [_json.loads(l)["text"]
for l in fh])
carry = {"system": sys_,
"Z_t": _torch.as_tensor(data["Z_src"]),
"Z_N": _torch.as_tensor(data["Z_N"]),
"labels": data["labels"], "corpus": corpus,
"checkpoint": str(ck1)}
tel.orchestration(event="phase1_restored", step=int(sys_.step))
else:
log.error("--start 2 but phase-1 checkpoint missing")
failures += 1
# ---------------- phase 2 ---------------------------------------------
if args.start <= 2 and args.phase <= 2 and failures == 0 and \
carry.get("system") is not None:
from hako.train.phase2_granite_joint import run as run2
log.info("=== PHASE 2: Granite joint fusion (budget %ds) ===",
cfg.budget_phase2)
try:
carry = run2(cfg, tel, mem, cfg.budget_phase2, carry)
tel.orchestration(event="phase2_ok")
except Exception as exc: # noqa: BLE001
log.exception("phase 2 failed")
tel.orchestration(event="phase2_failed", error=str(exc)[:300])
failures += 1
elif args.start == 3 and failures == 0:
import numpy as _np
import torch as _torch
from hako.sources import loader as _loader
from hako.train.restore import restore_system
ck2 = Path(cfg.ckpt_dir) / "hako_phase2.npz"
zw = Path(cfg.artifacts_dir) / "phase1_Z_workspace.npz"
if ck2.exists() and zw.exists():
sys_, _ = restore_system(cfg, tel, mem, ck2)
data = _np.load(zw)
from hako.nlp.datasets import Corpus
corpus = Corpus()
docs = _loader.fetch_datasets_streaming(
cfg, mem, {"pt_alpaca": 9000, "en": 6000, "zh": 2200})
for lang, dl in docs.items():
corpus.add_docs(lang, dl)
cache_file = Path(cfg.work_dir) / "cache" / "corpus_ptbr.jsonl"
if cache_file.exists():
import json as _json
with open(cache_file, "r", encoding="utf-8") as fh:
corpus.add_docs("pt", [_json.loads(l)["text"]
for l in fh])
from hako.sources.decompose import dequantize_int4
gnpz = _np.load(Path(cfg.artifacts_dir) /
"granite1b_decomposed.npz")
protos = dequantize_int4(gnpz["int4_proto_packed"],
gnpz["int4_proto_scales"],
int(gnpz["proto_dim_for_unpack"][0]))
carry = {"system": sys_,
"Z_t": _torch.as_tensor(data["Z_src"]),
"Z_N": _torch.as_tensor(data["Z_N"]),
"labels": data["labels"], "corpus": corpus,
"granite_protos": _torch.as_tensor(protos),
"checkpoint": str(ck2)}
tel.orchestration(event="phase2_restored", step=int(sys_.step))
else:
log.error("--start 3 but phase-2 checkpoint missing")
failures += 1
# ---------------- phase 3 ---------------------------------------------
if args.start <= 3 and args.phase <= 3 and failures == 0 and \
carry.get("system") is not None:
from hako.train.phase3_nlp import run as run3
log.info("=== PHASE 3: NLP/NLG (budget %ds) ===", cfg.budget_phase3)
try:
carry = run3(cfg, tel, mem, cfg.budget_phase3, carry)
tel.orchestration(event="phase3_ok")
except Exception as exc: # noqa: BLE001
log.exception("phase 3 failed")
tel.orchestration(event="phase3_failed", error=str(exc)[:300])
failures += 1
# ---------------- phase 4: validation + publish -----------------------
if args.phase <= 4:
log.info("=== PHASE 4: final validation & publish ===")
from hako.publish.push_hf import publish, security_scan
viol = security_scan(ROOT)
if viol:
log.error("token leak detected in %d files -- NOT publishing", viol)
failures += 1
if failures == 0:
summary = {
"nlg_val_loss": carry.get("nlg_val_loss"),
"nlg_sample": (carry.get("nlg_sample") or "")[:300],
"telemetry": tel.summary(),
"hw": {"cpus": hw.cpu_count,
"ram_GiB": round(hw.ram_total / 2**30, 2)},
"published_from": str(ROOT),
}
(ROOT / "artifacts" / "final_summary.json").write_text(
json.dumps(summary, indent=2, default=str), encoding="utf-8")
repo = os.environ.get("HAKO_REPO", cfg.hf_repo)
if os.environ.get("HF_TOKEN"):
try:
url = publish(ROOT, repo, tel.summary())
log.info("PUBLISHED: %s", url)
tel.orchestration(event="published", url=url)
except Exception as exc: # noqa: BLE001
log.error("publish failed: %s", exc)
failures += 1
else:
log.warning("HF_TOKEN not set -- skipping publish "
"(state saved locally)")
# final state save
if carry.get("system") is not None:
from hako.checkpoint import save_state
save_state(ROOT / "checkpoints" / "hako_final.npz",
**carry["system"].state_blocks())
log.info("run finished with %d failure(s)", failures)
return 1 if failures else 0
if __name__ == "__main__":
sys.exit(main())