FST_code / scripts /eval /python_file.py
jasonfan's picture
2026-03-19
3b2d368 verified
Raw
History Blame Contribute Delete
4.79 kB
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)}")