import os import numpy as np # ============================================================ # 配置 # ============================================================ CHECKPOINT_DIRS = { "fst_353M": "/work/jf381/checkpoints/fst_353M_resume_new", "tf_353M": "/work/jf381/checkpoints/transformer_353M_resume", "tf_1.3B": "/work/jf381/checkpoints/transformer_1_3B_resume", "fst_1.3B": "/work/jf381/checkpoints/fst_1_3B_resume", } N_FOLDS = 5 N_SPLITS = 10 SEED = 2026 rng = np.random.default_rng(SEED) # ============================================================ # 解析 gsm8k_results_fold_n.txt # ============================================================ def parse_gsm8k_file(path): """ 返回: List[int], 1=correct, 0=wrong """ results = [] with open(path, "r", encoding="utf-8") as f: for line in f: if "Truth:" in line: if "✓" in line: results.append(1) elif "✗" in line: results.append(0) return results # ============================================================ # 读取 checkpoint 的完整 GSM8K 数据 # ============================================================ def load_full_dataset(ckpt_dir): all_results = [] for fold_id in range(1, N_FOLDS + 1): fname = f"gsm8k_results_fold_{fold_id}.txt" fpath = os.path.join(ckpt_dir, fname) if not os.path.exists(fpath): raise FileNotFoundError(f"Missing file: {fpath}") all_results.extend(parse_gsm8k_file(fpath)) return np.array(all_results) # ============================================================ # Jackknife fold std(你定义的 error bar) # ============================================================ def jackknife_fold_std(fold_accs): """ fold_accs: np.array shape (5,) 返回: std : jackknife std means: 5 个 leave-one-fold-out mean """ n = len(fold_accs) loo_means = [] for i in range(n): idx = [j for j in range(n) if j != i] loo_means.append(fold_accs[idx].mean()) loo_means = np.array(loo_means) return loo_means.std(), loo_means # ============================================================ # 主流程 # ============================================================ print("\n======================") print("Loading GSM8K datasets") print("======================") model_data = {} dataset_size = None for name, path in CHECKPOINT_DIRS.items(): data = load_full_dataset(path) model_data[name] = data if dataset_size is None: dataset_size = len(data) else: assert len(data) == dataset_size, "Dataset size mismatch!" print(f"{name:10s}: {len(data)} samples, ACC={data.mean():.4f}") # ============================================================ # 生成共享的 10 个随机 split # ============================================================ indices = np.arange(dataset_size) shared_splits = [] for _ in range(N_SPLITS): perm = rng.permutation(indices) folds = np.array_split(perm, N_FOLDS) shared_splits.append(folds) # ============================================================ # 计算 mean acc + jackknife error bar # ============================================================ print("\n======================") print("Per-split stats (jackknife over folds)") print("======================") results = {name: [] for name in CHECKPOINT_DIRS} for split_id, folds in enumerate(shared_splits): print(f"\nSplit {split_id + 1}") for name, data in model_data.items(): # 5 个 fold acc fold_accs = np.array([data[f].mean() for f in folds]) mean_acc = fold_accs.mean() err_std, loo_means = jackknife_fold_std(fold_accs) results[name].append({ "mean": mean_acc, "err": err_std, "fold_accs": fold_accs, "loo_means": loo_means }) print( f" {name:10s} " f"mean={mean_acc:.4f} " f"err(jackknife)={err_std:.4f} " f"loo_means={np.round(loo_means, 4)}" ) # ============================================================ # 汇总(用于画图 / 表格) # ============================================================ print("\n======================") print("Summary (average over splits)") print("======================") for name, vals in results.items(): means = np.array([v["mean"] for v in vals]) errs = np.array([v["err"] for v in vals]) print(f"\n{name}") print(f" Mean ACC : {means.mean():.4f}") print(f" Mean jackknife err : {errs.mean():.4f}") print(f" ACC per split : {np.round(means, 4)}") print(f" Error bar per split : {np.round(errs, 4)}")