| |
| """ |
| 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 |
|
|
| |
| try: |
| from lmr.checkpointing import Checkpointing |
| from lmr.ddp import unwrap_model |
| except Exception: |
| |
| Checkpointing = None |
|
|
| def unwrap_model(m): |
| return m |
|
|
| |
| |
| |
| 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": { |
| "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", |
| }, |
| |
| "hellaswag": { |
| "type": "multiple_choice", |
| "num_labels": 4, |
| "hf_path": "hellaswag", |
| "format": "mc", |
| }, |
| "openbookqa": { |
| "type": "multiple_choice", |
| "num_labels": 4, |
| "hf_path": "openbookqa", |
| "format": "mc", |
| }, |
| |
| "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_TASKS_RANDOM_RESTARTS = {"cola", "mrpc", "rte", "stsb"} |
| SMALL_TASKS_RANDOM_RESTARTS_EXTRA = set({"piqa", "boolq", "winogrande", "hellaswag"}) |
|
|
| 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", |
| } |
|
|
| |
| |
| |
| 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 |
|
|
| |
| |
| |
| 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 |
|
|
| |
| |
| |
| 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)) |
|
|
| |
| 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 |
|
|
| |
| 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 |
|
|
| |
| 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} |
|
|
| |
| |
| |
| 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_np.ndim == 1: |
| if ttype == "classification": |
| preds = (logits_np > 0.5).astype(int) |
| else: |
| preds = logits_np.astype(float) |
| return preds |
|
|
| |
| 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_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() |
|
|
| |
| |
| |
| 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 |
|
|
| |
| |
| |
| 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 |
|
|
| |
| |
| |
| |
|
|
| 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 |
|
|
| |
| 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 |
|
|
| |
| 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() |
|
|
| |
| if isinstance(v, dict): |
| return v |
|
|
| |
| if isinstance(v, (list, tuple)) and len(v) == 1: |
| return _maybe_unwrap(v[0]) |
|
|
| |
| 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)): |
| |
| 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): |
| |
| 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) |
|
|
| |
| keys = {k.lower(): v for k, v in ex.items()} |
|
|
| |
| for k in ("ctx", "context"): |
| if k in keys and keys[k]: |
| return _to_text(keys[k]).strip() |
|
|
| |
| 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() |
|
|
| |
| 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() |
|
|
| |
| 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() |
|
|
| |
| for k in ("passage", "article", "story", "sentence", "paragraph"): |
| if k in keys and keys[k]: |
| return _to_text(keys[k]).strip() |
|
|
| |
| 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 |
|
|
| |
| 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 |
| |
| 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 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 |
|
|
| |
| ex_norm = {} |
| for k, v in ex.items(): |
| try: |
| ex_norm[k] = _maybe_unwrap(v) |
| except Exception: |
| ex_norm[k] = v |
| ex = ex_norm |
|
|
| |
| ctx = _assemble_context_from_example(ex) |
|
|
| |
| 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 |
|
|
| |
| 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 |
|
|
| |
| 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 |
|
|
| |
| 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 |
|
|
| |
| 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 |
|
|
| |
| 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 |
|
|
| |
| lab = ex.get("label") or ex.get("answerKey") or ex.get("answer") |
| lab_idx = _map_label_to_int(lab, []) |
| return ctx, [], lab_idx |
|
|
|
|
| |
| |
| |
| def detect_num_choices(raw_split, task: str, max_samples: int = 200): |
| max_s = min(len(raw_split), max_samples) |
| |
| 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 |
|
|
| |
| |
| |
| 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 |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| 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]) |
|
|
| |
| 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 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) |
| 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 |
| |
| |
| |
| 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" |
|
|
| |
| |
| |
| 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) |
|
|
| |
| |
| |
| def _infer_hidden_size_from_model(m): |
| |
| 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: |
| 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 |
|
|
| |
| if is_mc: |
| |
| 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.") |
|
|
| |
| 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) |
| |
| new_head = nn.Linear(hidden_size, 1).to(device) |
| |
| 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: |
| |
| pass |
|
|
| if not replaced: |
| |
| 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)") |
|
|
| |
| |
| |
| 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"] |
|
|
| |
| 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 = 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.") |
|
|
| |
| 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 |
|
|
| |
| 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: |
| |
| 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, :] |
| else: |
| |
| pass |
|
|
| |
| |
| 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: |
| |
| |
| |
| feat = flat_logits |
| if hasattr(model, "_mc_choice_head"): |
| |
| try: |
| in_f = model._mc_choice_head.in_features |
| if feat.shape[1] == in_f: |
| scalar_flat = model._mc_choice_head(feat) |
| else: |
| |
| 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: |
| |
| scalar_flat = feat.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("[ExtraTrain] no replacement head found on model to project pooled -> scalar") |
| |
| try: |
| scalar_flat = head_found(pooled) |
| except Exception as e: |
| |
| if pooled.ndim == 2: |
| |
| 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") |
|
|
| |
| 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 Cf != expected_num_labels: |
| |
| if not hasattr(model, "_mc_choice_to_label"): |
| model._mc_choice_to_label = nn.Linear(Cf, expected_num_labels).to(device) |
| |
| 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)) |
|
|
| |
| 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 |
|
|
|
|
| |
| |
| |
| 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() |
|
|
| |
| |
| |
| 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 |
|
|
| |
| 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) |
|
|
| |
| 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 |
|
|
| |
| 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") |
|
|
| |
| 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): |
| |
| 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) |
| 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 |
|
|
| |
| |
| |
|
|
| |
| |
| |
| 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() |
|
|
| |
| 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}") |
|
|
| |
| 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 |
|
|
|
|
| |
| |
| |
| 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 |
| |
| |
| |
| 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() |
|
|
| |
| 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") |
|
|
| |
| 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 |
|
|
|
|
| |
| |
| |
| 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'] |
| |
| |
| |
| |
| |
| |
| |
| 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 = [] |
|
|
| |
| 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)) |
|
|
| |
| 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) |
|
|
|
|
| |
| |
| |
| 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" |
| ) |
|
|