# 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 "" 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})")