Repository navigation
Expand file tree
/
Copy pathoptimize_threshold.py
More file actions
120 lines (97 loc) · 4.44 KB
/
Copy pathoptimize_threshold.py
File metadata and controls
120 lines (97 loc) · 4.44 KB
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
"""
Phase 11 — Threshold Optimisation
===================================
Sweep classification thresholds on the validation set to find the
threshold that maximises Micro-F1.
"""
import numpy as np
import torch
from torch.utils.data import DataLoader
from tqdm import tqdm
from sklearn.metrics import f1_score
import config
import utils
from dataset import MeSHDataset
from train import MeSHClassifier
def run(model_dir: str = None):
"""
Find the optimal classification threshold on the validation set.
Steps:
1. Load the best model checkpoint.
2. Run inference on the validation set.
3. Sweep thresholds [0.10, 0.15, ..., 0.75].
4. Select the threshold with the highest Micro-F1.
5. Save the result.
"""
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# ── Resolve model directory ──────────────────────────────────────────
if model_dir is None:
dirs = sorted(config.MODELS_DIR.iterdir())
if not dirs:
print("ERROR: No model directories found. Run train.py first.")
return
model_dir = dirs[-1]
else:
from pathlib import Path
model_dir = Path(model_dir)
model_cfg = utils.load_json(model_dir / "model_config.json")
num_labels = model_cfg["num_labels"]
top_k = model_cfg["top_k"]
label2id = utils.load_label2id(top_k=top_k)
splits = utils.load_splits()
from transformers import AutoTokenizer
tokenizer = AutoTokenizer.from_pretrained(model_dir / "tokenizer")
print(f"Loading validation dataset...")
val_ds = MeSHDataset(splits["validation"], label2id, tokenizer, weighted=False)
val_loader = DataLoader(val_ds, batch_size=config.EVAL_BATCH_SIZE, shuffle=False)
# ── Load model ───────────────────────────────────────────────────────
model = MeSHClassifier(config.MODEL_NAME, num_labels)
model.load_state_dict(torch.load(model_dir / "best_model.pt", map_location=device))
model.to(device)
model.eval()
# ── Get all predictions ──────────────────────────────────────────────
all_probs = []
all_labels = []
with torch.no_grad():
for batch in tqdm(val_loader, desc="Inference"):
input_ids = batch["input_ids"].to(device)
attention_mask = batch["attention_mask"].to(device)
labels = batch["labels"]
outputs = model(input_ids, attention_mask)
probs = torch.sigmoid(outputs["logits"]).cpu().numpy()
all_probs.append(probs)
all_labels.append(labels.numpy())
all_probs = np.concatenate(all_probs)
all_labels = np.concatenate(all_labels)
all_labels_binary = (all_labels > 0).astype(int)
# ── Sweep thresholds ─────────────────────────────────────────────────
thresholds = np.arange(0.10, 0.76, 0.05)
results = []
print(f"\n{'Threshold':>10} {'Micro-F1':>10} {'Macro-F1':>10}")
print("-" * 35)
best_threshold = 0.5
best_f1 = 0.0
for th in thresholds:
preds = (all_probs > th).astype(int)
micro = f1_score(all_labels_binary, preds, average="micro", zero_division=0)
macro = f1_score(all_labels_binary, preds, average="macro", zero_division=0)
results.append({"threshold": round(float(th), 2), "micro_f1": round(micro, 4), "macro_f1": round(macro, 4)})
print(f"{th:>10.2f} {micro:>10.4f} {macro:>10.4f}")
if micro > best_f1:
best_f1 = micro
best_threshold = round(float(th), 2)
report = {
"best_threshold": best_threshold,
"best_micro_f1": round(best_f1, 4),
"sweep_results": results,
}
utils.save_json(report, config.THRESHOLD_REPORT_PATH)
print(f"\n[OK] Best threshold: {best_threshold} (Micro-F1 = {best_f1:.4f})")
print(f" Report saved to {config.THRESHOLD_REPORT_PATH}")
return report
if __name__ == "__main__":
import argparse
parser = argparse.ArgumentParser(description="Optimise classification threshold")
parser.add_argument("--model-dir", type=str, default=None)
args = parser.parse_args()
run(model_dir=args.model_dir)