peter2000 commited on
Commit
fe8e23c
·
verified ·
1 Parent(s): eca9e1b

Add full-metric evaluation script (accuracy + F1 for laya zero-shot, laya fine-tuned, setfit)

Browse files
Files changed (1) hide show
  1. eval_full.py +163 -0
eval_full.py ADDED
@@ -0,0 +1,163 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ os.environ.setdefault("USE_TF", "0")
3
+ os.environ.setdefault("USE_TORCH", "1")
4
+ os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
5
+ os.environ.setdefault("HF_HUB_DISABLE_XET", "1")
6
+
7
+ import json
8
+ import time
9
+
10
+ import numpy as np
11
+ import pandas as pd
12
+ import torch
13
+
14
+ from huggingface_hub import HfApi, snapshot_download
15
+ from laya.agent import _fix_tokenizer_config
16
+ from sklearn.metrics import f1_score, hamming_loss, precision_score, recall_score
17
+ from sklearn.model_selection import train_test_split
18
+ from transformers import AutoTokenizer
19
+
20
+ import laya
21
+
22
+ BASE_MODEL_ID = "convaiinnovations/laya"
23
+ FT_REPO = "peter2000/laya-vulnerability-groups"
24
+ SETFIT_REPO = "peter2000/setfit-vulnerability-groups"
25
+ PARQUET_URL = "https://huggingface.co/datasets/GIZ/vulnerability_training_data_full/resolve/refs%2Fconvert%2Fparquet/default/train/0000.parquet"
26
+ LABELS = [
27
+ "Agricultural communities", "Coastal communities", "Ethnic, racial or other minorities",
28
+ "Fishery communities", "Informal sector workers", "Members of indigenous and local communities",
29
+ "Migrants and displaced persons", "Older persons", "Other", "Persons living in poverty",
30
+ "Persons with disabilities", "Persons with pre-existing health conditions",
31
+ "Residents of drought-prone regions", "Rural populations", "Sexual minorities (LGBTQI+)",
32
+ "Urban populations", "Women and other genders",
33
+ ]
34
+ QIDS = [f"g{i}" for i in range(len(LABELS))]
35
+
36
+ QUESTIONS = {
37
+ qid: {
38
+ "type": "noul",
39
+ "instructions": f"Does this text indicate that {label} are targeted, supported, or affected as a vulnerable group? Answer true or false.",
40
+ }
41
+ for qid, label in zip(QIDS, LABELS)
42
+ }
43
+
44
+ def load_data():
45
+ df = pd.read_parquet(PARQUET_URL)
46
+ assert len(df) == 475, f"expected 475 rows, got {len(df)}"
47
+ Y = df[LABELS].values.astype(np.int64)
48
+ nlab = Y.sum(1)
49
+ idx_tr, idx_te = train_test_split(
50
+ np.arange(len(df)), test_size=0.2, random_state=42, stratify=np.minimum(nlab, 3)
51
+ )
52
+ return df, Y, np.asarray(idx_tr), np.asarray(idx_te)
53
+
54
+ def ece(conf, correct, n_bins=15):
55
+ conf = np.asarray(conf, dtype=np.float64)
56
+ corr = np.asarray(correct, dtype=np.float64)
57
+ bins = np.linspace(0.0, 1.0, n_bins + 1)
58
+ e = 0.0
59
+ for lo, hi in zip(bins[:-1], bins[1:]):
60
+ m = (conf > lo) & (conf <= hi)
61
+ if m.sum() > 0:
62
+ e += m.mean() * abs(corr[m].mean() - conf[m].mean())
63
+ return float(e)
64
+
65
+ def evaluate_full(Y_true, P_pred, threshold=0.5):
66
+ pred = (P_pred >= threshold).astype(int)
67
+ pl_f1 = f1_score(Y_true, pred, average=None, zero_division=0)
68
+ pl_prec = precision_score(Y_true, pred, average=None, zero_division=0)
69
+ pl_rec = recall_score(Y_true, pred, average=None, zero_division=0)
70
+ per_label_acc = (pred == Y_true).mean(axis=0)
71
+ conf = np.where(pred == 1, P_pred, 1.0 - P_pred)
72
+ corr = (pred == Y_true).astype(np.float64)
73
+ return {
74
+ "threshold": threshold,
75
+ "exact_match_accuracy": float(((pred == Y_true).all(axis=1)).mean()),
76
+ "hamming_accuracy": float(1.0 - hamming_loss(Y_true, pred)),
77
+ "hamming_loss": float(hamming_loss(Y_true, pred)),
78
+ "macro_f1": float(f1_score(Y_true, pred, average="macro", zero_division=0)),
79
+ "micro_f1": float(f1_score(Y_true, pred, average="micro", zero_division=0)),
80
+ "weighted_f1": float(f1_score(Y_true, pred, average="weighted", zero_division=0)),
81
+ "macro_precision": float(precision_score(Y_true, pred, average="macro", zero_division=0)),
82
+ "micro_precision": float(precision_score(Y_true, pred, average="micro", zero_division=0)),
83
+ "macro_recall": float(recall_score(Y_true, pred, average="macro", zero_division=0)),
84
+ "micro_recall": float(recall_score(Y_true, pred, average="micro", zero_division=0)),
85
+ "ece": ece(conf, corr),
86
+ "per_label_f1": {LABELS[i]: round(float(pl_f1[i]), 4) for i in range(len(LABELS))},
87
+ "per_label_precision": {LABELS[i]: round(float(pl_prec[i]), 4) for i in range(len(LABELS))},
88
+ "per_label_recall": {LABELS[i]: round(float(pl_rec[i]), 4) for i in range(len(LABELS))},
89
+ "per_label_accuracy": {LABELS[i]: round(float(per_label_acc[i]), 4) for i in range(len(LABELS))},
90
+ "test_positives": {LABELS[i]: int(Y_true[:, i].sum()) for i in range(len(LABELS))},
91
+ }
92
+
93
+ def probs_from_answers(res):
94
+ return np.array([res["answers"][qid]["noul"] for qid in QIDS], dtype=np.float64)
95
+
96
+ def eval_agent(agent, texts, Y_true):
97
+ t0 = time.time()
98
+ P = np.stack([probs_from_answers(agent.predict(t, QUESTIONS)) for t in texts])
99
+ m = evaluate_full(Y_true, P)
100
+ m["eval_seconds"] = round(time.time() - t0, 1)
101
+ m["probabilities"] = P.round(4).tolist()
102
+ return m
103
+
104
+ def main():
105
+ df, Y, idx_tr, idx_te = load_data()
106
+ texts = df["text"].tolist()
107
+ X_te = [texts[i] for i in idx_te]
108
+ Y_te = Y[idx_te]
109
+ print(f"test rows: {len(X_te)}, labels: {len(LABELS)}", flush=True)
110
+ device = "cuda"
111
+ results = {}
112
+
113
+ print("== laya fine-tuned ==", flush=True)
114
+ ft_dir = snapshot_download(FT_REPO, ignore_patterns=["*.py"])
115
+ agent = laya.load(ft_dir, device=device)
116
+ m = eval_agent(agent, X_te, Y_te)
117
+ results["laya_fine_tuned"] = m
118
+ print(json.dumps({k: v for k, v in m.items() if k != "probabilities"}, indent=2), flush=True)
119
+ del agent
120
+ torch.cuda.empty_cache()
121
+
122
+ print("== laya base zero-shot ==", flush=True)
123
+ base_dir = snapshot_download(BASE_MODEL_ID, ignore_patterns=["multilingual/*", "typed-decisions/*", "assets/*", "eval/*", "*.py"])
124
+ _fix_tokenizer_config(base_dir)
125
+ agent = laya.load(base_dir, device=device)
126
+ m = eval_agent(agent, X_te, Y_te)
127
+ results["laya_base_zero_shot"] = m
128
+ print(json.dumps({k: v for k, v in m.items() if k != "probabilities"}, indent=2), flush=True)
129
+ del agent
130
+ torch.cuda.empty_cache()
131
+
132
+ print("== setfit ==", flush=True)
133
+ from setfit import SetFitModel
134
+ sf = SetFitModel.from_pretrained(SETFIT_REPO)
135
+ t0 = time.time()
136
+ P = np.asarray(sf.predict_proba(X_te))
137
+ m = evaluate_full(Y_te, P)
138
+ m["eval_seconds"] = round(time.time() - t0, 1)
139
+ m["probabilities"] = P.round(4).tolist()
140
+ results["setfit"] = m
141
+ print(json.dumps({k: v for k, v in m.items() if k != "probabilities"}, indent=2), flush=True)
142
+
143
+ out = {
144
+ "dataset": "GIZ/vulnerability_training_data_full",
145
+ "split": "train_test_split(random_state=42, test_size=0.2, stratify=min(n_labels,3)); n_test=95",
146
+ "models": results,
147
+ }
148
+ api = HfApi(token=os.environ.get("HF_TOKEN"))
149
+ for repo in (FT_REPO, SETFIT_REPO):
150
+ api.upload_file(
151
+ path_or_fileobj=json.dumps(out, indent=2).encode(),
152
+ path_in_repo="metrics_full.json",
153
+ repo_id=repo,
154
+ repo_type="model",
155
+ commit_message="Full accuracy+F1 metrics: laya zero-shot, laya fine-tuned, setfit (95 test rows)",
156
+ )
157
+ print("uploaded metrics_full.json to", FT_REPO, "and", SETFIT_REPO)
158
+ print("DONE")
159
+
160
+ if __name__ == "__main__":
161
+ t0 = time.time()
162
+ main()
163
+ print(f"elapsed {time.time()-t0:.0f}s")