File size: 5,792 Bytes
1ea7ba6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Aggregate test metrics across multi-seed runs and run McNemar significance.



Compares strong-aug (3 seeds) against the raw baseline on the same test split,

reports per-seed Top-1 + mean ± std, and McNemar's chi-square + p-value for

strong-aug vs raw baseline on a per-sample basis.

"""
import sys, os
sys.path.insert(0, "/mnt/d/SpiceNet" if os.path.exists("/mnt/d/SpiceNet") else "D:/SpiceNet")

import json
import numpy as np
import torch
from scipy.stats import chi2_contingency

import config
from src.dataset import get_dataloaders
from src.model import load_checkpoint


@torch.no_grad()
def collect_predictions(ckpt_path, device):
    model, *_ = load_checkpoint(ckpt_path, device)
    model.eval()
    _, _, test_loader, _, _ = get_dataloaders(multimodal=True)
    y_true, y_pred = [], []
    for imgs, tex, col, lbl in test_loader:
        imgs = imgs.to(device); tex = tex.to(device); col = col.to(device)
        logits, _ = model.forward_fusion(imgs, tex, col)
        y_true.extend(lbl.tolist())
        y_pred.extend(logits.argmax(1).cpu().tolist())
    return np.array(y_true), np.array(y_pred)


def mcnemar_test(y_true, pred_A, pred_B):
    """McNemar's exact test (continuity-corrected chi^2).



    b = A correct, B wrong; c = A wrong, B correct.

    chi2 = (|b - c| - 1)^2 / (b + c)

    """
    correct_A = pred_A == y_true
    correct_B = pred_B == y_true
    b = int(((correct_A) & (~correct_B)).sum())
    c = int(((~correct_A) & (correct_B)).sum())

    if b + c == 0:
        return {"b": 0, "c": 0, "chi2": 0.0, "p_value": 1.0,
                "verdict": "no disagreement"}

    chi2 = (abs(b - c) - 1) ** 2 / (b + c)
    # Exact via scipy: McNemar's is a 2x2 chi2 test on the discordant cells.
    # We'll report both the continuity-corrected chi2 and the exact binomial p.
    from scipy.stats import binom
    # Two-sided exact binomial p-value:
    n = b + c
    k = min(b, c)
    p_exact = 2 * binom.cdf(k, n, 0.5)
    if p_exact > 1:
        p_exact = 1.0
    return {
        "b_A_correct_B_wrong": b,
        "c_A_wrong_B_correct": c,
        "chi2_continuity_corrected": chi2,
        "p_value_exact_binomial": p_exact,
        "verdict": "significant" if p_exact < 0.05 else "not significant",
    }


def main():
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    ck = config.CHECKPOINT_DIR

    runs = {
        "raw_baseline":     str(ck / "best_original.pth"),
        "sam":              str(ck / "best_sam.pth"),
        "strongaug_seed42": str(ck / "best_strongaug.pth"),
        "strongaug_seed1337": str(ck / "seed1337" / "best.pth"),
        "strongaug_seed2024": str(ck / "seed2024" / "best.pth"),
    }

    preds = {}
    for name, path in runs.items():
        if not os.path.exists(path):
            print(f"  SKIP {name}: {path} not found")
            continue
        print(f"Evaluating {name}...")
        y_true, y_pred = collect_predictions(path, device)
        preds[name] = (y_true, y_pred)

    # Top-1 per run
    print("\n=== Per-run Top-1 ===")
    per_run = {}
    for name, (yt, yp) in preds.items():
        acc = float((yt == yp).mean())
        per_run[name] = acc
        print(f"  {name:>22s}: {acc:.4f}")

    # Mean ± std across strongaug seeds
    sa_keys = [k for k in per_run if k.startswith("strongaug_seed")]
    if sa_keys:
        accs = np.array([per_run[k] for k in sa_keys])
        print(f"\n=== Strong-aug seeds aggregate ({len(sa_keys)} seeds) ===")
        print(f"  Mean   : {accs.mean():.4f}")
        print(f"  Std    : {accs.std(ddof=1):.4f}")
        print(f"  Min/Max: {accs.min():.4f} / {accs.max():.4f}")

    # McNemar's test: strongaug_seed42 vs raw_baseline
    if "strongaug_seed42" in preds and "raw_baseline" in preds:
        yt_a, yp_a = preds["raw_baseline"]
        yt_b, yp_b = preds["strongaug_seed42"]
        assert np.array_equal(yt_a, yt_b), "Test sets must match across runs"
        m = mcnemar_test(yt_a, yp_a, yp_b)
        print("\n=== McNemar's test (raw baseline vs strong-aug seed 42) ===")
        for k, v in m.items():
            print(f"  {k}: {v}")

    # Pairwise across all strongaug seeds (sanity — should all be NS)
    print("\n=== Strong-aug seed agreement (pairwise McNemar) ===")
    if len(sa_keys) >= 2:
        for i, a in enumerate(sa_keys):
            for b in sa_keys[i+1:]:
                yt_a, yp_a = preds[a]
                yt_b, yp_b = preds[b]
                m = mcnemar_test(yt_a, yp_a, yp_b)
                if "verdict" in m and m["verdict"] == "no disagreement":
                    print(f"  {a} vs {b}: identical predictions")
                else:
                    print(f"  {a} vs {b}: b={m['b_A_correct_B_wrong']}, "
                          f"c={m['c_A_wrong_B_correct']}, p={m['p_value_exact_binomial']:.4f}")

    # Save full results
    out = {
        "per_run_top1": per_run,
        "strongaug_seeds_summary": {
            "n": len(sa_keys),
            "mean": float(np.mean([per_run[k] for k in sa_keys])) if sa_keys else None,
            "std": float(np.std([per_run[k] for k in sa_keys], ddof=1)) if len(sa_keys) > 1 else None,
            "min": float(np.min([per_run[k] for k in sa_keys])) if sa_keys else None,
            "max": float(np.max([per_run[k] for k in sa_keys])) if sa_keys else None,
        },
        "mcnemar_raw_vs_strongaug": m if ("strongaug_seed42" in preds and "raw_baseline" in preds) else None,
    }
    out_path = config.OUTPUT_DIR / "multiseed_aggregate.json"
    with open(out_path, "w") as f:
        json.dump(out, f, indent=2)
    print(f"\nSaved -> {out_path}")


if __name__ == "__main__":
    main()