| 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) |
|
|
| |
| |
| |
| 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 |
|
|
|
|
| |
| |
| |
| 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) |
|
|
|
|
| |
| |
| |
| 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}") |
|
|
| |
| |
| |
| 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) |
|
|
| |
| |
| |
| 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(): |
| |
| 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)}") |
|
|