from __future__ import annotations import os, json, time import random from typing import Dict, List, Tuple, Callable import yaml import numpy as np import torch import torch.nn.functional as F from rich.console import Console from rich.table import Table from bihmoe.models.dense import DenseModel from bihmoe.models.structured import StructuredBiHMoE from bihmoe.utils.misc import set_seed, fmt_bytes from bihmoe.utils.metrics import accuracy from bihmoe.render.english import ( render_bind_query_q, render_bind_query_a, render_twohop_q, render_twohop_a, ) console = Console() def _fmt(x, nd=4): try: if x != x: return "NaN" return f"{float(x):.{nd}f}" except Exception: return str(x) def print_metrics_pretty(metrics: dict) -> None: step = metrics.get("step", "?") def pair(split: str) -> str: s_acc = float(metrics.get(f"{split}/acc_s", 0.0)) d_acc = float(metrics.get(f"{split}/acc_d", 0.0)) delta = s_acc - d_acc return f"{split} {s_acc:.3f}/{d_acc:.3f} (Δ {delta:+.3f})" # One-line summary (corner-of-eye) line = " ".join([ f"step={step}", pair("iid"), pair("ood"), pair("ood_long"), pair("braid"), pair("pert"), ]) console.rule(f"[bold]Step {step}[/bold]") print("KEYLINE " + line) console.print(line) table = Table(show_header=True, header_style="bold") table.add_column("split") table.add_column("acc_s", justify="right") table.add_column("acc_d", justify="right") table.add_column("loss_s", justify="right") table.add_column("loss_d", justify="right") table.add_column("disagree μ", justify="right") table.add_column("disagree(c)", justify="right") table.add_column("disagree(!c)", justify="right") for split in ["iid", "ood", "ood_long", "braid", "pert"]: table.add_row( split, _fmt(metrics.get(f"{split}/acc_s"), 4), _fmt(metrics.get(f"{split}/acc_d"), 4), _fmt(metrics.get(f"{split}/loss_s"), 4), _fmt(metrics.get(f"{split}/loss_d"), 4), _fmt(metrics.get(f"{split}/disagree_mean"), 4), _fmt(metrics.get(f"{split}/disagree_correct_mean"), 4), _fmt(metrics.get(f"{split}/disagree_incorrect_mean"), 4), ) console.print(table) console.print( f"train loss: S={_fmt(metrics.get('loss_s_train'),4)} D={_fmt(metrics.get('loss_d_train'),4)} " f"cuda_peak={metrics.get('cuda_peak_h','?')} reserved={metrics.get('cuda_reserved_h','?')}" ) def now_run_id(tag: str) -> str: return time.strftime("%Y%m%d-%H%M%S") + f"_{tag}" def amp_setup(device: torch.device, mode: str): mode = str(mode).lower() if device.type != "cuda": return (False, None, None) if mode == "bf16": try: if hasattr(torch.cuda, "is_bf16_supported") and not torch.cuda.is_bf16_supported(): console.print("WARN: bf16 not supported; falling back to fp16 AMP") mode = "fp16" except Exception: mode = "fp16" if mode == "fp16": return (True, torch.float16, torch.amp.GradScaler('cuda', enabled=True)) if mode == "bf16": return (True, torch.bfloat16, torch.cuda.amp.GradScaler(enabled=False)) return (False, None, None) def compute_match_dense(struct_cfg: Dict, round_multiple: int = 256) -> Tuple[int,int]: Ls = int(struct_cfg["stem_layers"]) Lh = int(struct_cfg["hemi_layers"]) dff_dense = int(struct_cfg["dff_dense"]) dff_expert = int(struct_cfg["dff_expert"]) topk = int(struct_cfg["topk"]) dense_layers = Ls + 2 * Lh moe_per_hemi = Lh // 2 dense_per_hemi = Lh - moe_per_hemi ffn_units = (Ls * dff_dense) + (2 * dense_per_hemi * dff_dense) + (2 * moe_per_hemi * topk * dff_expert) dense_dff = ffn_units / dense_layers dense_dff_round = max(round_multiple, int(round_multiple * round(dense_dff / round_multiple))) return dense_layers, dense_dff_round def collate(batch_ids: List[List[int]], pad_id: int, max_len: int) -> torch.Tensor: B = len(batch_ids) T = min(max(len(x) for x in batch_ids), max_len) out = torch.full((B, T), pad_id, dtype=torch.long) for i, ids in enumerate(batch_ids): ids = ids[:T] out[i, :len(ids)] = torch.tensor(ids, dtype=torch.long) return out @torch.no_grad() def eval_split( model_s, model_d, device, use_amp: bool, amp_dtype, split_name: str, records: List[Dict], cfg: Dict, encode_obj: Callable[[Dict,int,int,int], Tuple[List[int], int]], ) -> Dict: vocab = int(cfg["task"]["vocab_size"]) max_len = int(cfg["train"]["seq_max_len"]) noise_vocab = int(cfg["task"]["noise_vocab"]) eval_bs = int(cfg["train"].get("eval_batch_size", 64)) y_gold, y_pred_s, y_pred_d = [], [], [] sum_loss_s, sum_loss_d = 0.0, 0.0 disag, incorrect = [], [] autocast_ctx = torch.autocast(device_type="cuda", dtype=amp_dtype, enabled=(use_amp and device.type=="cuda")) n_total = 0 for i in range(0, len(records), eval_bs): chunk = records[i:i+eval_bs] batch_ids, batch_tgt = [], [] for r in chunk: ids, tgt = encode_obj(r, vocab, noise_vocab, max_len) batch_ids.append(ids) batch_tgt.append(tgt) inp = collate(batch_ids, pad_id=0, max_len=max_len).to(device) tgt = torch.tensor(batch_tgt, device=device, dtype=torch.long) with autocast_ctx: logits_s, aux = model_s(inp, return_aux=True, global_step=10**9) logits_d = model_d(inp) loss_s = F.cross_entropy(logits_s.float(), tgt, reduction="sum") loss_d = F.cross_entropy(logits_d.float(), tgt, reduction="sum") ps = torch.argmax(logits_s, dim=-1).detach().cpu().numpy().astype(np.int64) pd = torch.argmax(logits_d, dim=-1).detach().cpu().numpy().astype(np.int64) y_gold.extend(batch_tgt) y_pred_s.extend(ps.tolist()) y_pred_d.extend(pd.tolist()) dk = aux["sym_kl"].float().detach().cpu().numpy() disag.extend(dk.tolist()) incorrect.extend((ps != np.array(batch_tgt, dtype=np.int64)).astype(np.int64).tolist()) bs = len(batch_tgt) n_total += bs sum_loss_s += float(loss_s.item()) sum_loss_d += float(loss_d.item()) y_gold_np = np.array(y_gold, dtype=np.int64) acc_s = accuracy(np.array(y_pred_s), y_gold_np) acc_d = accuracy(np.array(y_pred_d), y_gold_np) disag_np = np.array(disag, dtype=np.float64) incorrect_np = np.array(incorrect, dtype=np.int64) return { f"{split_name}/loss_s": float(sum_loss_s / max(1, n_total)), f"{split_name}/loss_d": float(sum_loss_d / max(1, n_total)), f"{split_name}/acc_s": acc_s, f"{split_name}/acc_d": acc_d, f"{split_name}/disagree_mean": float(disag_np.mean()), f"{split_name}/disagree_correct_mean": float(disag_np[incorrect_np==0].mean()) if (incorrect_np==0).any() else float("nan"), f"{split_name}/disagree_incorrect_mean": float(disag_np[incorrect_np==1].mean()) if (incorrect_np==1).any() else float("nan"), } def build_eval_sets(cfg: Dict, seed: int, make_record_fn, record_to_json_fn) -> Dict[str, List[Dict]]: rng = random.Random(seed) tcfg = cfg["task"] scfg = cfg["splits"] # IDs vary by task; the make_record_fn will validate usage key_ids = list(range(int(tcfg["key_start"]), int(tcfg["key_start"]) + int(tcfg["key_count"]))) val_ids = list(range(int(tcfg["val_start"]), int(tcfg["val_start"]) + int(tcfg["val_count"]))) mid_ids = list(range(int(tcfg.get("mid_start", 1024)), int(tcfg.get("mid_start", 1024)) + int(tcfg.get("mid_count", 128)))) def mk(n, nmin, nmax, fmt="train", gap_max=0): out = [] for _ in range(n): npairs = rng.randint(nmin, nmax) rec = make_record_fn(rng, npairs, key_ids, mid_ids, val_ids, fmt=fmt, gap_max=gap_max) \ if tcfg["name"] == "twohop_bind" \ else make_record_fn(rng, npairs, key_ids, val_ids, fmt=fmt, gap_max=gap_max) out.append(record_to_json_fn(rec)) return out n_iid = int(cfg["train"].get("eval_size_iid", 256)) n_ood = int(cfg["train"].get("eval_size_ood", 256)) n_pert = int(cfg["train"].get("eval_size_pert", 256)) n_long = int(cfg["train"].get("eval_size_long", n_ood)) n_braid = int(cfg["train"].get("eval_size_braid", n_pert)) return { "iid": mk(n_iid, int(scfg["train_pairs_min"]), int(scfg["train_pairs_max"]), fmt="train", gap_max=0), "ood": mk(n_ood, int(scfg["ood_pairs_min"]), int(scfg["ood_pairs_max"]), fmt="train", gap_max=0), "ood_long": mk( n_long, int(scfg.get("long_pairs_min", scfg["ood_pairs_min"])), int(scfg.get("long_pairs_max", scfg["ood_pairs_max"])), fmt="perturb_gap", gap_max=int(scfg.get("long_gap_max", max(12, int(scfg.get("pert_gap_max", 6))*2))), ), "braid": mk( n_braid, int(scfg["train_pairs_min"]), int(scfg["train_pairs_max"]), fmt="braid", gap_max=int(scfg.get("braid_gap_max", 0)), ), "pert": mk( n_pert, int(scfg["train_pairs_min"]), int(scfg["train_pairs_max"]), fmt="perturb_gap", gap_max=int(scfg.get("pert_gap_max", 6)), ), } def dump_english_examples(out_dir: str, eval_sets: Dict[str, List[Dict]], n: int, render_q, render_a, solve_from_json): os.makedirs(out_dir, exist_ok=True) for split, recs in eval_sets.items(): path = os.path.join(out_dir, f"examples_{split}.txt") with open(path, "w", encoding="utf-8") as f: for r in recs[:n]: f.write(render_q(r)) f.write("\n") gold = solve_from_json(r) f.write(render_a(gold)) f.write("\n" + ("-"*60) + "\n") def mem_snapshot(): if not torch.cuda.is_available(): return {} return { "cuda_alloc": int(torch.cuda.memory_allocated()), "cuda_reserved": int(torch.cuda.memory_reserved()), "cuda_peak": int(torch.cuda.max_memory_allocated()), } def main(): import argparse ap = argparse.ArgumentParser() ap.add_argument("--config", required=True) args = ap.parse_args() cfg = yaml.safe_load(open(args.config, "r", encoding="utf-8")) seed = int(cfg["run"]["seed"]) set_seed(seed) device = torch.device("cuda" if torch.cuda.is_available() else "cpu") if device.type == "cuda": torch.backends.cuda.matmul.allow_tf32 = True use_amp, amp_dtype, scaler = amp_setup(device, cfg["model_common"]["dtype"]) vocab = int(cfg["task"]["vocab_size"]) run_id = now_run_id(str(cfg["run"]["tag"])) out_dir = os.path.join(str(cfg["run"]["out_root"]), run_id) os.makedirs(out_dir, exist_ok=True) with open(os.path.join(out_dir, "config.yaml"), "w", encoding="utf-8") as f: yaml.safe_dump(cfg, f, sort_keys=False) # --- Task selection --- task_name = str(cfg["task"]["name"]) if task_name == "bind_query": from bihmoe.tasks.bind_query import make_record as mk, record_to_json as r2j, encode as enc, BindQueryRecord, solve def encode_obj(r: Dict, vocab: int, noise_vocab: int, max_len: int): rec = BindQueryRecord(pairs=[tuple(p) for p in r["pairs"]], query_k=int(r["query_k"]), fmt=str(r.get("fmt","train")), gap_max=int(r.get("gap_max",0))) return enc(rec, vocab, noise_vocab=noise_vocab, max_len=max_len) def solve_from_json(r: Dict) -> int: rec = BindQueryRecord(pairs=[tuple(p) for p in r["pairs"]], query_k=int(r["query_k"]), fmt=str(r.get("fmt","train")), gap_max=int(r.get("gap_max",0))) return solve(rec) render_q, render_a = render_bind_query_q, render_bind_query_a elif task_name == "twohop_bind": from bihmoe.tasks.twohop_bind import make_record as mk2, record_to_json as r2j2, encode as enc2, TwoHopBindRecord, solve as solve2 def encode_obj(r: Dict, vocab: int, noise_vocab: int, max_len: int): rec = TwoHopBindRecord( f_pairs=[tuple(p) for p in r["f_pairs"]], g_pairs=[tuple(p) for p in r["g_pairs"]], query_k=int(r["query_k"]), fmt=str(r.get("fmt","train")), gap_max=int(r.get("gap_max",0)), ) return enc2(rec, vocab, noise_vocab=noise_vocab, max_len=max_len) def solve_from_json(r: Dict) -> int: rec = TwoHopBindRecord( f_pairs=[tuple(p) for p in r["f_pairs"]], g_pairs=[tuple(p) for p in r["g_pairs"]], query_k=int(r["query_k"]), fmt=str(r.get("fmt","train")), gap_max=int(r.get("gap_max",0)), ) return solve2(rec) mk, r2j, enc = mk2, r2j2, enc2 render_q, render_a = render_twohop_q, render_twohop_a else: raise SystemExit(f"Unknown task.name={task_name}") # Eval sets (fixed) eval_sets = build_eval_sets(cfg, seed=seed + 999, make_record_fn=mk, record_to_json_fn=r2j) with open(os.path.join(out_dir, "eval_sets.json"), "w", encoding="utf-8") as f: json.dump(eval_sets, f) dump_english_examples(out_dir, eval_sets, n=8, render_q=render_q, render_a=render_a, solve_from_json=solve_from_json) # Dense compute-match if bool(cfg["dense"].get("auto_from_struct", True)): dense_layers, dense_dff = compute_match_dense(cfg["structured"], int(cfg["dense"].get("round_multiple",256))) else: dense_layers = int(cfg["dense"]["layers"]) dense_dff = int(cfg["dense"]["dff"]) d_model = int(cfg["model_common"]["d_model"]) n_heads = int(cfg["model_common"]["n_heads"]) model_d = DenseModel( vocab_size=vocab, d_model=d_model, n_heads=n_heads, n_layers=dense_layers, d_ff=dense_dff, dropout=0.0, head_mode="cls", pool=str(cfg.get("dense", {}).get("pool", "last")), pad_id=0 ).to(device) model_s = StructuredBiHMoE( vocab_size=vocab, d_model=d_model, n_heads=n_heads, n_layers_stem=int(cfg["structured"]["stem_layers"]), n_layers_hemi=int(cfg["structured"]["hemi_layers"]), d_ff_dense=int(cfg["structured"]["dff_dense"]), d_ff_expert=int(cfg["structured"]["dff_expert"]), n_experts=int(cfg["structured"]["experts"]), top_k=int(cfg["structured"]["topk"]), workspace_tokens=int(cfg["structured"]["workspace"]), reconcile_every=int(cfg["structured"]["reconcile_every"]), dropout=0.0, pad_id=0, fuse="mean", moe_warmup_steps=int(cfg.get('structured', {}).get('moe_warmup_steps', 0)), left_local_window=int(cfg.get('structured', {}).get('left_local_window', 0)), right_noise_std=float(cfg.get('structured', {}).get('right_noise_std', 0.0)), callosum_competitive=bool(cfg.get('structured', {}).get('callosum_competitive', True)), callosum_tau=float(cfg.get('structured', {}).get('callosum_tau', 1.0)), chiasm_enabled=bool(cfg.get('structured', {}).get('chiasm_enabled', False)), noise_vocab=int(cfg.get('task', {}).get('noise_vocab', 16)), key_start=int(cfg.get('task', {}).get('key_start', -1)), key_count=int(cfg.get('task', {}).get('key_count', 0)), mid_start=int(cfg.get('task', {}).get('mid_start', -1)), mid_count=int(cfg.get('task', {}).get('mid_count', 0)), val_start=int(cfg.get('task', {}).get('val_start', -1)), val_count=int(cfg.get('task', {}).get('val_count', 0)), output_gate=bool(cfg.get('structured', {}).get('output_gate', True)), output_tau=float(cfg.get('structured', {}).get('output_tau', 1.0)), ).to(device) opt_d = torch.optim.AdamW(model_d.parameters(), lr=float(cfg["train"]["lr"])) opt_s = torch.optim.AdamW(model_s.parameters(), lr=float(cfg["train"]["lr"])) # ID pools tcfg = cfg["task"] key_ids = list(range(int(tcfg["key_start"]), int(tcfg["key_start"]) + int(tcfg["key_count"]))) val_ids = list(range(int(tcfg["val_start"]), int(tcfg["val_start"]) + int(tcfg["val_count"]))) mid_ids = list(range(int(tcfg.get("mid_start", 1024)), int(tcfg.get("mid_start", 1024)) + int(tcfg.get("mid_count", 128)))) scfg = cfg["splits"] train_min = int(scfg["train_pairs_min"]) mix_pert_prob = float(cfg['train'].get('mix_pert_prob', 0.0)) mix_long_prob = float(cfg['train'].get('mix_long_prob', 0.0)) train_max = int(scfg["train_pairs_max"]) batch = int(cfg["train"]["batch_size"]) max_len = int(cfg["train"]["seq_max_len"]) noise_vocab = int(cfg["task"]["noise_vocab"]) steps = int(cfg["train"]["steps"]) eval_every = int(cfg["train"]["eval_every"]) grad_clip = float(cfg["train"].get("grad_clip", 1.0)) lb_weight = float(cfg["train"].get("lb_weight", 0.0)) metrics_path = os.path.join(out_dir, "metrics.jsonl") console.print(f"run_id: {run_id}") console.print(f"out_dir: {out_dir}") console.print(f"task: {task_name}") console.print(f"device: {device} amp: {use_amp} amp_dtype: {amp_dtype}") console.print(f"dense_layers/dff: {dense_layers} {dense_dff}") autocast_ctx = torch.autocast(device_type="cuda", dtype=amp_dtype, enabled=(use_amp and device.type=="cuda")) def make_train_batch(step: int): rng = random.Random((seed * 1_000_000) + step) cur = cfg.get("train", {}).get("curriculum", None) # Base knobs (fallback if no curriculum) mix_long = float(cfg["train"].get("mix_long_prob", 0.0)) mix_pert = float(cfg["train"].get("mix_pert_prob", 0.0)) mix_braid = float(cfg["train"].get("mix_braid_prob", 0.0)) long_pairs_min = int(scfg.get("long_pairs_min", scfg["ood_pairs_min"])) long_pairs_max = int(scfg.get("long_pairs_max", scfg["ood_pairs_max"])) long_gap_max = int(scfg.get("long_gap_max", max(12, int(scfg.get("pert_gap_max", 6))*2))) pert_gap_max = int(scfg.get("pert_gap_max", 6)) braid_gap_max = int(scfg.get("braid_gap_max", 0)) def lerp(a, b, t): return a + (b - a) * t if isinstance(cur, dict): p1 = int(cur.get("phase1_steps", 2000)) ramp = int(cur.get("ramp_steps", 0)) # Phase 1 mix_long1 = float(cur.get("mix_long1", mix_long)) mix_pert1 = float(cur.get("mix_pert1", mix_pert)) mix_braid1 = float(cur.get("mix_braid1", mix_braid)) # Phase 2 mix_long2 = float(cur.get("mix_long2", mix_long)) mix_pert2 = float(cur.get("mix_pert2", mix_pert)) mix_braid2 = float(cur.get("mix_braid2", mix_braid)) long_pairs_min2 = int(cur.get("long_pairs_min2", long_pairs_min)) long_pairs_max2 = int(cur.get("long_pairs_max2", long_pairs_max)) long_gap_max2 = int(cur.get("long_gap_max2", long_gap_max)) pert_gap_max2 = int(cur.get("pert_gap_max2", pert_gap_max)) braid_gap_max2 = int(cur.get("braid_gap_max2", braid_gap_max)) if ramp <= 0: if step <= p1: mix_long, mix_pert, mix_braid = mix_long1, mix_pert1, mix_braid1 else: mix_long, mix_pert, mix_braid = mix_long2, mix_pert2, mix_braid2 long_pairs_min, long_pairs_max = long_pairs_min2, long_pairs_max2 long_gap_max, pert_gap_max, braid_gap_max = long_gap_max2, pert_gap_max2, braid_gap_max2 else: if step <= p1: t = 0.0 elif step >= p1 + ramp: t = 1.0 else: t = (step - p1) / float(ramp) mix_long = float(lerp(mix_long1, mix_long2, t)) mix_pert = float(lerp(mix_pert1, mix_pert2, t)) mix_braid = float(lerp(mix_braid1, mix_braid2, t)) long_pairs_min = int(round(lerp(long_pairs_min, long_pairs_min2, t))) long_pairs_max = int(round(lerp(long_pairs_max, long_pairs_max2, t))) long_gap_max = int(round(lerp(long_gap_max, long_gap_max2, t))) pert_gap_max = int(round(lerp(pert_gap_max, pert_gap_max2, t))) braid_gap_max = int(round(lerp(braid_gap_max, braid_gap_max2, t))) # Normalize so total hard prob <= 0.95 mix_long = max(0.0, mix_long) mix_pert = max(0.0, mix_pert) mix_braid = max(0.0, mix_braid) total = mix_long + mix_pert + mix_braid if total > 0.95: scale = 0.95 / total mix_long *= scale mix_pert *= scale mix_braid *= scale recs, tgts = [], [] for _ in range(batch): r = rng.random() fmt = "train" gap = 0 if r < mix_long: npairs = rng.randint(long_pairs_min, long_pairs_max) fmt = "perturb_gap" gap = long_gap_max elif r < (mix_long + mix_pert): npairs = rng.randint(train_min, train_max) fmt = "perturb_gap" gap = pert_gap_max elif r < (mix_long + mix_pert + mix_braid): npairs = rng.randint(train_min, train_max) fmt = "braid" gap = braid_gap_max else: npairs = rng.randint(train_min, train_max) fmt = "train" gap = 0 if task_name == "twohop_bind": rec = mk(rng, npairs, key_ids, mid_ids, val_ids, fmt=fmt, gap_max=gap) else: rec = mk(rng, npairs, key_ids, val_ids, fmt=fmt, gap_max=gap) ids, tgt = enc(rec, vocab, noise_vocab=noise_vocab, max_len=max_len) recs.append(ids) tgts.append(tgt) return recs, tgts for step in range(1, steps + 1): recs, tgts = make_train_batch(step) inp = collate(recs, pad_id=0, max_len=max_len).to(device) tgt = torch.tensor(tgts, device=device, dtype=torch.long) # Structured step opt_s.zero_grad(set_to_none=True) with autocast_ctx: logits_s, aux_s = model_s(inp, return_aux=True, global_step=step) loss_s_main = F.cross_entropy(logits_s.float(), tgt) loss_s = loss_s_main + (lb_weight * aux_s.get('lb_loss', 0.0)) if not torch.isfinite(loss_s): console.print(f"NON-FINITE loss_s at step {step} = {float(loss_s)}") break if scaler is not None and scaler.is_enabled(): scaler.scale(loss_s).backward() scaler.unscale_(opt_s) torch.nn.utils.clip_grad_norm_(model_s.parameters(), grad_clip) scaler.step(opt_s) else: loss_s.backward() torch.nn.utils.clip_grad_norm_(model_s.parameters(), grad_clip) opt_s.step() # Dense step opt_d.zero_grad(set_to_none=True) with autocast_ctx: logits_d = model_d(inp) loss_d = F.cross_entropy(logits_d.float(), tgt) if not torch.isfinite(loss_d): console.print(f"NON-FINITE loss_d at step {step} = {float(loss_d)}") break if scaler is not None and scaler.is_enabled(): scaler.scale(loss_d).backward() scaler.unscale_(opt_d) torch.nn.utils.clip_grad_norm_(model_d.parameters(), grad_clip) scaler.step(opt_d) scaler.update() else: loss_d.backward() torch.nn.utils.clip_grad_norm_(model_d.parameters(), grad_clip) opt_d.step() if step % eval_every == 0 or step == 1: model_s.eval() model_d.eval() if torch.cuda.is_available(): torch.cuda.reset_peak_memory_stats() lbv = aux_s.get("lb_loss", 0.0) if hasattr(lbv, "detach"): lbv = float(lbv.detach().float().item()) else: lbv = float(lbv) metrics = {"step": step, "loss_s_train": float(loss_s.item()), "loss_d_train": float(loss_d.item()), "lb_loss": lbv, "lb_weight": float(lb_weight)} for split in ["iid","ood","ood_long","braid","pert"]: metrics.update(eval_split(model_s, model_d, device, use_amp, amp_dtype, split, eval_sets[split], cfg, encode_obj)) if torch.cuda.is_available(): metrics.update({ "cuda_alloc": int(torch.cuda.memory_allocated()), "cuda_reserved": int(torch.cuda.memory_reserved()), "cuda_peak": int(torch.cuda.max_memory_allocated()), }) metrics["cuda_alloc_h"] = fmt_bytes(metrics.get("cuda_alloc",0)) metrics["cuda_reserved_h"] = fmt_bytes(metrics.get("cuda_reserved",0)) metrics["cuda_peak_h"] = fmt_bytes(metrics.get("cuda_peak",0)) with open(metrics_path, "a", encoding="utf-8") as f: f.write(json.dumps(metrics) + "\n") print_metrics_pretty(metrics) model_s.train() model_d.train() console.print(f"done: {out_dir}") if __name__ == "__main__": main()