File size: 4,403 Bytes
89bb297
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
"""
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()