FST_code / src /lmr /glue_benchmark /glue_benchmark_cls_last.py
jasonfan's picture
2026-03-19
3b2d368 verified
Raw
History Blame Contribute Delete
52.4 kB
# lmr/glue_benchmark.py
"""
Standalone GLUE benchmark runner for your project.
Modifications:
- Move tokenizer-produced CLS token (if present) to the last non-pad position
for each example (before padding), so the wrapped model can pool from the
last non-pad token when cls_token_at_end=True.
- Robust extraction of logits/pooling during training and evaluation to ensure
we're using the last-token CLS representation.
- Minor bug fixes and debug prints to verify transformations.
"""
import os
import json
import re
from pathlib import Path
from typing import Optional, List, Tuple
from typing import Optional
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)
from lmr.checkpointing import Checkpointing
from lmr.ddp import unwrap_model
# ---------------------------------------------------------------------
# GLUE 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"},
}
# ---------------------------------------------------------------------
# Task-aware example field extraction (fixes empty-text issues)
# ---------------------------------------------------------------------
def _get_text_pair_from_example(task: str, ex: dict):
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 _move_cls_to_end_before_padding_list(input_ids_list: List[List[int]], tokenizer, max_length: Optional[int] = None, debug: bool = False):
"""
For each example (list of token ids, not padded), move the first occurrence
of cls_token_id to the end (so it becomes the last non-pad token). If max_length
is provided and the resulting length would exceed it, truncate content while
ensuring CLS is the final token.
"""
cls_id = getattr(tokenizer, "cls_token_id", None)
if cls_id is None:
if debug:
print("[move_cls] tokenizer has no cls_token_id; skipping move.")
return input_ids_list
out = []
for idx, ids in enumerate(input_ids_list):
if not isinstance(ids, list):
# tolerant: convert tensor to list if needed
try:
ids = list(ids)
except Exception:
out.append(ids)
continue
if cls_id not in ids:
# no CLS in this sequence, possibly truncate
if max_length is not None and len(ids) > max_length:
if debug:
print(f"[move_cls] example {idx}: no CLS, truncating to max_length {max_length}")
out.append(ids[:max_length])
else:
out.append(ids)
continue
# remove the first CLS occurrence
new_ids = []
removed = False
for tid in ids:
if tid == cls_id and not removed:
removed = True
continue
new_ids.append(tid)
# ensure final token is CLS, respecting max_length
if max_length is not None:
if len(new_ids) + 1 > max_length:
if debug:
print(f"[move_cls] example {idx}: truncated to keep CLS at end and respect max_length {max_length}")
new_ids = new_ids[: max_length - 1]
new_ids.append(cls_id)
if debug:
orig_pos = ids.index(cls_id)
print(f"[move_cls] example {idx}: moved CLS (id={cls_id}) from pos {orig_pos} to pos {len(new_ids)-1} (len after move {len(new_ids)})")
out.append(new_ids)
return out
def _batch_tokenize(tokenizer, texts: List[Tuple[Optional[str], Optional[str]]], max_length: int = 128, move_cls_to_end: bool = False, debug: bool = False):
"""
Batch-tokenize texts. If move_cls_to_end is True, move the tokenizer's CLS
token to each example's final non-pad position (before padding/truncation).
Returns dict with "input_ids" (list of lists) and "attention_mask" (list of lists).
"""
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, return_attention_mask=True)
ids = enc.get("input_ids", None)
masks = enc.get("attention_mask", None)
# convert tensors to lists if needed
if isinstance(ids, torch.Tensor):
ids = ids.tolist()
if isinstance(masks, torch.Tensor):
masks = masks.tolist()
# Optionally move CLS before padding
if move_cls_to_end and isinstance(ids, list):
ids = _move_cls_to_end_before_padding_list(ids, tokenizer, max_length=max_length, debug=debug)
masks = [ [1]*len(x) for x in ids ]
return {"input_ids": ids, "attention_mask": masks}
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)
ids = enc.get("input_ids", None)
masks = enc.get("attention_mask", None) or enc.get("mask", None)
if isinstance(ids, torch.Tensor):
ids = ids.tolist()
if isinstance(masks, torch.Tensor):
masks = masks.tolist()
if move_cls_to_end and isinstance(ids, list):
ids = _move_cls_to_end_before_padding_list(ids, tokenizer, max_length=max_length, debug=debug)
masks = [ [1]*len(x) for x in ids ]
return {"input_ids": ids, "attention_mask": masks}
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), return_attention_mask=True)
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), return_attention_mask=True)
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}")
if move_cls_to_end:
input_ids_list = _move_cls_to_end_before_padding_list(input_ids_list, tokenizer, max_length=max_length, debug=debug)
attention_mask_list = [ [1]*len(x) for x in input_ids_list ]
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):
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()
# ---------------------------------------------------------------------
# Model wrapping helper (robust) - inline pooling, retry for hidden_states
# ---------------------------------------------------------------------
def make_wrapped_model_if_needed(model, hidden_size: Optional[int], num_labels: int, force_num_labels: Optional[int] = None, cls_token_at_end: bool = True):
import torch.nn as nn
import re
base_model = model
def _detect_head_dim(m):
try:
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
except Exception:
pass
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:
try:
cand = getattr(base_model, "config", None)
if cand is not None and hasattr(cand, "hidden_size"):
inferred_hidden = int(cand.hidden_size)
except Exception:
inferred_hidden = None
if inferred_hidden is None:
try:
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 cand > 1 and cand < 1000000:
inferred_hidden = int(cand)
break
except Exception:
inferred_hidden = None
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, cls_token_at_end=True):
super().__init__()
self.base = base
self.classifier = nn.Linear(hidden_size, num_labels)
self.logits_projector = None
self.cls_token_at_end = bool(cls_token_at_end)
def pool_from_last_hidden(self, last_hidden, attention_mask):
# last_hidden: (B, T, H)
if not self.cls_token_at_end:
# CLS at start
return last_hidden[:, 0, :]
# CLS at end (we expect CLS token to be placed at last non-pad position)
if attention_mask is None:
return last_hidden[:, -1, :]
if attention_mask.dtype != torch.long and attention_mask.dtype != torch.int:
attention_mask = attention_mask.long()
lengths = attention_mask.sum(dim=1)
lengths = torch.clamp(lengths, min=1)
idx = (lengths - 1).unsqueeze(1).to(torch.long) # (B,1)
batch_size, seq_len, hidden = last_hidden.size()
idx_exp = idx.unsqueeze(-1).expand(-1, -1, hidden) # (B,1,H)
pooled = last_hidden.gather(1, idx_exp).squeeze(1)
return pooled
def forward(self, input_ids=None, attention_mask=None, labels=None, **kwargs):
# Try calling base normally first
try:
out = self.base(input_ids=input_ids, attention_mask=attention_mask, **kwargs)
except TypeError:
# Some bases accept positional only
out = self.base(input_ids)
# 1) If we have last_hidden_state -> pool from it
last_hidden = getattr(out, "last_hidden_state", None)
if last_hidden is not None:
pooled = self.pool_from_last_hidden(last_hidden, attention_mask)
logits = self.classifier(pooled)
return type("Out", (), {"logits": logits, "loss": None, "last_hidden_state": last_hidden})
# 2) If out[0] is a (B,T,H) tensor -> pool from it
if isinstance(out, (tuple, list)) and len(out) > 0:
cand = out[0]
if torch.is_tensor(cand):
if cand.ndim == 3:
pooled = self.pool_from_last_hidden(cand, attention_mask)
logits = self.classifier(pooled)
return type("Out", (), {"logits": logits, "loss": None, "last_hidden_state": cand})
# If cand is (B, C) assume already logits
if cand.ndim == 2 and cand.shape[1] == num_labels:
return type("Out", (), {"logits": cand, "loss": None})
# 3) If hidden_states attr exists, use last layer
hidden_states = getattr(out, "hidden_states", None)
if hidden_states is not None:
if isinstance(hidden_states, (list, tuple)):
last_hidden = hidden_states[-1]
else:
last_hidden = hidden_states
if torch.is_tensor(last_hidden) and last_hidden.ndim == 3:
pooled = self.pool_from_last_hidden(last_hidden, attention_mask)
logits = self.classifier(pooled)
return type("Out", (), {"logits": logits, "loss": None, "last_hidden_state": last_hidden})
# 4) Try to re-call base with output_hidden_states=True to get hidden states
try:
new_out = self.base(input_ids=input_ids, attention_mask=attention_mask, output_hidden_states=True, return_dict=True, **kwargs)
new_last_hidden = getattr(new_out, "last_hidden_state", None)
if new_last_hidden is None:
new_hidden_states = getattr(new_out, "hidden_states", None)
if new_hidden_states is not None:
if isinstance(new_hidden_states, (list, tuple)):
new_last_hidden = new_hidden_states[-1]
else:
new_last_hidden = new_hidden_states
if new_last_hidden is not None and torch.is_tensor(new_last_hidden) and new_last_hidden.ndim == 3:
pooled = self.pool_from_last_hidden(new_last_hidden, attention_mask)
logits = self.classifier(pooled)
return type("Out", (), {"logits": logits, "loss": None, "last_hidden_state": new_last_hidden})
except Exception:
pass
# 5) Fallback: if logits present and appear token-level (B, T, C) -> reduce to (B, C)
logits = getattr(out, "logits", None)
if isinstance(logits, torch.Tensor):
# token-level logits case: (B, T, C)
if logits.ndim == 3 and logits.shape[2] == num_labels:
# attempt to pick last non-pad token's logits per sample if attention_mask present
if attention_mask is not None:
if attention_mask.dtype != torch.long and attention_mask.dtype != torch.int:
am = attention_mask.long()
else:
am = attention_mask
lengths = am.sum(dim=1)
lengths = torch.clamp(lengths, min=1)
idx = (lengths - 1).unsqueeze(1).unsqueeze(-1).expand(-1, -1, logits.size(-1))
pooled_logits = logits.gather(1, idx).squeeze(1)
return type("Out", (), {"logits": pooled_logits, "loss": getattr(out, "loss", None)})
else:
pooled_logits = logits.mean(dim=1)
return type("Out", (), {"logits": pooled_logits, "loss": getattr(out, "loss", None)})
# if logits are already (B, C) and correct -> return
if logits.ndim == 2 and logits.shape[1] == num_labels:
return type("Out", (), {"logits": logits, "loss": getattr(out, "loss", None)})
# if logits are (T, C) and batch likely 1, try to collapse
if logits.ndim == 2 and input_ids is not None and input_ids.ndim == 2 and input_ids.size(0) == 1:
if attention_mask is not None:
lengths = attention_mask.sum(dim=1)
last_idx = int(max(1, lengths[0].item()) - 1)
pooled_logits = logits[last_idx:last_idx+1, :].squeeze(0).unsqueeze(0)
return type("Out", (), {"logits": pooled_logits, "loss": getattr(out, "loss", None)})
else:
pooled_logits = logits.mean(dim=0, keepdim=True)
return type("Out", (), {"logits": pooled_logits, "loss": getattr(out, "loss", None)})
raise RuntimeError("Wrapped base model did not return recognizable hidden states or logits")
return _WrappedModel(base_model, inferred_hidden, num_labels, cls_token_at_end=cls_token_at_end), True
# ---------------------------------------------------------------------
# Tokenize HF dataset split into tensors (used for finetune)
# ---------------------------------------------------------------------
def _tokenize_hf_split_to_tensors(task: str, tokenizer, raw_split, cfg_task, max_length=128, batch_tokenize_size=512, debug: bool = False):
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):
batch_texts = texts[i:i+batch_tokenize_size]
# IMPORTANT: we set move_cls_to_end=True so CLS is moved to end before padding
enc = _batch_tokenize(tokenizer, batch_texts, max_length=max_length, move_cls_to_end=True, debug=debug)
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: classification -> long, regression -> float
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
# ---------------------------------------------------------------------
# Train full fine-tune (entire model) for a GLUE task
# ---------------------------------------------------------------------
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,
cls_token_at_end: bool = True, debug: bool = False):
cfg_task = GLUE_TASKS[task]
device = torch.device(device if torch.cuda.is_available() else "cpu")
PREFERRED_METRIC_KEY = {
"cola": "matthews_correlation",
"sst2": "accuracy",
"mrpc": "accuracy",
"stsb": "pearson",
"qqp": "accuracy",
"mnli": "accuracy",
"qnli": "accuracy",
"rte": "accuracy",
"wnli": "accuracy",
}
preferred_key = PREFERRED_METRIC_KEY.get(task, None)
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
wrapped_model, wrapped_flag = make_wrapped_model_if_needed(model, hidden_size, cfg_task["num_labels"], force_num_labels=cfg_task["num_labels"], cls_token_at_end=cls_token_at_end)
model = wrapped_model
model.to(device)
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)
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)
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
# ignore_index=-100 helps if you used -100 as "no label" sentinel; safe default
loss_fn = torch.nn.CrossEntropyLoss(ignore_index=-100) if cfg_task["type"] == "classification" else torch.nn.MSELoss()
best_metric_res = {}
best_score = None
best_epoch = -1
model.train()
global_step = 0
for epoch in range(epochs):
running_loss = 0.0
for step, batch in enumerate(tqdm(train_loader, desc=f"Train {task} epoch {epoch+1}")):
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)
# Forward and robust logits extraction / pooling:
out = model(input_ids=ids_b, attention_mask=mask_b, labels=None,cls_token_at_end=True)
logits = getattr(out, "logits", None)
# If logits missing, try tuple/list first element
if logits is None and isinstance(out, (tuple, list)) and len(out) > 0:
cand = out[0]
if torch.is_tensor(cand):
if cand.ndim == 2:
logits = cand
elif cand.ndim == 3:
# token-level logits: pick last non-pad token per sample
am = mask_b
if am is None:
logits = cand.mean(dim=1)
else:
if am.dtype != torch.long and am.dtype != torch.int:
am = am.long()
lengths = am.sum(dim=1).clamp(min=1)
idx = (lengths - 1).unsqueeze(1).unsqueeze(-1).expand(-1, -1, cand.size(-1))
logits = cand.gather(1, idx).squeeze(1)
else:
logits = None
# If still no logits, but we have last_hidden_state / hidden_states, pool + classifier
if logits is None:
last_hidden = getattr(out, "last_hidden_state", None)
if last_hidden is 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 last_hidden is not None and torch.is_tensor(last_hidden) and last_hidden.ndim == 3:
try:
pooled = model.pool_from_last_hidden(last_hidden, mask_b)
logits = model.classifier(pooled)
except Exception:
# fallback pooling
am = mask_b
if am is None:
pooled = last_hidden.mean(dim=1)
else:
if am.dtype != torch.long and am.dtype != torch.int:
am = am.long()
lengths = am.sum(dim=1).clamp(min=1)
idx = (lengths - 1).unsqueeze(1).unsqueeze(-1).expand(-1, -1, last_hidden.size(-1))
pooled = last_hidden.gather(1, idx).squeeze(1)
logits = model.classifier(pooled)
if logits is None:
raise RuntimeError("Model did not return usable logits/hidden_states during finetune")
# Now compute loss
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()
global_step += 1
running_loss += loss.item() * ids_b.size(0)
# 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); mask_b = mask_b.to(device); labs_b = labs_b.to(device)
out = model(input_ids=ids_b, attention_mask=mask_b, labels=None,cls_token_at_end=True)
logits = getattr(out, "logits", None)
if logits is None and isinstance(out, (tuple, list)) and len(out) > 0:
cand = out[0]
if torch.is_tensor(cand):
if cand.ndim == 2:
logits = cand
elif cand.ndim == 3:
am = mask_b
if am is None:
logits = cand.mean(dim=1)
else:
if am.dtype != torch.long and am.dtype != torch.int:
am = am.long()
lengths = am.sum(dim=1).clamp(min=1)
idx = (lengths - 1).unsqueeze(1).unsqueeze(-1).expand(-1, -1, cand.size(-1))
logits = cand.gather(1, idx).squeeze(1)
if logits is None:
last_hidden = getattr(out, "last_hidden_state", None)
if last_hidden is 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 last_hidden is not None and torch.is_tensor(last_hidden) and last_hidden.ndim == 3:
try:
pooled = model.pool_from_last_hidden(last_hidden, mask_b)
logits = model.classifier(pooled)
except Exception:
am = mask_b
if am is None:
pooled = last_hidden.mean(dim=1)
else:
if am.dtype != torch.long and am.dtype != torch.int:
am = am.long()
lengths = am.sum(dim=1).clamp(min=1)
idx = (lengths - 1).unsqueeze(1).unsqueeze(-1).expand(-1, -1, last_hidden.size(-1))
pooled = last_hidden.gather(1, idx).squeeze(1)
logits = model.classifier(pooled)
if logits is None:
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")
print(f"[FT] {task} epoch {epoch+1} val_loss={avg_val_loss:.6f} metric={metric_res}")
# best selection
score = None
if isinstance(metric_res, dict) and metric_res:
if preferred_key is not None and preferred_key in metric_res:
try:
score = float(metric_res[preferred_key])
except Exception:
score = None
if score is None:
for k, v in metric_res.items():
try:
score = float(v)
break
except Exception:
continue
if score is None or (isinstance(metric_res, dict) and "error" in metric_res):
try:
score = -float(avg_val_loss)
except Exception:
score = float(epoch)
is_better = False
if best_score is None:
is_better = True
else:
try:
if float(score) > float(best_score):
is_better = True
except Exception:
is_better = True if float(score) != float(best_score) else False
if is_better:
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}")
print(f"[FT] Best validation epoch for task '{task}': epoch={best_epoch}, score={best_score}, metrics={best_metric_res}")
return model, best_metric_res
# ---------------------------------------------------------------------
# Main per-task runner (evaluation only)
# ---------------------------------------------------------------------
def run_glue_task(task: str, tokenizer, model, checkpointing: Optional[Checkpointing] = None,
device: str = "cuda", batch_size: int = 64, max_length: int = 128,
output_dir: str = "glue_output", cls_token_at_end: bool = True):
assert task in GLUE_TASKS, f"Unknown GLUE task: {task}"
cfg = GLUE_TASKS[task]
hf = load_dataset("glue", cfg["hf_name"])
if task == "mnli":
val_splits = ["validation_matched", "validation_mismatched"]
else:
val_splits = ["validation"]
results_by_split = {}
for split in val_splits:
raw = hf[split]
print(f"[GLUE] Task={task} split={split} samples={len(raw)}")
texts = []
labels = []
for ex in raw:
s1, s2 = _get_text_pair_from_example(task, ex)
texts.append((s1, s2))
labels.append(ex.get("label") if "label" in ex else -100)
BATCH = 512
input_ids_all = []
attention_all = []
for i in range(0, len(texts), BATCH):
batch_texts = texts[i:i+BATCH]
# Move CLS to end for evaluation as well
enc = _batch_tokenize(tokenizer, batch_texts, max_length=max_length, move_cls_to_end=True, debug=False)
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, attention_mask = _pad_and_tensorize(input_ids_all, attention_all, pad_id)
labels_t = torch.tensor(labels, dtype=torch.long if cfg["type"]=="classification" else torch.float)
ds = TensorDataset(input_ids, attention_mask, labels_t)
loader = DataLoader(ds, batch_size=batch_size, shuffle=False, pin_memory=True)
if checkpointing is not None:
try:
checkpointing.load_model_states("recent")
except Exception:
pass
device = torch.device(device if torch.cuda.is_available() else "cpu")
model.to(device)
model.eval()
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
force = None
if cfg["type"] == "regression":
force = 1
else:
try:
import torch.nn as nn
existing_dim = None
if hasattr(model, "classifier") and isinstance(getattr(model, "classifier"), nn.Linear):
existing_dim = getattr(model, "classifier").out_features
elif hasattr(model, "lm_head") and isinstance(getattr(model, "lm_head"), nn.Linear):
existing_dim = getattr(model, "lm_head").out_features
if existing_dim is not None and existing_dim != cfg["num_labels"]:
force = cfg["num_labels"]
except Exception:
force = cfg["num_labels"]
wrapped_model, wrapped = make_wrapped_model_if_needed(model, hidden_size, cfg["num_labels"], force_num_labels=force, cls_token_at_end=cls_token_at_end)
wrapped_model.to(device)
wrapped_model.eval()
all_logits = []
all_labels = []
with torch.no_grad():
for batch in tqdm(loader, desc=f"Eval {task}:{split}"):
ids_b, mask_b, labels_b = batch
ids_b = ids_b.to(device)
mask_b = mask_b.to(device)
out = wrapped_model(input_ids=ids_b, attention_mask=mask_b, labels=None,cls_token_at_end=True)
logits = getattr(out, "logits", None)
if logits is None:
if isinstance(out, (tuple, list)):
cand = out[0]
if torch.is_tensor(cand):
if cand.ndim == 3:
am = mask_b
if am is None:
logits = cand.mean(dim=1)
else:
if am.dtype != torch.long and am.dtype != torch.int:
am = am.long()
lengths = am.sum(dim=1).clamp(min=1)
idx = (lengths - 1).unsqueeze(1).unsqueeze(-1).expand(-1, -1, cand.size(-1))
logits = cand.gather(1, idx).squeeze(1)
elif cand.ndim == 2:
logits = cand
else:
raise RuntimeError("Model forward did not return logits")
logits_np = logits.detach().cpu().numpy()
all_logits.append(logits_np)
all_labels.append(labels_b.detach().cpu().numpy())
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,))
preds = _postprocess_predictions(task, all_logits, cfg)
if cfg["type"] == "classification":
preds_out = preds.astype(int).tolist()
refs_out = all_labels.astype(int).tolist()
else:
preds_out = preds.astype(float).tolist()
refs_out = all_labels.astype(float).tolist()
metric = evaluate.load("glue", cfg["hf_name"])
try:
metric_res = metric.compute(predictions=preds_out, references=refs_out)
except Exception:
try:
metric_res = metric.compute(predictions=np.array(preds_out), references=np.array(refs_out))
except Exception as e:
metric_res = {"error": str(e)}
os.makedirs(output_dir, exist_ok=True)
out_json = Path(output_dir) / f"{task}_{split}_results.json"
with open(out_json, "w", encoding="utf-8") as f:
json.dump({"task": task, "split": split, "metrics": metric_res}, f, indent=2)
csv_p = Path(output_dir) / f"{task}_{split}_preds.csv"
pd.DataFrame({"pred": preds_out, "label": refs_out}).to_csv(csv_p, index=False)
results_by_split[split] = metric_res
return results_by_split
# ---------------------------------------------------------------------
# Top-level runner: multiple tasks + (optional) auto-train
# ---------------------------------------------------------------------
def run_glue_benchmark(
config,
tokenizer,
model,
checkpointing: Optional[Checkpointing] = None,
out_dir: str = "glue_outputs",
):
tasks = getattr(
config,
"glue_tasks",
["sst2"],
)
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", {})
# default per-task epochs (only used if not passed in config.train_epochs_per_task)
default_train_epochs_per_task = {
"cola": 5,
"mrpc": 3,
"rte": 5,
"stsb": 3,
"sst2": 3,
"qqp": 3,
"qnli": 3,
"mnli": 3,
"wnli": 5,
}
train_lrs_per_task = {
"mnli": 4e-5,
"qqp": 4e-5,
"qnli": 2e-5,
"sst2": 4e-5,
"cola": 3e-5,
"mrpc": 2e-5,
"rte": 4e-5,
"stsb": 2e-5,
"wnli": 4e-5,
}
train_batch_size = getattr(config, "train_batch_size", 32)
train_lr = getattr(config, "train_lr", 2e-6)
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)
cls_token_at_end = getattr(config, "cls_token_at_end", True)
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 = []
original_state = None
try:
original_state = unwrap_model(model).state_dict()
except Exception:
try:
original_state = model.state_dict()
except Exception:
original_state = None
for task in tasks:
print(f"\n==== Running GLUE task: {task} ====")
epochs_this_task = train_epochs_per_task.get(task, train_epochs) if isinstance(train_epochs_per_task, dict) and train_epochs_per_task.get(task) is not None else default_train_epochs_per_task.get(task, train_epochs)
train_lr_per_task = train_lrs_per_task.get(task, train_lr)
print(f"[GLUE] Epochs for task '{task}': {epochs_this_task}")
hf = load_dataset("glue", GLUE_TASKS[task]["hf_name"])
train_raw = hf["train"]
if task == "mnli":
val_raw = hf["validation_matched"]
else:
val_raw = hf["validation"]
if auto_train:
print(
f"[GLUE] Auto-training enabled. "
f"Fine-tuning task '{task}' from pretrained checkpoint."
)
if original_state is not None:
try:
unwrap_model(model).load_state_dict(
original_state, strict=False
)
except Exception:
try:
model.load_state_dict(original_state, strict=False)
except Exception:
pass
if checkpointing is not None:
try:
checkpointing.load_model_states("recent")
except Exception:
pass
task_ckpt_dir = checkpoint_out_root / task
task_ckpt_dir.mkdir(parents=True, exist_ok=True)
model, metric_res = 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=train_lr_per_task,
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(task_ckpt_dir),
cls_token_at_end=cls_token_at_end,
)
print(
f"[GLUE] Finished fine-tuning for task '{task}'. "
f"Val metric: {metric_res}"
)
else:
if checkpointing is not None:
try:
checkpointing.load_model_states("recent")
except Exception:
pass
task_out_dir = out_dir / task
task_out_dir.mkdir(parents=True, exist_ok=True)
res = run_glue_task(
task=task,
tokenizer=tokenizer,
model=model,
checkpointing=None,
device=device,
batch_size=batch_size,
max_length=max_length,
output_dir=str(task_out_dir),
cls_token_at_end=cls_token_at_end
)
for split, metrics in res.items():
rows.append(
{
"task": task,
"split": split,
"epochs": epochs_this_task,
"metrics": json.dumps(metrics),
}
)
summary_csv = out_dir / "glue_summary.csv"
pd.DataFrame(rows).to_csv(summary_csv, index=False)
print(f"\n[GLUE] Summary saved to: {summary_csv}")
return pd.DataFrame(rows)
# ---------------------------------------------------------------------
# CLI for quick testing (optional)
# ---------------------------------------------------------------------
if __name__ == "__main__":
import argparse
parser = argparse.ArgumentParser()
parser.add_argument("--tasks", type=str, default="sst2", help="comma separated glue tasks")
parser.add_argument("--batch_size", type=int, default=64)
parser.add_argument("--max_length", type=int, default=128)
parser.add_argument("--device", type=str, default="cuda")
parser.add_argument("--out_dir", type=str, default="glue_outputs")
parser.add_argument("--cls_token_at_end", action="store_true", help="If set, pooled CLS is taken from last non-padded token (default False).")
args = parser.parse_args()
class C: pass
cfg = C()
cfg.glue_tasks = args.tasks.split(",")
cfg.batch_size = args.batch_size
cfg.max_length = args.max_length
cfg.device = args.device
cfg.train_epochs = 3
cfg.train_batch_size = 32
cfg.train_lr = 2e-6
cfg.train_warmup_steps = 100
cfg.train_weight_decay = 0.01
cfg.train_grad_accum_steps = 1
cfg.auto_train = False
# Use the flag value directly (default False unless --cls_token_at_end passed)
cfg.cls_token_at_end = args.cls_token_at_end
print("This module is intended to be invoked from your project's main, which provides tokenizer/model/checkpointing.")
print(f"Example usage in your main: run_glue_benchmark(config.benchmark, tokenizer, model, checkpointing, out_dir={args.out_dir})")