smartcore-v1 / code /kod /faz3_ema.py
kdirgul's picture
post-hoc checkpoint EMA (IMU-1): son N checkpoints_350m weight-avg -> ema/ckpt.pt (faz5 kiyasi), saf-mantik 3/3
89bb297 verified
Raw
History Blame Contribute Delete
4.4 kB
"""
Faz 3 sonrası — post-hoc checkpoint EMA (weight averaging). IMU-1 (arXiv 2602.02522): sub-500M'de bedava kalite↑.
checkpoints_350m son N ckpt'nin MODEL ağırlıklarını ortalar (uniform=soup / --decay=EMA son-ağır) →
ema/ckpt.pt kaydeder+push. faz5 `--ckpt ema/ckpt.pt` ile eval → tek-final vs EMA kıyası (faz5 GATE'i).
SADECE model ağırlıkları ortalanır (optimizer state DEĞİL); cfg en son ckpt'ten; fp32 topla→bf16 kaydet.
CPU yeter (GPU gerekmez). Yalnız aynı-mimari (aynı koşu) ckpt'ler → anahtar/şekil birebir.
Kullanım:
HF_TOKEN=hf_xxx python faz3_ema.py --subdir checkpoints_350m --last 5
HF_TOKEN=hf_xxx python faz3_ema.py --last 5 --decay 0.6 # EMA: son ckpt daha ağır
"""
import os, sys, argparse
import torch
REPO = "kdirgul/smartcore-v1"
# ───────────── saf-mantık (yerelde test edilebilir) ─────────────
def ema_weights(n, decay):
"""n ckpt için normalize ağırlık. decay=0 → uniform (soup); >0 → w_i ∝ decay^(n-1-i) (son ckpt en ağır)."""
w = [decay ** (n - 1 - i) for i in range(n)] if decay > 0 else [1.0] * n
s = sum(w)
return [x / s for x in w]
def average_states(states, weights):
"""state_dict listesi + normalize ağırlıklar → ağırlıklı-ortalama. Float fp32'de topla→bf16; non-float sonuncu."""
avg, nonfloat = {}, {}
for wi, sd in zip(weights, states):
for k, v in sd.items():
if v.is_floating_point():
avg[k] = (v.float() * wi) if k not in avg else (avg[k] + v.float() * wi)
else:
nonfloat[k] = v # int/bool buffer → ortalama anlamsız, sonuncuyu tut
out = {k: v.to(torch.bfloat16) for k, v in avg.items()}
out.update(nonfloat)
return out
# ───────────── HF ─────────────
def list_ckpts(api, subdir):
fs = [f for f in api.list_repo_files(REPO, repo_type="model")
if f.startswith(f"{subdir}/step_") and f.endswith("ckpt.pt")]
return sorted(fs) # step_NNNNNN sıfır-dolgulu → lexikografik = sayısal
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--subdir", default="checkpoints_350m", help="pretrain ckpt namespace")
ap.add_argument("--last", type=int, default=5, help="son N ckpt ortalanır")
ap.add_argument("--decay", type=float, default=0.0, help="0=uniform(soup); >0 EMA (son ckpt en ağır)")
ap.add_argument("--out_subdir", default="ema", help="HF çıktı alt-dizini (ema/ckpt.pt)")
ap.add_argument("--out_local", default="/content/ema/ckpt.pt")
ap.add_argument("--no_push", action="store_true")
args = ap.parse_args()
from huggingface_hub import HfApi, hf_hub_download, get_token
token = os.environ.get("HF_TOKEN") or get_token()
api = HfApi(token=token)
fs = list_ckpts(api, args.subdir)
assert fs, f"'{args.subdir}/' altında ckpt yok"
sel = fs[-args.last:]
w = ema_weights(len(sel), args.decay)
print(f"[ema] {len(fs)} ckpt bulundu → son {len(sel)} ortalanıyor (decay={args.decay}):", flush=True)
for f, wi in zip(sel, w):
print(f" {f} (w={wi:.3f})", flush=True)
states, cfg = [], None
for f in sel:
p = hf_hub_download(REPO, f, repo_type="model", token=token)
st = torch.load(p, map_location="cpu", weights_only=False)
cfg = st["cfg"]; states.append(st["model"])
avg = average_states(states, w)
os.makedirs(os.path.dirname(args.out_local), exist_ok=True)
steps = [f.split("step_")[1].split("/")[0] for f in sel]
torch.save({"model": avg, "cfg": cfg, "ema": True, "steps": steps, "decay": args.decay}, args.out_local)
print(f"[ema] kaydedildi → {args.out_local} (adımlar {steps[0]}..{steps[-1]})", flush=True)
if not args.no_push and token:
api.upload_file(path_or_fileobj=args.out_local, path_in_repo=f"{args.out_subdir}/ckpt.pt",
repo_id=REPO, repo_type="model",
commit_message=f"post-hoc EMA: son {len(sel)} {args.subdir} ckpt (decay={args.decay})")
print(f"[ema] HF push OK → {args.out_subdir}/ckpt.pt", flush=True)
print(f" faz5 kıyas: python faz5_eval.py --ckpt {args.out_subdir}/ckpt.pt --tasks xcopa,belebele,hellaswag,xnli,turkishmmlu --limit 300", flush=True)
if __name__ == "__main__":
main()