File size: 4,792 Bytes
3b2d368 | 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 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 | 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)}")
|