FST_code / src /lmr /glue_benchmark /glue_benchmark copy 5.py
jasonfan's picture
2026-03-19
3b2d368 verified
Raw
History Blame Contribute Delete
96.5 kB
# lmr/glue_benchmark_grid_full.py
"""
Grid-search GLUE + extra tasks runner (full single-file).
Features:
- LR sweep across bert_lr_candidates
- Random restarts for small/unstable tasks
- Save per-run checkpoints and run_meta.json
- Save all_runs.csv and best_overall per task
- Evaluate best model and compute errorbars for GLUE
- Supports EXTRA_TASKS (boolq/piqa/winogrande + hellaswag/openbookqa/arc) with MC/pair handling
"""
import os
import json
import re
import math
import random
import shutil
import time
from pathlib import Path
from typing import Optional, List, Tuple, Dict, Any
import torch
import numpy as np
import pandas as pd
from datasets import load_dataset
from torch.utils.data import DataLoader, TensorDataset
from tqdm import tqdm
import evaluate
# Project imports: adjust if your package layout differs
try:
from lmr.checkpointing import Checkpointing
from lmr.ddp import unwrap_model
except Exception:
# If these modules are not available, provide lightweight fallbacks to avoid import errors
Checkpointing = None
def unwrap_model(m):
return m
# ---------------------------------------------------------------------
# Tasks config
# ---------------------------------------------------------------------
GLUE_TASKS = {
"cola": {"type": "classification", "num_labels": 2, "hf_name": "cola"},
"sst2": {"type": "classification", "num_labels": 2, "hf_name": "sst2"},
"mrpc": {"type": "classification", "num_labels": 2, "hf_name": "mrpc"},
"stsb": {"type": "regression", "num_labels": 1, "hf_name": "stsb"},
"qqp": {"type": "classification", "num_labels": 2, "hf_name": "qqp"},
"mnli": {"type": "classification", "num_labels": 3, "hf_name": "mnli"},
"qnli": {"type": "classification", "num_labels": 2, "hf_name": "qnli"},
"rte": {"type": "classification", "num_labels": 2, "hf_name": "rte"},
"wnli": {"type": "classification", "num_labels": 2, "hf_name": "wnli"},
}
# Extra tasks (BoolQ, PIQA, Winogrande, HellaSwag, OpenBookQA, ARC variants)
EXTRA_TASKS = {
"boolq": {
"type": "classification",
"num_labels": 2,
"hf_path": "boolq",
"format": "pair",
},
"piqa": {
"type": "multiple_choice",
"num_labels": 2,
"hf_path": "piqa",
"format": "mc",
},
"winogrande": {
"type": "multiple_choice",
"num_labels": 2,
"hf_path": "winogrande",
"hf_config": "winogrande_xl",
"format": "mc",
},
# Added tasks below
"hellaswag": {
"type": "multiple_choice",
"num_labels": 4,
"hf_path": "hellaswag",
"format": "mc",
},
"openbookqa": {
"type": "multiple_choice",
"num_labels": 4,
"hf_path": "openbookqa",
"format": "mc",
},
# AI2 ARC splits — align names with common usage
"arc_easy": {
"type": "multiple_choice",
"num_labels": 4,
"hf_path": "ai2_arc",
"hf_config": "ARC-Easy",
"format": "mc",
},
"arc_challenge": {
"type": "multiple_choice",
"num_labels": 4,
"hf_path": "ai2_arc",
"hf_config": "ARC-Challenge",
"format": "mc",
},
}
ALL_TASKS = {**GLUE_TASKS, **EXTRA_TASKS}
# Small/unstable tasks for extra random restarts
SMALL_TASKS_RANDOM_RESTARTS = {"cola", "mrpc", "rte", "stsb"}
SMALL_TASKS_RANDOM_RESTARTS_EXTRA = set({"piqa", "boolq", "winogrande", "hellaswag"}) # adjust as desired
BERT_LR_CANDIDATES = [2e-5, 3e-5, 4e-5, 5e-5]
PREFERRED_METRIC_KEY = {
"cola": "matthews_correlation",
"sst2": "accuracy",
"mrpc": "accuracy",
"stsb": "pearson",
"qqp": "accuracy",
"mnli": "accuracy",
"qnli": "accuracy",
"rte": "accuracy",
"wnli": "accuracy",
"boolq": "accuracy",
"piqa": "accuracy",
"winogrande": "accuracy",
"hellaswag": "accuracy",
"openbookqa": "accuracy",
"arc_easy": "accuracy",
"arc_challenge": "accuracy",
}
# ---------------------------------------------------------------------
# Repro helpers
# ---------------------------------------------------------------------
def _set_all_seeds(seed: int):
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
try:
torch.cuda.manual_seed_all(seed)
except Exception:
pass
try:
torch.backends.cudnn.deterministic = True
torch.backends.cudnn.benchmark = False
except Exception:
pass
def _json_dump(obj: Any, path: Path):
path.parent.mkdir(parents=True, exist_ok=True)
with open(path, "w", encoding="utf-8") as f:
json.dump(obj, f, indent=2, ensure_ascii=False)
def _safe_float(x):
try:
if isinstance(x, (np.generic,)):
return float(x.item())
return float(x)
except Exception:
return None
def _metric_to_scalar(task: str, metric_res: Dict[str, Any], fallback_val_loss: Optional[float] = None) -> float:
if isinstance(metric_res, dict) and metric_res:
pref = PREFERRED_METRIC_KEY.get(task)
if pref is not None and pref in metric_res:
v = _safe_float(metric_res.get(pref))
if v is not None and not math.isnan(v):
return float(v)
for _, v in metric_res.items():
fv = _safe_float(v)
if fv is not None and not math.isnan(fv):
return float(fv)
if fallback_val_loss is not None:
try:
return -float(fallback_val_loss)
except Exception:
pass
return -1e9
# ---------------------------------------------------------------------
# Task-aware example field extraction (robust)
# ---------------------------------------------------------------------
def _get_text_pair_from_example(task: str, ex: dict):
"""
Robustly extract (s1, s2) from a HF GLUE/example dict `ex` depending on task.
Returns (s1:str, s2:Optional[str]) where s2 can be None for single-sentence tasks.
"""
task_field_map = {
"cola": ("sentence", None),
"sst2": ("sentence", None),
"mrpc": ("sentence1", "sentence2"),
"stsb": ("sentence1", "sentence2"),
"qqp": ("question1", "question2"),
"mnli": ("premise", "hypothesis"),
"qnli": ("question", "sentence"),
"rte": ("sentence1", "sentence2"),
"wnli": ("sentence1", "sentence2"),
}
f1, f2 = task_field_map.get(task, (None, None))
def _try_keys(keys):
for k in keys:
if k in ex and ex.get(k) is not None:
return ex.get(k)
return None
s1_candidates = []
s2_candidates = []
if f1:
s1_candidates.append(f1)
s1_candidates += ["sentence1", "premise", "question", "sentence", "text", "question1"]
if f2:
s2_candidates.append(f2)
s2_candidates += ["sentence2", "hypothesis", "question2", "question1", "text2"]
s1 = _try_keys(s1_candidates)
s2 = _try_keys(s2_candidates)
if s1 is None:
s1 = ex.get("sentence") or ex.get("premise") or ex.get("question") or ex.get("text")
if s2 is None:
s2 = ex.get("sentence2") or ex.get("hypothesis") or ex.get("question2")
s1 = "" if s1 is None else (s1 if isinstance(s1, str) else str(s1))
s2 = None if s2 is None else (s2 if isinstance(s2, str) else str(s2))
return s1, s2
# ---------------------------------------------------------------------
# Tokenization helpers (robust to different tokenizer APIs)
# ---------------------------------------------------------------------
def _pad_and_tensorize(input_ids_list, attention_mask_list, pad_token_id: int):
max_len = max(len(x) for x in input_ids_list) if input_ids_list else 0
ids_padded = [x + [pad_token_id] * (max_len - len(x)) for x in input_ids_list]
mask_padded = [m + [0] * (max_len - len(m)) for m in attention_mask_list]
input_ids = torch.tensor(ids_padded, dtype=torch.long)
attention_mask = torch.tensor(mask_padded, dtype=torch.long)
return input_ids, attention_mask
def _batch_tokenize(tokenizer, texts: List[Tuple[Optional[str], Optional[str]]], max_length: int = 128):
"""
Robust batch tokenization for a variety of tokenizer APIs.
- texts: list of (s1, s2) where s2 may be None.
- Try HF tokenizer(...) first, then various batch methods, then per-example fallback.
Returns dict with 'input_ids' (list of lists) and 'attention_mask'.
"""
sanitized = []
for a, b in texts:
a_s = "" if a is None else (a if isinstance(a, str) else str(a))
b_s = None if b is None else (b if isinstance(b, str) else str(b))
sanitized.append((a_s, b_s))
# 1) Try HF-like tokenizer(...) first
try:
flat = [(a if b is None else (a, b)) for a, b in sanitized]
enc = tokenizer(flat, truncation=True, padding=False, max_length=max_length)
if isinstance(enc.get("input_ids", None), torch.Tensor):
enc["input_ids"] = enc["input_ids"].tolist()
if isinstance(enc.get("attention_mask", None), torch.Tensor):
enc["attention_mask"] = enc["attention_mask"].tolist()
return enc
except Exception:
pass
# 2) Try other batch-like methods
for method_name in ("batch_encode", "encode_batch", "batch_encode_plus", "encode_batch_pair", "encode_batch_items"):
fn = getattr(tokenizer, method_name, None)
if fn is None:
continue
try:
try:
enc = fn(sanitized, max_length=max_length, truncation=True, padding=False)
except TypeError:
enc = fn(sanitized)
if isinstance(enc.get("input_ids", None), torch.Tensor):
enc["input_ids"] = enc["input_ids"].tolist()
if isinstance(enc.get("attention_mask", None), torch.Tensor):
enc["attention_mask"] = enc["attention_mask"].tolist()
return enc
except Exception:
continue
# 3) Fallback per-example
input_ids_list = []
attention_mask_list = []
for a, b in sanitized:
try:
if b is None:
try:
single = tokenizer.encode(a)
except TypeError:
single = tokenizer.encode([a])
else:
single = None
try:
single = tokenizer.encode((a, b))
except Exception:
try:
single = tokenizer.encode(a, b)
except Exception:
single = tokenizer(a if b is None else (a, b))
if isinstance(single, dict):
ids = single.get("input_ids") or single.get("ids") or []
mask = single.get("attention_mask") or single.get("mask") or [1] * len(ids)
elif isinstance(single, torch.Tensor):
ids = single.tolist()
mask = [1] * len(ids)
elif isinstance(single, list):
ids = single
mask = [1] * len(ids)
else:
tmp = tokenizer(a if b is None else (a, b))
if isinstance(tmp, dict):
ids = tmp.get("input_ids") or tmp.get("ids") or []
mask = tmp.get("attention_mask") or tmp.get("mask") or [1] * len(ids)
elif torch.is_tensor(tmp):
ids = tmp.tolist()
mask = [1] * len(ids)
else:
ids = list(tmp)
mask = [1] * len(ids)
if len(ids) > max_length:
ids = ids[:max_length]
mask = mask[:max_length]
input_ids_list.append(ids)
attention_mask_list.append(mask)
except Exception as e:
snippet = (a[:80] + "...") if a else "<empty>"
raise RuntimeError(f"Tokenizer fallback encode failed for example '{snippet}': {e}")
return {"input_ids": input_ids_list, "attention_mask": attention_mask_list}
# ---------------------------------------------------------------------
# Postprocess preds to the right shapes/types (fixes metric mismatches)
# ---------------------------------------------------------------------
def _postprocess_predictions(task: str, logits_np: np.ndarray, cfg_task: dict):
"""
Take logits (N, C) or (N,) or (N,1) and produce preds array ready for evaluate.compute:
- classification -> 1D ints (class indices or binary 0/1)
- regression -> 1D floats (for stsb typically 0..5)
"""
ttype = cfg_task["type"]
num_labels = cfg_task["num_labels"]
if logits_np is None or logits_np.size == 0:
return np.array([])
# If logits are shape (N, ) -> treat as single score per example (binary/regression)
if logits_np.ndim == 1:
if ttype == "classification":
preds = (logits_np > 0.5).astype(int)
else:
preds = logits_np.astype(float)
return preds
# If logits shape (N, 1)
if logits_np.ndim == 2 and logits_np.shape[1] == 1:
col = logits_np[:, 0]
if ttype == "classification":
preds = (col > 0.5).astype(int)
else:
preds = col.astype(float)
return preds
# If logits shape (N, C)
if logits_np.ndim == 2 and logits_np.shape[1] >= 1:
if ttype == "classification":
preds = np.argmax(logits_np, axis=-1).astype(int)
return preds
else:
if logits_np.shape[1] == 1:
preds = logits_np[:, 0].astype(float)
else:
preds = logits_np.mean(axis=1).astype(float)
if task == "stsb":
preds = np.clip(preds, 0.0, 5.0)
return preds
return logits_np.ravel()
# ---------------------------------------------------------------------
# Model wrapping helper (robust)
# ---------------------------------------------------------------------
def make_wrapped_model_if_needed(model, hidden_size: Optional[int], num_labels: int, force_num_labels: Optional[int] = None):
"""
Robust wrapper factory with resilient hidden_size inference.
Returns (model_or_wrapper, wrapped_flag)
"""
import torch.nn as nn
base_model = model
def _detect_head_dim(m):
if hasattr(m, "classifier") and isinstance(getattr(m, "classifier"), nn.Linear):
return getattr(m, "classifier").out_features
if hasattr(m, "lm_head") and isinstance(getattr(m, "lm_head"), nn.Linear):
return getattr(m, "lm_head").out_features
if hasattr(m, "get_output_embeddings"):
out_emb = m.get_output_embeddings()
if out_emb is not None:
if isinstance(out_emb, nn.Embedding):
return out_emb.embedding_dim if hasattr(out_emb, "embedding_dim") else out_emb.num_embeddings
if isinstance(out_emb, nn.Linear):
return out_emb.out_features
return None
if force_num_labels is None:
head_dim = _detect_head_dim(base_model)
if head_dim is not None and head_dim == num_labels:
return base_model, False
inferred_hidden = hidden_size
if inferred_hidden is None:
cand = getattr(base_model, "config", None)
if cand is not None and hasattr(cand, "hidden_size"):
try:
inferred_hidden = int(cand.hidden_size)
except Exception:
inferred_hidden = None
if inferred_hidden is None:
un = unwrap_model(base_model)
sd = un.state_dict()
for k, v in sd.items():
if re.search(r"embed|embedding|word_embeddings|token_embedding|embed_tokens", k, re.I):
if hasattr(v, "shape") and len(v.shape) == 2:
inferred_hidden = int(v.shape[1])
break
if re.search(r"q_proj|k_proj|v_proj|o_proj|dense|fc|linear|proj", k, re.I):
if hasattr(v, "shape") and len(v.shape) == 2:
cand = max(v.shape)
if 1 < cand < 1_000_000:
inferred_hidden = int(cand)
break
if inferred_hidden is None:
raise RuntimeError(
"Cannot infer hidden_size for wrapped classifier head. "
"Please set `model.config.hidden_size` or pass `hidden_size` explicitly."
)
class _WrappedModel(nn.Module):
def __init__(self, base, hidden_size, num_labels):
super().__init__()
self.base = base
self.classifier = nn.Linear(hidden_size, num_labels)
self.logits_projector = None
def forward(self, input_ids=None, attention_mask=None, labels=None, **kwargs):
out = None
try:
out = self.base(input_ids=input_ids, attention_mask=attention_mask, **kwargs)
except TypeError:
out = self.base(input_ids)
last_hidden = getattr(out, "last_hidden_state", None)
if last_hidden is not None:
pooled = last_hidden[:, 0, :]
logits = self.classifier(pooled)
return type("Out", (), {"logits": logits, "loss": None})
if isinstance(out, (tuple, list)) and len(out) > 0:
cand = out[0]
if torch.is_tensor(cand):
if cand.ndim == 3:
pooled = cand[:, 0, :]
logits = self.classifier(pooled)
return type("Out", (), {"logits": logits, "loss": None})
if cand.ndim == 2 and cand.shape[1] == num_labels:
return type("Out", (), {"logits": cand, "loss": None})
logits = getattr(out, "logits", None)
if logits is not None:
if logits.ndim == 2 and logits.shape[1] == num_labels:
return type("Out", (), {"logits": logits, "loss": getattr(out, "loss", None)})
exist_dim = logits.shape[1]
if self.logits_projector is None or self.logits_projector.weight.shape[1] != exist_dim:
self.logits_projector = nn.Linear(exist_dim, num_labels).to(logits.device)
projected = self.logits_projector(logits)
return type("Out", (), {"logits": projected, "loss": getattr(out, "loss", None)})
hidden_states = getattr(out, "hidden_states", None)
if hidden_states is not None:
last_hidden = hidden_states[-1] if isinstance(hidden_states, (list, tuple)) else hidden_states
if torch.is_tensor(last_hidden) and last_hidden.ndim == 3:
pooled = last_hidden[:, 0, :]
logits = self.classifier(pooled)
return type("Out", (), {"logits": logits, "loss": None})
raise RuntimeError("Wrapped base model did not return recognizable hidden states or logits")
return _WrappedModel(base_model, inferred_hidden, num_labels), True
# ---------------------------------------------------------------------
# Tokenize HF split to tensors (for finetune)
# ---------------------------------------------------------------------
def _tokenize_hf_split_to_tensors(task: str, tokenizer, raw_split, cfg_task, max_length=128, batch_tokenize_size=512):
texts = []
labels = []
empty_s1 = 0
empty_s2 = 0
for ex in raw_split:
s1, s2 = _get_text_pair_from_example(task, ex)
texts.append((s1, s2))
labels.append(ex.get("label") if "label" in ex else -100)
if not s1 or (isinstance(s1, str) and s1.strip() == ""):
empty_s1 += 1
if s2 is not None and (not s2 or (isinstance(s2, str) and s2.strip() == "")):
empty_s2 += 1
total = len(texts)
print(
f"[tokenize] task={task} samples={total} empty_s1={empty_s1} empty_s2={empty_s2} "
f"({(empty_s1/total if total>0 else 0):.2%}, {(empty_s2/total if total>0 else 0):.2%})"
)
input_ids_all = []
attention_all = []
for i in range(0, len(texts), batch_tokenize_size):
enc = _batch_tokenize(tokenizer, texts[i:i+batch_tokenize_size], max_length=max_length)
ids = enc.get("input_ids")
masks = enc.get("attention_mask") or enc.get("mask") or enc.get("masks")
if isinstance(ids, torch.Tensor):
ids = ids.tolist()
if isinstance(masks, torch.Tensor):
masks = masks.tolist()
input_ids_all.extend(ids)
attention_all.extend(masks)
pad_id = getattr(tokenizer, "pad_token_id", None)
if pad_id is None:
try:
pad_id = tokenizer.token_to_id("[PAD]")
except Exception:
pad_id = 0
input_ids_t, attention_mask_t = _pad_and_tensorize(input_ids_all, attention_all, pad_id)
labels_t = torch.tensor(labels, dtype=torch.long if cfg_task["type"] == "classification" else torch.float)
return input_ids_t, attention_mask_t, labels_t
# ---------------------------------------------------------------------
# Generic multiple-choice extractor & tokenizer
# ---------------------------------------------------------------------
# Improved unwrapping + robust context assembly and MC example extractor
def _maybe_unwrap(v):
"""
Improved unwrap:
- numpy / torch scalars -> Python scalars
- single-element lists/tuples -> unwrap
- dicts are left as-is (they may represent structured choices)
- numeric strings -> int when appropriate (helps labels like '3')
- lists of strings left as-is (caller may join)
"""
try:
import numpy as _np
import torch as _torch
except Exception:
_np = None
_torch = None
if v is None:
return None
# torch tensor scalar or 0-d array
if _torch is not None and isinstance(v, _torch.Tensor):
if v.ndim == 0:
return v.item()
if v.numel() == 1:
return v.view(-1).tolist()[0]
return v
# numpy scalar / array
if _np is not None and isinstance(v, _np.ndarray):
if v.shape == () or v.size == 1:
return v.flatten().tolist()[0]
return v.tolist()
# dict likely meaningful (choices dict) -> keep as-is
if isinstance(v, dict):
return v
# list/tuple with single element -> unwrap to element
if isinstance(v, (list, tuple)) and len(v) == 1:
return _maybe_unwrap(v[0])
# string that represents an integer -> convert to int (helps labels like '3')
if isinstance(v, str):
s = v.strip()
if s.isdigit():
try:
return int(s)
except Exception:
pass
return v
return v
def _assemble_context_from_example(ex: dict):
"""
Build a single string context from many possible fields.
Priority / heuristics:
1. If 'ctx' or 'context' present and non-empty -> use it.
2. If 'ctx_a' and 'ctx_b' (or 'ctxA'/'ctxB') present -> join them.
3. If 'premise' and 'hypothesis' present -> join them (suitable for MNLI-like).
4. Use question-oriented fields if that is the best we can do: question_stem, question, stem.
5. Fallback: join a selection of textual fields in a sensible order.
Converts lists/tuples of strings to joined text.
"""
def _to_text(x):
if x is None:
return ""
if isinstance(x, (list, tuple)):
# attempt to join list of strings or dicts
parts = []
for it in x:
if isinstance(it, dict):
txt = it.get("text") or it.get("choice") or it.get("label") or str(it)
parts.append(str(txt))
else:
parts.append(str(it))
return " ".join([p.strip() for p in parts if p is not None and str(p).strip() != ""])
if isinstance(x, dict):
# pick likely textual fields
for k in ("text", "content", "question", "sentence", "passage", "context"):
if k in x and x[k]:
return _to_text(x[k])
return str(x)
return str(x)
# common direct fields
keys = {k.lower(): v for k, v in ex.items()}
# 1) explicit ctx / context
for k in ("ctx", "context"):
if k in keys and keys[k]:
return _to_text(keys[k]).strip()
# 2) ctx_a + ctx_b variants
a_keys = ("ctx_a", "ctxA", "context_a", "contextA", "ctxa")
b_keys = ("ctx_b", "ctxB", "context_b", "contextB", "ctxb")
a_val = None
b_val = None
for ka in a_keys:
if ka in keys and keys[ka]:
a_val = keys[ka]
break
for kb in b_keys:
if kb in keys and keys[kb]:
b_val = keys[kb]
break
if a_val is not None or b_val is not None:
parts = []
if a_val is not None:
parts.append(_to_text(a_val))
if b_val is not None:
parts.append(_to_text(b_val))
return " ".join([p.strip() for p in parts if p and p.strip() != ""]).strip()
# 3) premise + hypothesis
if "premise" in keys or "hypothesis" in keys:
p = keys.get("premise")
h = keys.get("hypothesis")
parts = []
if p:
parts.append(_to_text(p))
if h:
parts.append(_to_text(h))
return " ".join([p.strip() for p in parts if p and p.strip() != ""]).strip()
# 4) question stem / question / stem
for k in ("question_stem", "questionStem", "question", "stem", "prompt", "goal"):
if k.lower() in keys and keys[k.lower()]:
return _to_text(keys[k.lower()]).strip()
# 5) passage/article/story/sentence
for k in ("passage", "article", "story", "sentence", "paragraph"):
if k in keys and keys[k]:
return _to_text(keys[k]).strip()
# 6) try to build a concat from any textual fields in a sensible order
fallback_fields = [
"context", "question", "passage", "article", "story", "sentence",
"prompt", "goal", "stem", "query", "description", "narration"
]
parts = []
for f in fallback_fields:
if f in keys and keys[f]:
parts.append(_to_text(keys[f]))
combined = " ".join([p.strip() for p in parts if p and p.strip() != ""]).strip()
if combined:
return combined
# final fallback: try first non-empty string-ish field
for k, v in ex.items():
if isinstance(v, str) and v.strip():
return v.strip()
if isinstance(v, (list, tuple)) and len(v) > 0 and all(isinstance(x, str) for x in v):
return " ".join(v).strip()
return ""
def _extract_mc_example(task: str, ex: dict):
"""
Robust extractor for multiple-choice examples across HF dataset schemas.
Returns (context:str, options_list:List[str], label_index_or_None)
Improvements over earlier versions:
- Uses _maybe_unwrap on fields
- Assembles context robustly via _assemble_context_from_example (handles ctx_a+ctx_b, premise+hypothesis, etc.)
- Accepts choices as:
* dict {'text': [...], 'label': [...]} (OpenBookQA / ARC)
* list of dicts or strings
* 'endings' list (HellaSwag)
* explicit fields choice1/option1 etc.
- Normalizes label forms: 'A'..'D' -> 0..3, '3' -> int(3), torch/numpy scalars handled
"""
def _map_label_to_int(label, opts):
if label is None:
return None
# unwrap containers
if isinstance(label, (list, tuple)) and len(label) > 0:
label = label[0]
try:
import numpy as _np
import torch as _torch
except Exception:
_np = None
_torch = None
if _torch is not None and isinstance(label, _torch.Tensor):
try:
return int(label.item())
except Exception:
pass
if _np is not None and isinstance(label, _np.ndarray):
if label.size == 1:
return int(label.flatten().tolist()[0])
if isinstance(label, str):
s = label.strip()
if len(s) == 1 and s.isalpha():
return ord(s.upper()) - ord("A")
if s.isdigit():
try:
return int(s)
except Exception:
pass
# if label equals one of option texts, return that index
if opts:
for i, o in enumerate(opts):
if isinstance(o, str) and s == o:
return i
return None
if isinstance(label, (int, np.integer)):
return int(label)
return None
# shallow normalized copy
ex_norm = {}
for k, v in ex.items():
try:
ex_norm[k] = _maybe_unwrap(v)
except Exception:
ex_norm[k] = v
ex = ex_norm
# assemble context early
ctx = _assemble_context_from_example(ex)
# 1) OpenBookQA / ARC-style: 'choices' is dict with 'text' list
if "choices" in ex and isinstance(ex["choices"], dict):
chd = ex["choices"]
texts = chd.get("text") or chd.get("texts") or chd.get("choice") or None
if isinstance(texts, (list, tuple)) and len(texts) > 0:
opts = [str(x) for x in texts]
lab = ex.get("answerKey") or ex.get("answer") or ex.get("correctAnswer") or ex.get("label")
lab_idx = _map_label_to_int(lab, opts)
return ctx, opts, lab_idx
# 2) HellaSwag style 'endings'
if "endings" in ex and isinstance(ex["endings"], list) and len(ex["endings"]) > 0:
opts = [str(x) for x in ex["endings"]]
lab = ex.get("label") or ex.get("answerKey") or ex.get("answer")
lab_idx = _map_label_to_int(lab, opts)
return ctx, opts, lab_idx
# 3) 'choices' as list (strings or dicts)
if "choices" in ex and ex["choices"] is not None:
ch = ex["choices"]
if isinstance(ch, list) and len(ch) > 0:
opts = []
for c in ch:
if isinstance(c, dict):
opts.append(c.get("text") or c.get("choice") or c.get("label") or str(c))
else:
opts.append(str(c))
lab = ex.get("answerKey") or ex.get("answer") or ex.get("label") or ex.get("correct")
lab_idx = _map_label_to_int(lab, opts)
return ctx, opts, lab_idx
# 4) explicit option fields like 'choice1','choice2' or 'option1'..
opts = []
for prefix in ("choice", "option", "ending", "answer"):
i = 1
found = False
while True:
key = f"{prefix}{i}"
if key in ex:
val = _maybe_unwrap(ex.get(key))
opts.append(str(val))
found = True
i += 1
else:
break
if found:
lab = ex.get("label") or ex.get("answerKey") or ex.get("answer")
lab_idx = _map_label_to_int(lab, opts)
return ctx, opts, lab_idx
# 5) fallback: look for list-valued fields that look like options
for k, v in ex.items():
if isinstance(v, (list, tuple)) and 2 <= len(v) <= 10 and all(isinstance(x, (str, dict)) for x in v):
opts = [x.get("text") if isinstance(x, dict) and x.get("text") else str(x) for x in v]
lab = ex.get("answer") or ex.get("label") or ex.get("answerKey")
lab_idx = _map_label_to_int(lab, opts)
return ctx, opts, lab_idx
# 6) collect candidate option-fields heuristically
candidate_opts = []
for k in sorted(ex.keys()):
if any(tok in k.lower() for tok in ("option", "choice", "ending", "answer", "alt", "sol", "choices")):
val = _maybe_unwrap(ex.get(k))
candidate_opts.append(str(val))
if candidate_opts:
lab = ex.get("label") or ex.get("answer") or ex.get("answerKey")
lab_idx = _map_label_to_int(lab, candidate_opts)
return ctx, candidate_opts, lab_idx
# last resort: no options discovered
lab = ex.get("label") or ex.get("answerKey") or ex.get("answer")
lab_idx = _map_label_to_int(lab, [])
return ctx, [], lab_idx
# ---------------------------
# 新增:detect_num_choices
# ---------------------------
def detect_num_choices(raw_split, task: str, max_samples: int = 200):
max_s = min(len(raw_split), max_samples)
# 强制拿出前 max_s 个示例为「标量形式」:select 确保每个示例是 dict of scalars 而不是子 Dataset
subset = raw_split.select(range(max_s)) if hasattr(raw_split, "select") else raw_split[:max_s]
counts = {}
total = 0
kept = 0
bad_label = 0
for ex in subset:
total += 1
try:
ctx, opts, lab = _extract_mc_example(task, ex)
except Exception:
counts.setdefault("no_opts", 0)
counts["no_opts"] += 1
continue
if not opts:
counts.setdefault("no_opts", 0)
counts["no_opts"] += 1
else:
kept += 1
counts.setdefault(len(opts), 0)
counts[len(opts)] += 1
if lab is None or (isinstance(lab, int) and lab < 0):
bad_label += 1
numeric = {k: v for k, v in counts.items() if isinstance(k, int)}
detected = None
if numeric:
detected = max(numeric.items(), key=lambda x: x[1])[0]
summary = {
"task": task,
"total_scanned": total,
"kept": kept,
"counts": counts,
"detected_num_choices": detected,
"skipped_bad_label": bad_label,
}
print(f"[detect_num_choices][{task}] scanned={total} kept={kept} detected={detected} counts={counts} bad_label={bad_label}")
return summary
# ---------------------------
# 更新:_tokenize_generic_mc_split_to_tensors
# ---------------------------
def _tokenize_generic_mc_split_to_tensors(task: str, tokenizer, raw_split, max_length=128, batch_tokenize_size=256, expect_num_choices: Optional[int] = None):
"""
Generic MC tokenizer with optional expect_num_choices enforcement.
Returns (input_ids_t (N,C,L), attention_t (N,C,L), labels_t (N,)), and prints debug counts.
If expect_num_choices is provided, examples whose number of options != expect_num_choices are considered 'inconsistent' and skipped.
This version does not swallow exceptions: tokenizer errors and unexpected shapes will raise.
"""
contexts = []
options = []
labels = []
total = 0
skipped_no_opts = 0
skipped_inconsistent = 0
skipped_bad_label = 0
kept = 0
for ex in raw_split:
ctx, opts, lab = _extract_mc_example(task, ex)
total += 1
# print(task)
# print(ex)
# print(ctx)
# print(opts)
# print(lab)
# import pdb
# pdb.set_trace()
# print(ex)
# print(ctx)
# print(opts)
# print(lab)
if not opts:
skipped_no_opts += 1
continue
if expect_num_choices is not None and len(opts) != expect_num_choices:
skipped_inconsistent += 1
continue
if lab is None:
skipped_bad_label += 1
continue
try:
lab_i = int(lab)
except Exception:
skipped_bad_label += 1
continue
if lab_i < 0 or lab_i >= len(opts):
skipped_bad_label += 1
continue
contexts.append(ctx if ctx is not None else "")
options.append(opts)
labels.append(lab_i)
kept += 1
print(f"[tokenize_generic_mc][{task}] total={total} kept={kept} skipped_no_opts={skipped_no_opts} skipped_inconsistent={skipped_inconsistent} skipped_bad_label={skipped_bad_label}")
if len(contexts) == 0:
return torch.zeros((0, 1, 1), dtype=torch.long), torch.zeros((0, 1, 1), dtype=torch.long), torch.tensor([], dtype=torch.long)
pad_id = getattr(tokenizer, "pad_token_id", None)
if pad_id is None:
try:
pad_id = tokenizer.token_to_id("[PAD]")
except Exception:
pad_id = 0
num_choices = len(options[0])
# Build flat pairs for batch tokenization and assert tokenizer returns expected length
input_ids_rows = []
attention_rows = []
for i in range(0, len(contexts), batch_tokenize_size):
chunk_ctx = contexts[i:i+batch_tokenize_size]
chunk_opts = options[i:i+batch_tokenize_size]
flat_pairs = []
for c, opts in zip(chunk_ctx, chunk_opts):
for o in opts:
flat_pairs.append((c, o))
enc = tokenizer(flat_pairs, truncation=True, padding=False, max_length=max_length)
ids_flat = enc.get("input_ids")
masks_flat = enc.get("attention_mask") or enc.get("mask") or enc.get("masks")
if isinstance(ids_flat, torch.Tensor):
ids_flat = ids_flat.tolist()
if isinstance(masks_flat, torch.Tensor):
masks_flat = masks_flat.tolist()
expected = len(chunk_ctx) * num_choices
if len(ids_flat) != expected:
# raise explicit error so user sees where tokenizer mismatched
raise RuntimeError(f"[tokenize_generic_mc][{task}] tokenizer returned {len(ids_flat)} items but expected {expected} (i={i}). "
f"First flat_pairs sample: {flat_pairs[0] if len(flat_pairs)>0 else None}")
per_example = []
per_mask_example = []
idx = 0
for _ in range(len(chunk_ctx)):
row = []
row_mask = []
for _ in range(num_choices):
row.append(ids_flat[idx])
row_mask.append(masks_flat[idx])
idx += 1
per_example.append(row)
per_mask_example.append(row_mask)
input_ids_rows.extend(per_example)
attention_rows.extend(per_mask_example)
max_len = max(len(seq) for row in input_ids_rows for seq in row) if input_ids_rows else 1
input_ids_padded = [
[ seq + [pad_id] * (max_len - len(seq)) for seq in row ]
for row in input_ids_rows
]
attention_padded = [
[ mask + [0] * (max_len - len(mask)) for mask in row ]
for row in attention_rows
]
input_ids_t = torch.tensor(input_ids_padded, dtype=torch.long) # (N, C, L)
attention_t = torch.tensor(attention_padded, dtype=torch.long)
labels_t = torch.tensor(labels, dtype=torch.long)
return input_ids_t, attention_t, labels_t
# ---------------------------
# train_full_finetune_extra(替换版:用替换 classifier 输出 1 的方法)
# ---------------------------
def train_full_finetune_extra(task: str, tokenizer, model, raw_train, raw_val,
device: str = "cuda", epochs: int = 3, batch_size: int = 16,
lr: float = 2e-5, weight_decay: float = 0.01, warmup_steps: int = 100,
max_length: int = 128, grad_accum_steps: int = 1, out_checkpoint_dir: Optional[str] = None,
expect_num_choices: Optional[int] = None):
"""
Fine-tune for EXTRA_TASKS with replacing the model's classification head to output 1 scalar:
- Replace model.classifier / model.lm_head / model.score / similar with Linear(hidden, 1)
- For MC tasks: flatten (B, C, L) -> (B*C, L), forward, obtain (B*C,1) -> reshape (B, C) -> CE loss
- Ensures new head parameters are created before optimizer so they are optimized.
This variant surfaces errors (no silent swallowing).
"""
import torch.nn as nn
cfg = EXTRA_TASKS[task]
device = torch.device(device if torch.cuda.is_available() else "cpu")
model.to(device)
is_mc = cfg["format"] == "mc"
# ---------------------------
# Tokenize splits
# ---------------------------
if is_mc:
train_ids, train_mask, train_labels = _tokenize_generic_mc_split_to_tensors(
task, tokenizer, raw_train, max_length=max_length, batch_tokenize_size=256, expect_num_choices=expect_num_choices
)
val_ids, val_mask, val_labels = _tokenize_generic_mc_split_to_tensors(
task, tokenizer, raw_val, max_length=max_length, batch_tokenize_size=256, expect_num_choices=expect_num_choices
)
if train_ids.ndim != 3:
raise RuntimeError(f"[ExtraTrain][{task}] train_ids expected 3 dims (N,C,L) got {train_ids.ndim}")
if val_ids.ndim != 3 and len(val_ids) > 0:
raise RuntimeError(f"[ExtraTrain][{task}] val_ids expected 3 dims (N,C,L) got {val_ids.ndim}")
train_ds = TensorDataset(train_ids, train_mask, train_labels)
val_ds = TensorDataset(val_ids, val_mask, val_labels)
else:
train_ids, train_mask, train_labels = _tokenize_hf_split_to_tensors(task, tokenizer, raw_train, cfg, max_length=max_length)
val_ids, val_mask, val_labels = _tokenize_hf_split_to_tensors(task, tokenizer, raw_val, cfg, max_length=max_length)
train_ds = TensorDataset(train_ids, train_mask, train_labels)
val_ds = TensorDataset(val_ids, val_mask, val_labels)
if len(train_ds) == 0:
raise RuntimeError(f"No training samples after tokenization for task={task}; aborting.")
train_loader = DataLoader(train_ds, batch_size=batch_size, shuffle=True, pin_memory=True)
val_loader = DataLoader(val_ds, batch_size=max(64, batch_size), shuffle=False, pin_memory=True)
# ---------------------------
# Replace model head to output scalar per example (for MC we'll apply it to each choice)
# ---------------------------
def _infer_hidden_size_from_model(m):
# try config.hidden_size first
hidden = None
try:
if hasattr(m, "config") and hasattr(m.config, "hidden_size"):
hidden = int(m.config.hidden_size)
except Exception:
hidden = None
if hidden is not None:
return hidden
# try to inspect state_dict looking for embedding/proj dimensions
try:
un = unwrap_model(m)
sd = un.state_dict()
for k, v in sd.items():
if re.search(r"embed|embedding|word_embeddings|token_embedding|embed_tokens", k, re.I):
if hasattr(v, "shape") and len(v.shape) == 2:
return int(v.shape[1])
if re.search(r"q_proj|k_proj|v_proj|o_proj|dense|fc|linear|proj", k, re.I):
if hasattr(v, "shape") and len(v.shape) == 2:
cand = max(v.shape)
if 1 < cand < 1_000_000:
return int(cand)
except Exception:
pass
return None
# Only replace head for MC and for pair tasks we keep existing pipeline (but we still may want to ensure classifier shape)
if is_mc:
# detect suitable hidden size
hidden_size = _infer_hidden_size_from_model(model)
if hidden_size is None:
raise RuntimeError("Cannot infer hidden size from model; please set model.config.hidden_size or pass hidden_size explicitly.")
# Candidate head attribute names to replace
head_attrs = ["classifier", "lm_head", "score", "classifier_head", "head"]
replaced = False
for attr in head_attrs:
if hasattr(model, attr):
try:
old = getattr(model, attr)
# create new linear head that maps hidden -> 1
new_head = nn.Linear(hidden_size, 1).to(device)
# init reasonably
try:
nn.init.xavier_uniform_(new_head.weight)
if new_head.bias is not None:
nn.init.zeros_(new_head.bias)
except Exception:
pass
setattr(model, attr, new_head)
print(f"[ExtraTrain][{task}] Replaced model.{attr} with Linear({hidden_size},1)")
replaced = True
break
except Exception:
# ignore and try next
pass
if not replaced:
# fallback: attach as attribute _mc_choice_head
model._mc_choice_head = nn.Linear(hidden_size, 1).to(device)
try:
nn.init.xavier_uniform_(model._mc_choice_head.weight)
if model._mc_choice_head.bias is not None:
nn.init.zeros_(model._mc_choice_head.bias)
except Exception:
pass
print(f"[ExtraTrain][{task}] Attached model._mc_choice_head = Linear({hidden_size},1)")
# ---------------------------
# Create optimizer AFTER head replacement so head params are included
# ---------------------------
optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=weight_decay)
total_steps = max(1, (len(train_loader) // max(1, grad_accum_steps)) * epochs)
from transformers import get_cosine_schedule_with_warmup
scheduler = get_cosine_schedule_with_warmup(optimizer, num_warmup_steps=warmup_steps, num_training_steps=total_steps)
ce_loss_fn = torch.nn.CrossEntropyLoss()
mse_loss_fn = torch.nn.MSELoss()
model.train()
global_step = 0
final_metric_res = {}
expected_num_labels = cfg["num_labels"]
# training loop
for epoch in range(epochs):
running_loss = 0.0
for step, batch in enumerate(tqdm(train_loader, desc=f"[ExtraTrain] {task} epoch {epoch+1}")):
if is_mc:
ids_b, mask_b, labs_b = batch # ids_b: (B, C, L)
ids_b = ids_b.to(device); mask_b = mask_b.to(device); labs_b = labs_b.to(device)
if labs_b.dim() > 1:
labs_b = labs_b.view(-1)
B = ids_b.size(0)
if labs_b.numel() != B:
raise RuntimeError(f"[ExtraTrain][FATAL] label count ({labs_b.numel()}) != batch size ({B}) at step {step}; aborting.")
# flatten (B,C,L) -> (B*C, L)
Bf, Cf, Lf = ids_b.shape
flat_ids = ids_b.view(Bf * Cf, Lf).to(device)
flat_mask = mask_b.view(Bf * Cf, Lf).to(device) if mask_b is not None else None
# forward
out_flat = None
try:
if flat_mask is not None:
out_flat = model(input_ids=flat_ids, attention_mask=flat_mask, labels=None)
else:
out_flat = model(input_ids=flat_ids, labels=None)
except TypeError:
# Some models accept single positional tensor
out_flat = model(flat_ids)
# get a pooled representation or logits from out_flat
flat_logits = getattr(out_flat, "logits", None)
pooled = None
if flat_logits is None:
# try pooler_output or last_hidden_state
pooled = getattr(out_flat, "pooler_output", None)
if pooled is None:
lh = getattr(out_flat, "last_hidden_state", None)
if lh is not None:
pooled = lh[:, 0, :]
else:
# If model already returns logits of shape (B*C, D) or (B*C,1), we can still feed them to head or reduce
pass
# compute scalar per flat example using the replaced head
# Case 1: model already returned logits as scalar (B*C,1) or (B*C,)
scalar_flat = None
if flat_logits is not None:
# flat_logits could be (B*C,) or (B*C,1) or (B*C,D)
if flat_logits.ndim == 1:
scalar_flat = flat_logits.view(-1, 1)
elif flat_logits.ndim == 2 and flat_logits.shape[1] == 1:
scalar_flat = flat_logits.view(-1, 1)
elif flat_logits.ndim == 2:
# project via the replaced head (we expect a head that maps hidden->1, so treat flat_logits as features)
# If we replaced model.classifier, it's likely expecting hidden vectors, but we have flat_logits features:
# attempt to use model._mc_choice_head if present; else take mean over dims then linear
feat = flat_logits
if hasattr(model, "_mc_choice_head"):
# ensure dimensions align: if feat dim matches head.in_features, use it
try:
in_f = model._mc_choice_head.in_features
if feat.shape[1] == in_f:
scalar_flat = model._mc_choice_head(feat)
else:
# reduce feat to hidden via mean
reduced = feat.mean(dim=1)
scalar_flat = model._mc_choice_head(reduced.unsqueeze(1) if len(reduced.shape)==1 else reduced)
except Exception:
reduced = feat.mean(dim=1)
scalar_flat = model._mc_choice_head(reduced.unsqueeze(1) if len(reduced.shape)==1 else reduced)
else:
# fallback: mean over feature dims to scalar
scalar_flat = feat.mean(dim=1, keepdim=True)
elif flat_logits.ndim == 3:
# reduce seq dim and feature dim to scalar
scalar_flat = flat_logits.mean(dim=tuple(range(1, flat_logits.ndim))).view(-1, 1)
else:
scalar_flat = flat_logits.reshape(flat_logits.size(0), -1).mean(dim=1, keepdim=True)
elif pooled is not None:
# pooled is (B*C, hidden), apply the replaced head
# find which head attribute we replaced
head_found = None
for attr in ("classifier", "lm_head", "score", "classifier_head", "head", "_mc_choice_head"):
if hasattr(model, attr):
head_found = getattr(model, attr)
break
if head_found is None:
raise RuntimeError("[ExtraTrain] no replacement head found on model to project pooled -> scalar")
# head_found may expect input dim hidden_size
try:
scalar_flat = head_found(pooled)
except Exception as e:
# if mismatch, try a linear mapping by reducing pooled
if pooled.ndim == 2:
# If head expects 1-d input, ensure shape matches
try:
scalar_flat = head_found(pooled)
except Exception:
scalar_flat = pooled.mean(dim=1, keepdim=True)
else:
scalar_flat = pooled.mean(dim=1, keepdim=True)
else:
raise RuntimeError("[ExtraTrain] cannot obtain features or logits from model forward to compute scalar per choice")
# Now scalar_flat should be (B*C, 1) -> reshape to (B, C)
if scalar_flat is None:
raise RuntimeError("[ExtraTrain] scalar_flat is None after processing model output")
if scalar_flat.ndim == 1:
scalar_flat = scalar_flat.view(-1, 1)
choice_scores = scalar_flat.view(Bf, Cf)
# If dataset's number of choices (Cf) differs from expected_num_labels, project to expected
if Cf != expected_num_labels:
# create a mapping layer if not exist (choice_count -> expected_num_labels)
if not hasattr(model, "_mc_choice_to_label"):
model._mc_choice_to_label = nn.Linear(Cf, expected_num_labels).to(device)
# register params in optimizer
try:
optimizer.add_param_group({'params': model._mc_choice_to_label.parameters()})
except Exception:
pass
print(f"[ExtraTrain][{task}] created _mc_choice_to_label: {Cf}->{expected_num_labels}")
used_logits = model._mc_choice_to_label(choice_scores)
else:
used_logits = choice_scores
if used_logits.ndim != 2 or used_logits.size(0) != labs_b.size(0):
raise RuntimeError(f"Logits/labels shape mismatch used_logits={tuple(used_logits.shape)} labels={tuple(labs_b.shape)} at step {step}")
loss = ce_loss_fn(used_logits, labs_b.long())
else:
ids_b, mask_b, labs_b = batch
ids_b = ids_b.to(device); mask_b = mask_b.to(device); labs_b = labs_b.to(device)
out = model(input_ids=ids_b, attention_mask=mask_b, labels=None)
logits = getattr(out, "logits", None)
if logits is None and isinstance(out, (tuple, list)):
logits = out[0]
if logits is None:
raise RuntimeError("Model did not return logits for pair task")
if logits.ndim == 2 and logits.shape[1] != expected_num_labels:
exist_dim = logits.shape[1]
model._mc_projector = nn.Linear(exist_dim, expected_num_labels).to(device)
try:
optimizer.add_param_group({'params': model._mc_projector.parameters()})
except Exception:
pass
used_logits = model._mc_projector(logits)
else:
used_logits = logits
if cfg["type"] == "classification":
loss = ce_loss_fn(used_logits, labs_b.long())
else:
if used_logits.ndim == 2 and used_logits.shape[1] == 1:
preds = used_logits.squeeze(1)
elif used_logits.ndim == 2:
preds = used_logits.mean(dim=1)
else:
preds = used_logits
loss = mse_loss_fn(preds, labs_b.float())
loss = loss / max(1, grad_accum_steps)
loss.backward()
if (step + 1) % max(1, grad_accum_steps) == 0:
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
scheduler.step()
optimizer.zero_grad()
global_step += 1
running_loss += float(loss.item()) * (ids_b.size(0) if is_mc else ids_b.size(0))
# validation
model.eval()
all_logits = []
all_labels = []
with torch.no_grad():
for batch in tqdm(val_loader, desc=f"[ExtraVal] {task} epoch {epoch+1}", leave=False):
if is_mc:
ids_b, mask_b, labs_b = batch
if labs_b.dim() > 1:
labs_b = labs_b.view(-1)
B = ids_b.size(0)
if labs_b.numel() != B:
raise RuntimeError(f"[ExtraVal][FATAL] val labels count ({labs_b.numel()}) != batch size ({B})")
ids_b = ids_b.to(device); mask_b = mask_b.to(device); labs_b = labs_b.to(device)
Bf, Cf, Lf = ids_b.shape
flat_ids = ids_b.view(Bf*Cf, Lf).to(device)
flat_mask = mask_b.view(Bf*Cf, Lf).to(device) if mask_b is not None else None
try:
if flat_mask is not None:
out_flat = model(input_ids=flat_ids, attention_mask=flat_mask, labels=None)
else:
out_flat = model(input_ids=flat_ids, labels=None)
except TypeError:
out_flat = model(flat_ids)
flat_logits = getattr(out_flat, "logits", None)
pooled = None
if flat_logits is None:
pooled = getattr(out_flat, "pooler_output", None)
if pooled is None:
lh = getattr(out_flat, "last_hidden_state", None)
if lh is not None:
pooled = lh[:, 0, :]
scalar_flat = None
if flat_logits is not None:
if flat_logits.ndim == 1:
scalar_flat = flat_logits.view(-1, 1)
elif flat_logits.ndim == 2 and flat_logits.shape[1] == 1:
scalar_flat = flat_logits.view(-1, 1)
elif flat_logits.ndim == 2:
if hasattr(model, "_mc_choice_head"):
try:
in_f = model._mc_choice_head.in_features
if flat_logits.shape[1] == in_f:
scalar_flat = model._mc_choice_head(flat_logits)
else:
scalar_flat = flat_logits.mean(dim=1, keepdim=True)
except Exception:
scalar_flat = flat_logits.mean(dim=1, keepdim=True)
else:
scalar_flat = flat_logits.mean(dim=1, keepdim=True)
elif flat_logits.ndim == 3:
scalar_flat = flat_logits.mean(dim=tuple(range(1, flat_logits.ndim))).view(-1, 1)
else:
scalar_flat = flat_logits.reshape(flat_logits.size(0), -1).mean(dim=1, keepdim=True)
elif pooled is not None:
head_found = None
for attr in ("classifier", "lm_head", "score", "classifier_head", "head", "_mc_choice_head"):
if hasattr(model, attr):
head_found = getattr(model, attr)
break
if head_found is None:
raise RuntimeError("[ExtraVal] no replacement head found on model to project pooled -> scalar")
try:
scalar_flat = head_found(pooled)
except Exception:
scalar_flat = pooled.mean(dim=1, keepdim=True)
else:
raise RuntimeError("Model did not return logits/poolable outputs during MC validation.")
choice_scores = scalar_flat.view(Bf, Cf)
if choice_scores.shape[1] != expected_num_labels:
exist_dim = choice_scores.shape[1]
if not hasattr(model, "_mc_choice_to_label"):
model._mc_choice_to_label = nn.Linear(exist_dim, expected_num_labels).to(device)
try:
optimizer.add_param_group({'params': model._mc_choice_to_label.parameters()})
except Exception:
pass
used_logits = model._mc_choice_to_label(choice_scores)
else:
used_logits = choice_scores
all_logits.append(used_logits.detach().cpu().numpy())
all_labels.append(labs_b.detach().cpu().numpy())
else:
ids_b, mask_b, labs_b = batch
ids_b = ids_b.to(device); mask_b = mask_b.to(device); labs_b = labs_b.to(device)
out = model(input_ids=ids_b, attention_mask=mask_b, labels=None)
logits = getattr(out, "logits", None)
if logits is None and isinstance(out, (tuple, list)):
logits = out[0]
if logits.ndim == 2 and logits.shape[1] != expected_num_labels:
exist_dim = logits.shape[1]
if not hasattr(model, "_mc_projector"):
model._mc_projector = nn.Linear(exist_dim, expected_num_labels).to(device)
try:
optimizer.add_param_group({'params': model._mc_projector.parameters()})
except Exception:
pass
used_logits = model._mc_projector(logits)
else:
used_logits = logits
all_logits.append(used_logits.detach().cpu().numpy())
all_labels.append(labs_b.detach().cpu().numpy())
model.train()
all_logits = np.concatenate(all_logits, axis=0) if all_logits else np.zeros((0, cfg["num_labels"]))
all_labels = np.concatenate(all_labels, axis=0) if all_labels else np.zeros((0,))
if cfg["type"] == "classification" or cfg["type"] == "multiple_choice":
preds = np.argmax(all_logits, axis=-1) if all_logits.size else np.array([])
else:
preds = _postprocess_predictions(task, all_logits, {"type": cfg["type"], "num_labels": cfg["num_labels"]})
metric = evaluate.load("accuracy")
metric_res = metric.compute(predictions=preds.tolist(), references=all_labels.tolist())
print(f"[Extra FT] {task} epoch {epoch+1} metric={metric_res}")
final_metric_res = metric_res
if out_checkpoint_dir:
outp = Path(out_checkpoint_dir)
outp.mkdir(parents=True, exist_ok=True)
fname = outp / "finetuned_extra.pt"
sd = unwrap_model(model).state_dict()
torch.save(sd, str(fname))
print(f"[Extra FT] Saved finetuned model to: {fname}")
return model, final_metric_res
# ---------------------------
# 辅助:打印 HellaSwag 样例(用来“给我看一下”)
# ---------------------------
def print_hellaswag_samples(n: int = 5):
"""
加载 hellaswag 数据集并打印前 n 个样例经 _extract_mc_example 解析后的 context/options/label。
运行时会在 stdout 中输出,方便你检查数据格式。
"""
ds = load_dataset("hellaswag")
split = ds.get("validation") or ds.get("train") or ds.get("test")
if split is None:
print("[print_hellaswag_samples] no split")
return
for i, ex in enumerate(split[:n]):
ctx, opts, lab = _extract_mc_example("hellaswag", ex)
print(f"=== sample {i} ===")
print("context:", (ctx[:400] + "...") if ctx is not None and len(str(ctx))>400 else ctx)
print("num_options:", len(opts))
for j, o in enumerate(opts):
print(f" [{j}] {o[:200]}{'...' if len(str(o))>200 else ''}")
print("label:", lab)
print()
# ---------------------------
# _tokenize_mc_split_to_tensors (general MC tokenizer)
# ---------------------------
def _tokenize_mc_split_to_tensors(task: str, tokenizer, raw_split, max_length=128, batch_tokenize_size=256):
"""
Generalized multiple-choice tokenizer that supports:
- PIQA (goal/sol1/sol2)
- Winogrande (sentence/option1/option2)
- ARC-like (question + choices, answerKey)
- RACE/other formats with 'options'/'choices' (list of strings or list of dicts with 'text')
Returns:
input_ids: (N, C, L)
attention_mask: (N, C, L)
labels: (N,) with -1 for unknown/no-label examples
This function performs strict checks and will raise on inconsistencies (no silent catches).
"""
contexts = []
options = []
labels = []
def _extract_choices_text(choices_field):
out = []
for c in choices_field:
if isinstance(c, str):
out.append(c)
elif isinstance(c, dict):
txt = c.get("text") or c.get("content") or c.get("choice") or c.get("label")
out.append(txt if txt is not None else str(c))
else:
out.append(str(c))
return out
def _map_answer_to_index(ans, choice_texts):
if ans is None:
return -1
if isinstance(ans, (list, tuple)):
if len(ans) == 0:
return -1
ans = ans[0]
if isinstance(ans, str):
s = ans.strip()
if len(s) == 1 and s.isalpha():
idx = ord(s.upper()) - ord("A")
if 0 <= idx < len(choice_texts):
return idx
try:
ni = int(s)
if 0 <= ni < len(choice_texts):
return ni
if 1 <= ni <= len(choice_texts):
return ni - 1
except Exception:
pass
for i, ct in enumerate(choice_texts):
if isinstance(ct, str) and s == ct:
return i
for i, ct in enumerate(choice_texts):
if isinstance(ct, str) and s in ct:
return i
low = s.lower()
if low in ("true", "t", "yes", "y"):
return 1
if low in ("false", "f", "no", "n"):
return 0
return -1
if isinstance(ans, (int, np.integer)):
ai = int(ans)
if 0 <= ai < len(choice_texts):
return ai
if 1 <= ai <= len(choice_texts):
return ai - 1
return -1
return -1
num_choices = None
for ex in raw_split:
ctx = None
opts = None
lab = None
# PIQA
if "goal" in ex and ("sol1" in ex or "sol2" in ex):
ctx = ex.get("goal")
opts = [ex.get("sol1"), ex.get("sol2")]
lab = ex.get("label") or ex.get("answer") or ex.get("answerKey")
if isinstance(lab, (str, np.str_)) and str(lab).isdigit():
lab = int(lab)
# Winogrande
elif "sentence" in ex and ("option1" in ex or "option2" in ex):
ctx = ex.get("sentence")
opts = [ex.get("option1"), ex.get("option2")]
lab = ex.get("answer") or ex.get("label") or ex.get("answerKey")
if isinstance(lab, str) and lab.isdigit():
lab = int(lab) - 1 # convert 1-based -> 0-based
# ARC / RACE-style: question + choices (choices may be list of strings or list of dicts)
elif "question" in ex and ("choices" in ex or "options" in ex):
ctx = ex.get("question")
choices_field = ex.get("choices") or ex.get("options")
opts = _extract_choices_text(choices_field)
lab = ex.get("answerKey") or ex.get("answer") or ex.get("correct_answer") or ex.get("label")
# Generic fallback: try known pairs
elif "context" in ex and "options" in ex:
ctx = ex.get("context")
opts = _extract_choices_text(ex.get("options"))
lab = ex.get("label") or ex.get("answer")
else:
if "question" in ex:
ctx = ex.get("question")
elif "query" in ex:
ctx = ex.get("query")
elif "prompt" in ex:
ctx = ex.get("prompt")
else:
for k, v in ex.items():
if isinstance(v, str) and len(v) > 0:
ctx = v
break
for k, v in ex.items():
if isinstance(v, (list, tuple)) and len(v) > 1:
if all(isinstance(x, (str, dict)) for x in v):
opts = _extract_choices_text(v)
break
lab = ex.get("answer") or ex.get("label") or ex.get("answerKey") or ex.get("correct_answer")
if opts is None:
opts = [""]
if num_choices is None:
num_choices = len(opts)
else:
if len(opts) != num_choices:
if len(opts) < num_choices:
opts = opts + [""] * (num_choices - len(opts))
else:
opts = opts[:num_choices]
lab_idx = _map_answer_to_index(lab, opts)
contexts.append(ctx)
options.append(opts)
labels.append(int(lab_idx) if lab_idx is not None else -1)
if num_choices is None:
return torch.zeros((0, 1, 1), dtype=torch.long), torch.zeros((0, 1, 1), dtype=torch.long), torch.tensor([], dtype=torch.long)
input_ids_rows = []
attention_rows = []
pad_id = getattr(tokenizer, "pad_token_id", None)
if pad_id is None:
try:
pad_id = tokenizer.token_to_id("[PAD]")
except Exception:
pad_id = 0
for i in range(0, len(contexts), batch_tokenize_size):
chunk_ctx = contexts[i:i+batch_tokenize_size]
chunk_opts = options[i:i+batch_tokenize_size]
flat_pairs = []
for c, opts in zip(chunk_ctx, chunk_opts):
for o in opts:
flat_pairs.append((c if c is not None else "", o if o is not None else ""))
enc = tokenizer(flat_pairs, truncation=True, padding=False, max_length=max_length)
ids_flat = enc.get("input_ids")
masks_flat = enc.get("attention_mask") or enc.get("mask") or enc.get("masks")
if isinstance(ids_flat, torch.Tensor):
ids_flat = ids_flat.tolist()
if isinstance(masks_flat, torch.Tensor):
masks_flat = masks_flat.tolist()
per_example = []
per_mask_example = []
idx = 0
for _ in range(len(chunk_ctx)):
row = []
row_mask = []
for _ in range(num_choices):
if idx >= len(ids_flat):
# this will raise later if shapes inconsistent; better to raise now
raise RuntimeError(f"[tokenize_mc] tokenizer returned too few items at batch starting index {i} (got {len(ids_flat)}, needed at least {len(chunk_ctx)*num_choices})")
row.append(ids_flat[idx])
row_mask.append(masks_flat[idx])
idx += 1
per_example.append(row)
per_mask_example.append(row_mask)
input_ids_rows.extend(per_example)
attention_rows.extend(per_mask_example)
max_len = max(len(seq) for row in input_ids_rows for seq in row) if input_ids_rows else 1
input_ids_padded = [
[ seq + [pad_id] * (max_len - len(seq)) for seq in row ]
for row in input_ids_rows
]
attention_padded = [
[ mask + [0] * (max_len - len(mask)) for mask in row ]
for row in attention_rows
]
input_ids_t = torch.tensor(input_ids_padded, dtype=torch.long) # (N, C, L)
attention_t = torch.tensor(attention_padded, dtype=torch.long)
labels_t = torch.tensor(labels, dtype=torch.long)
return input_ids_t, attention_t, labels_t
# ---------------------------------------------------------------------
# Postprocess preds helper already defined above (_postprocess_predictions)
# ---------------------------------------------------------------------
# ---------------------------------------------------------------------
# Training for GLUE tasks (full fine-tune)
# ---------------------------------------------------------------------
def train_full_finetune(
task: str,
tokenizer,
model,
raw_train,
raw_val,
device: str = "cuda",
epochs: int = 3,
batch_size: int = 32,
lr: float = 2e-5,
weight_decay: float = 0.01,
warmup_steps: int = 100,
max_length: int = 128,
grad_accum_steps: int = 1,
out_checkpoint_dir: Optional[str] = None,
seed: Optional[int] = None,
):
cfg_task = GLUE_TASKS[task]
device_t = torch.device(device if torch.cuda.is_available() else "cpu")
if seed is not None:
_set_all_seeds(int(seed))
hidden_size = None
if hasattr(model, "config") and hasattr(model.config, "hidden_size"):
try:
hidden_size = int(model.config.hidden_size)
except Exception:
hidden_size = None
model, wrapped_flag = make_wrapped_model_if_needed(
model, hidden_size, cfg_task["num_labels"], force_num_labels=cfg_task["num_labels"]
)
model.to(device_t)
train_ids, train_mask, train_labels = _tokenize_hf_split_to_tensors(task, tokenizer, raw_train, cfg_task, max_length=max_length)
val_ids, val_mask, val_labels = _tokenize_hf_split_to_tensors(task, tokenizer, raw_val, cfg_task, max_length=max_length)
train_ds = TensorDataset(train_ids, train_mask, train_labels)
val_ds = TensorDataset(val_ids, val_mask, val_labels)
g = torch.Generator()
if seed is not None:
g.manual_seed(int(seed))
train_loader = DataLoader(train_ds, batch_size=batch_size, shuffle=True, pin_memory=True, generator=g)
val_loader = DataLoader(val_ds, batch_size=max(64, batch_size), shuffle=False, pin_memory=True)
optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=weight_decay)
total_steps = max(1, (len(train_loader) // max(1, grad_accum_steps)) * epochs)
try:
from transformers import get_cosine_schedule_with_warmup
scheduler = get_cosine_schedule_with_warmup(optimizer, num_warmup_steps=warmup_steps, num_training_steps=total_steps)
except Exception:
scheduler = None
loss_fn = torch.nn.CrossEntropyLoss() if cfg_task["type"] == "classification" else torch.nn.MSELoss()
best_metric_res: Dict[str, Any] = {}
best_score: Optional[float] = None
best_epoch = -1
model.train()
for epoch in range(epochs):
for step, batch in enumerate(tqdm(train_loader, desc=f"Train {task} epoch {epoch+1} (lr={lr:g})")):
ids_b, mask_b, labs_b = batch
ids_b = ids_b.to(device_t)
mask_b = mask_b.to(device_t)
labs_b = labs_b.to(device_t)
out = model(input_ids=ids_b, attention_mask=mask_b, labels=None)
logits = getattr(out, "logits", None)
if logits is None:
if isinstance(out, (tuple, list)):
logits = out[0]
else:
raise RuntimeError("Model did not return logits during finetune")
if cfg_task["type"] == "classification":
loss = loss_fn(logits, labs_b.long())
else:
if logits.ndim == 2 and logits.shape[1] == 1:
preds = logits.squeeze(1)
elif logits.ndim == 2:
preds = logits.mean(dim=1)
else:
preds = logits
loss = loss_fn(preds, labs_b.float())
loss = loss / max(1, grad_accum_steps)
loss.backward()
if (step + 1) % max(1, grad_accum_steps) == 0:
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
if scheduler is not None:
scheduler.step()
optimizer.zero_grad()
# validation
model.eval()
tot_val_loss = 0.0
all_logits = []
all_labels = []
with torch.no_grad():
for ids_b, mask_b, labs_b in tqdm(val_loader, desc=f"Validate {task} epoch {epoch+1}", leave=False):
ids_b = ids_b.to(device_t)
mask_b = mask_b.to(device_t)
labs_b = labs_b.to(device_t)
out = model(input_ids=ids_b, attention_mask=mask_b, labels=None)
logits = getattr(out, "logits", None)
if logits is None:
if isinstance(out, (tuple, list)):
logits = out[0]
else:
raise RuntimeError("Model did not return logits during validation")
if cfg_task["type"] == "classification":
l = loss_fn(logits, labs_b.long())
else:
if logits.ndim == 2 and logits.shape[1] == 1:
preds = logits.squeeze(1)
elif logits.ndim == 2:
preds = logits.mean(dim=1)
else:
preds = logits
l = loss_fn(preds, labs_b.float())
tot_val_loss += l.item() * ids_b.size(0)
all_logits.append(logits.detach().cpu().numpy())
all_labels.append(labs_b.detach().cpu().numpy())
model.train()
all_logits = np.concatenate(all_logits, axis=0) if all_logits else np.zeros((0, cfg_task["num_labels"]))
all_labels = np.concatenate(all_labels, axis=0) if all_labels else np.zeros((0,))
preds = _postprocess_predictions(task, all_logits, cfg_task)
metric = evaluate.load("glue", cfg_task["hf_name"])
try:
metric_res = metric.compute(predictions=preds.tolist(), references=all_labels.tolist())
except Exception:
try:
metric_res = metric.compute(predictions=preds, references=all_labels)
except Exception as e:
metric_res = {"error": str(e)}
avg_val_loss = tot_val_loss / len(val_ds) if len(val_ds) > 0 else float("nan")
score = _metric_to_scalar(task, metric_res, fallback_val_loss=avg_val_loss)
print(f"[FT] {task} epoch {epoch+1} lr={lr:g} val_loss={avg_val_loss:.6f} metric={metric_res} score={score:.6f}")
if best_score is None or float(score) > float(best_score):
best_score = float(score)
best_metric_res = metric_res
best_epoch = epoch + 1
if out_checkpoint_dir:
outp = Path(out_checkpoint_dir)
outp.mkdir(parents=True, exist_ok=True)
best_fname = outp / "best_finetuned.pt"
try:
sd = unwrap_model(model).state_dict()
except Exception:
sd = model.state_dict()
torch.save(sd, str(best_fname))
print(f"[FT] Saved best checkpoint (epoch {best_epoch}) to: {best_fname}")
# save final checkpoint
if out_checkpoint_dir:
outp = Path(out_checkpoint_dir)
outp.mkdir(parents=True, exist_ok=True)
fname = outp / "finetuned.pt"
try:
sd = unwrap_model(model).state_dict()
except Exception:
sd = model.state_dict()
torch.save(sd, str(fname))
print(f"[FT] Saved finetuned model to: {fname}")
meta = {
"task": task,
"lr": float(lr),
"seed": int(seed) if seed is not None else None,
"epochs": int(epochs),
"batch_size": int(batch_size),
"grad_accum_steps": int(grad_accum_steps),
"warmup_steps": int(warmup_steps),
"weight_decay": float(weight_decay),
"max_length": int(max_length),
"wrapped_flag": bool(wrapped_flag),
"best_epoch": int(best_epoch),
"best_score": float(best_score) if best_score is not None else None,
"best_metrics": best_metric_res,
}
_json_dump(meta, outp / "run_meta.json")
print(
f"[FT] Best validation for task '{task}' "
f"(lr={lr:g}, seed={seed}): epoch={best_epoch}, "
f"score={best_score}, metrics={best_metric_res}"
)
return model, best_metric_res, float(best_score) if best_score is not None else -1e9, best_epoch
# ---------------------------------------------------------------------
# Load finetuned checkpoint for eval (wrapped/unwrapped)
# ---------------------------------------------------------------------
def _load_finetuned_checkpoint_for_task(task: str, base_model, checkpoint_path: str):
if not checkpoint_path or not os.path.exists(checkpoint_path):
return base_model
cfg = GLUE_TASKS.get(task) or EXTRA_TASKS.get(task)
sd = torch.load(checkpoint_path, map_location="cpu")
keys = list(sd.keys()) if isinstance(sd, dict) else []
looks_wrapped = any(k.startswith("base.") for k in keys) or any(k.startswith("classifier.") for k in keys)
hidden_size = None
if hasattr(base_model, "config") and hasattr(base_model.config, "hidden_size"):
try:
hidden_size = int(base_model.config.hidden_size)
except Exception:
hidden_size = None
if looks_wrapped and cfg is not None:
wrapped_model, _ = make_wrapped_model_if_needed(
base_model,
hidden_size,
cfg["num_labels"],
force_num_labels=cfg["num_labels"],
)
try:
unwrap_model(wrapped_model).load_state_dict(sd, strict=False)
except Exception:
wrapped_model.load_state_dict(sd, strict=False)
return wrapped_model
try:
unwrap_model(base_model).load_state_dict(sd, strict=False)
except Exception:
base_model.load_state_dict(sd, strict=False)
return base_model
# ---------------------------------------------------------------------
# Evaluation only for EXTRA_TASKS (no training)
# ---------------------------------------------------------------------
def run_extra_task(
task: str,
tokenizer,
model,
device: str = "cuda",
batch_size: int = 32,
max_length: int = 128,
output_dir: Optional[str] = None,
expect_num_choices: Optional[int] = None,
):
"""
Evaluation-only runner for EXTRA_TASKS.
Uses the SAME MC flatten->reshape logic as train_full_finetune_extra.
Reports accuracy on validation split (lm-eval style).
"""
cfg = EXTRA_TASKS[task]
device_t = torch.device(device if torch.cuda.is_available() else "cpu")
model.to(device_t)
model.eval()
# load dataset
if "hf_config" in cfg:
ds = load_dataset(cfg["hf_path"], cfg["hf_config"])
else:
ds = load_dataset(cfg["hf_path"])
raw_eval = ds.get("validation") or ds.get("test")
if raw_eval is None:
raise RuntimeError(f"[run_extra_task] No validation/test split for task={task}")
is_mc = cfg["format"] == "mc"
if is_mc:
ids, masks, labels = _tokenize_generic_mc_split_to_tensors(
task,
tokenizer,
raw_eval,
max_length=max_length,
expect_num_choices=expect_num_choices,
)
if ids.ndim != 3:
raise RuntimeError(f"[run_extra_task][{task}] ids must be (N,C,L)")
ds_eval = TensorDataset(ids, masks, labels)
else:
ids, masks, labels = _tokenize_hf_split_to_tensors(
task,
tokenizer,
raw_eval,
cfg,
max_length=max_length,
)
ds_eval = TensorDataset(ids, masks, labels)
loader = DataLoader(ds_eval, batch_size=batch_size, shuffle=False)
all_preds = []
all_refs = []
with torch.no_grad():
for batch in tqdm(loader, desc=f"[ExtraEval] {task}"):
if is_mc:
ids_b, mask_b, labs_b = batch
ids_b = ids_b.to(device_t)
mask_b = mask_b.to(device_t)
labs_b = labs_b.to(device_t)
B, C, L = ids_b.shape
flat_ids = ids_b.view(B * C, L)
flat_mask = mask_b.view(B * C, L)
try:
out = model(input_ids=flat_ids, attention_mask=flat_mask, labels=None)
except TypeError:
out = model(flat_ids)
flat_logits = getattr(out, "logits", None)
if flat_logits is None and isinstance(out, (tuple, list)):
flat_logits = out[0]
if flat_logits is None:
raise RuntimeError("Model did not return logits during MC eval")
# reduce to (B, C)
if flat_logits.ndim == 1:
choice_scores = flat_logits.view(B, C)
elif flat_logits.ndim == 2:
D = flat_logits.shape[1]
if D == 1:
choice_scores = flat_logits.view(B, C).squeeze(-1)
else:
if not hasattr(model, "_mc_projector") or model._mc_projector.weight.shape[1] != D:
model._mc_projector = torch.nn.Linear(D, 1).to(device_t)
proj = model._mc_projector(flat_logits)
choice_scores = proj.view(B, C).squeeze(-1)
else:
choice_scores = flat_logits.mean(dim=tuple(range(1, flat_logits.ndim))).view(B, C)
preds = torch.argmax(choice_scores, dim=1)
else:
ids_b, mask_b, labs_b = batch
ids_b = ids_b.to(device_t)
mask_b = mask_b.to(device_t)
out = model(input_ids=ids_b, attention_mask=mask_b, labels=None)
logits = getattr(out, "logits", None)
if logits is None and isinstance(out, (tuple, list)):
logits = out[0]
preds = torch.argmax(logits, dim=-1)
all_preds.append(preds.cpu().numpy())
all_refs.append(labs_b.cpu().numpy())
all_preds = np.concatenate(all_preds, axis=0)
all_refs = np.concatenate(all_refs, axis=0)
acc = float((all_preds == all_refs).mean())
metric_res = {"accuracy": acc}
print(f"[ExtraEval] {task} accuracy={acc:.4f}")
if output_dir:
Path(output_dir).mkdir(parents=True, exist_ok=True)
with open(Path(output_dir) / "eval_metrics.json", "w") as f:
json.dump(metric_res, f, indent=2)
return metric_res
# ---------------------------------------------------------------------
# Grid-runner: LR sweep + restarts + checkpointing + eval
# ---------------------------------------------------------------------
def run_glue_benchmark(
config,
tokenizer,
model,
checkpointing: Optional[Checkpointing] = None,
out_dir: str = "glue_outputs_grid",
):
tasks = getattr(config, "glue_tasks", None)
if tasks is None:
tasks = list(ALL_TASKS.keys())
tasks = ['openbookqa','arc_easy','arc_challenge']
# ['hellaswag','openbookqa','arc_easy','arc_challenge']
# tasks = tasks[::-1]
# tasks = ["piqa", "boolq", "winogrande"]
# "hellaswag": "accuracy",
# "openbookqa": "accuracy",
# "arc_easy": "accuracy",
# "arc_challenge": "accuracy",
batch_size = getattr(config, "batch_size", 64)
max_length = getattr(config, "max_length", 128)
device = getattr(config, "device", "cuda")
auto_train = getattr(config, "auto_train", True)
train_epochs = getattr(config, "train_epochs", 3)
train_epochs_per_task = getattr(config, "train_epochs_per_task", {}) or {
t: train_epochs for t in tasks
}
train_batch_size = getattr(config, "train_batch_size", 32)
train_warmup_steps = getattr(config, "train_warmup_steps", 100)
train_weight_decay = getattr(config, "train_weight_decay", 0.01)
train_grad_accum_steps = getattr(config, "train_grad_accum_steps", 1)
lr_candidates = getattr(config, "bert_lr_candidates", BERT_LR_CANDIDATES)
random_restarts_small = int(getattr(config, "random_restarts_small", 1))
base_seed = int(getattr(config, "base_seed", 543211))
out_dir = Path(out_dir)
out_dir.mkdir(parents=True, exist_ok=True)
checkpoint_out_root = out_dir / "checkpoints"
checkpoint_out_root.mkdir(parents=True, exist_ok=True)
rows = []
# save original model state
try:
original_state = unwrap_model(model).state_dict()
except Exception:
original_state = model.state_dict()
def _reset_model_to_original():
try:
unwrap_model(model).load_state_dict(original_state, strict=False)
except Exception:
model.load_state_dict(original_state, strict=False)
for task in tasks:
print(f"\n==== Grid-running task: {task} ====")
epochs_this_task = int(train_epochs_per_task.get(task, train_epochs))
# load dataset
if task in EXTRA_TASKS:
cfg = EXTRA_TASKS[task]
if "hf_config" in cfg:
ds = load_dataset(cfg["hf_path"], cfg["hf_config"])
else:
ds = load_dataset(cfg["hf_path"])
train_raw = ds.get("train")
val_raw = ds.get("validation") or ds.get("test")
else:
hf = load_dataset("glue", GLUE_TASKS[task]["hf_name"])
train_raw = hf["train"]
val_raw = hf["validation_matched"] if task == "mnli" else hf["validation"]
task_ckpt_root = checkpoint_out_root / task
task_ckpt_root.mkdir(parents=True, exist_ok=True)
all_run_records = []
best_run = {"score": None, "lr": None, "restart": None, "seed": None, "best_ckpt_path": None}
small_flag = (task in SMALL_TASKS_RANDOM_RESTARTS) or (task in SMALL_TASKS_RANDOM_RESTARTS_EXTRA)
restarts_per_lr = random_restarts_small if small_flag else 1
if auto_train:
for lr in lr_candidates:
for restart_idx in range(restarts_per_lr):
seed = (
base_seed
+ (abs(hash(task)) % 10000) * 1000
+ restart_idx * 10
+ (int(round(lr * 1e7)) % 1000)
)
print(f"[SWEEP] task={task} lr={lr:g} restart={restart_idx} seed={seed}")
_reset_model_to_original()
run_dir = task_ckpt_root / f"lr_{lr:g}" / f"restart_{restart_idx}"
run_dir.mkdir(parents=True, exist_ok=True)
if checkpointing is not None:
try:
checkpointing.load_model_states("recent")
except Exception:
pass
if task in EXTRA_TASKS:
det = detect_num_choices(val_raw, task)
expect_num = det.get("detected_num_choices")
model, metric_res = train_full_finetune_extra(
task=task,
tokenizer=tokenizer,
model=model,
raw_train=train_raw,
raw_val=val_raw,
device=device,
epochs=epochs_this_task,
batch_size=train_batch_size,
lr=float(lr),
weight_decay=train_weight_decay,
warmup_steps=train_warmup_steps,
max_length=max_length,
grad_accum_steps=train_grad_accum_steps,
out_checkpoint_dir=str(run_dir),
expect_num_choices=expect_num,
)
sc = _metric_to_scalar(task, metric_res)
best_epoch = None
else:
model, metric_res, sc, best_epoch = train_full_finetune(
task=task,
tokenizer=tokenizer,
model=model,
raw_train=train_raw,
raw_val=val_raw,
device=device,
epochs=epochs_this_task,
batch_size=train_batch_size,
lr=float(lr),
weight_decay=train_weight_decay,
warmup_steps=train_warmup_steps,
max_length=max_length,
grad_accum_steps=train_grad_accum_steps,
out_checkpoint_dir=str(run_dir),
seed=seed,
)
rec = {
"task": task,
"lr": float(lr),
"restart": int(restart_idx),
"seed": int(seed),
"dev_score": float(sc) if sc is not None else None,
"run_dir": str(run_dir),
"best_ckpt_path": str(run_dir / "best_finetuned.pt"),
}
all_run_records.append(rec)
if best_run["score"] is None or (sc is not None and sc > best_run["score"]):
best_run.update(
{
"score": float(sc),
"lr": float(lr),
"restart": int(restart_idx),
"seed": int(seed),
"best_ckpt_path": rec["best_ckpt_path"],
}
)
pd.DataFrame(all_run_records).to_csv(task_ckpt_root / "all_runs.csv", index=False)
best_overall_dir = task_ckpt_root / "best_overall"
best_overall_dir.mkdir(parents=True, exist_ok=True)
if best_run["best_ckpt_path"] and os.path.exists(best_run["best_ckpt_path"]):
shutil.copy2(best_run["best_ckpt_path"], best_overall_dir / "best_finetuned.pt")
_json_dump(best_run, best_overall_dir / "best_meta.json")
_reset_model_to_original()
model_for_eval = _load_finetuned_checkpoint_for_task(
task, model, str(best_overall_dir / "best_finetuned.pt")
)
else:
model_for_eval = model
task_out_dir = out_dir / task
task_out_dir.mkdir(parents=True, exist_ok=True)
if task in EXTRA_TASKS:
det = detect_num_choices(val_raw, task)
expect_num = det.get("detected_num_choices")
metric_res = run_extra_task(
task=task,
tokenizer=tokenizer,
model=model_for_eval,
device=device,
batch_size=batch_size,
max_length=max_length,
output_dir=str(task_out_dir),
expect_num_choices=expect_num,
)
rows.append({"task": task, "split": "validation", "metrics": json.dumps(metric_res)})
else:
res = run_glue_task(
task=task,
tokenizer=tokenizer,
model=model_for_eval,
device=device,
batch_size=batch_size,
max_length=max_length,
output_dir=str(task_out_dir),
compute_errorbar=True,
)
for split, pack in res.items():
rows.append(
{
"task": task,
"split": split,
"metrics": json.dumps(pack["metrics"]),
"errorbar": json.dumps(pack["errorbar"]),
}
)
summary_csv = out_dir / "glue_summary.csv"
pd.DataFrame(rows).to_csv(summary_csv, index=False)
print(f"[Grid] Summary saved to: {summary_csv}")
return pd.DataFrame(rows)
# ---------------------------------------------------------------------
# CLI
# ---------------------------------------------------------------------
if __name__ == "__main__":
print(
"This module is intended to be imported and used by providing tokenizer/model objects.\n"
"Example:\n"
" tok = AutoTokenizer.from_pretrained('bert-base-uncased')\n"
" model = AutoModelForSequenceClassification.from_pretrained('bert-base-uncased')\n"
" run_glue_benchmark(cfg, tok, model)\n"
)