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