| |
| """ |
| 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 |
|
|
| |
| from lmr.checkpointing import Checkpointing |
| from lmr.ddp import unwrap_model |
|
|
| |
| |
| |
| 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"}, |
| } |
|
|
| |
| |
| |
| 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 |
|
|
| |
| |
| |
| 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): |
| |
| try: |
| ids = list(ids) |
| except Exception: |
| out.append(ids) |
| continue |
|
|
| if cls_id not in ids: |
| |
| 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 |
|
|
| |
| new_ids = [] |
| removed = False |
| for tid in ids: |
| if tid == cls_id and not removed: |
| removed = True |
| continue |
| new_ids.append(tid) |
|
|
| |
| 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)) |
|
|
| |
| 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) |
| |
| 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: |
| 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) |
| 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 |
|
|
| |
| 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} |
|
|
| |
| |
| |
| 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() |
|
|
| |
| |
| |
| 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): |
| |
| if not self.cls_token_at_end: |
| |
| return last_hidden[:, 0, :] |
| |
| 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) |
| batch_size, seq_len, hidden = last_hidden.size() |
| idx_exp = idx.unsqueeze(-1).expand(-1, -1, hidden) |
| pooled = last_hidden.gather(1, idx_exp).squeeze(1) |
| return pooled |
|
|
| def forward(self, input_ids=None, attention_mask=None, labels=None, **kwargs): |
| |
| 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 = 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}) |
|
|
| |
| 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.ndim == 2 and cand.shape[1] == num_labels: |
| return type("Out", (), {"logits": cand, "loss": None}) |
|
|
| |
| 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}) |
|
|
| |
| 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 |
|
|
| |
| logits = getattr(out, "logits", None) |
| if isinstance(logits, torch.Tensor): |
| |
| if logits.ndim == 3 and logits.shape[2] == num_labels: |
| |
| 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.ndim == 2 and logits.shape[1] == num_labels: |
| return type("Out", (), {"logits": logits, "loss": getattr(out, "loss", None)}) |
| |
| 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 |
|
|
| |
| |
| |
| 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] |
| |
| 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_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 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 |
|
|
| |
| 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) |
|
|
| |
| 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) |
| else: |
| logits = None |
|
|
| |
| 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 usable logits/hidden_states 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() |
| global_step += 1 |
| running_loss += loss.item() * ids_b.size(0) |
|
|
| |
| 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}") |
|
|
| |
| 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 |
|
|
| |
| |
| |
| 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] |
| |
| 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 |
|
|
| |
| |
| |
| 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_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) |
|
|
|
|
| |
| |
| |
| 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 |
| |
| 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})") |
|
|