| |
| """ |
| BertFineTuneTrainer |
| |
| Features: |
| - Supports classification (single-sentence), sentence-pair (GLUE-style), and next-sentence-prediction (synthesizes pairs). |
| - Loads HF datasets via `datasets.load_dataset`. |
| - Handles tokenization, dataloader creation, optimizer, scheduler, fp16, DDP (via torch.distributed if use_ddp True). |
| - Evaluates accuracy/precision/recall/f1 and reports. |
| - Saves best checkpoint by eval metric (F1 for binary/multi-class; accuracy fallback). |
| """ |
|
|
| from pathlib import Path |
| import random |
| import os |
| import time |
| import math |
| import json |
| import shutil |
| from typing import Optional, Dict, Any |
|
|
| import numpy as np |
| import torch |
| from torch.utils.data import DataLoader |
| from torch.utils.data.distributed import DistributedSampler |
| from torch.optim import AdamW |
| from torch.cuda.amp import GradScaler, autocast |
| from transformers import ( |
| AutoTokenizer, |
| AutoConfig, |
| AutoModelForSequenceClassification, |
| BertForNextSentencePrediction, |
| get_linear_schedule_with_warmup, |
| ) |
| from datasets import load_dataset, Dataset, DatasetDict |
| from tqdm import tqdm |
| from sklearn.metrics import accuracy_score, precision_recall_fscore_support |
|
|
| |
| def log(*args, **kwargs): |
| print(time.strftime("%Y-%m-%d %H:%M:%S"), "-", *args, **kwargs) |
|
|
| class BertFineTuneTrainer: |
| def __init__(self, cfg: Dict[str, Any], device: Optional[torch.device] = None): |
| """ |
| cfg: dictionary with finetune settings (see DEFAULT_CONFIG in main.py) |
| """ |
| self.cfg = cfg |
| self.device = device or (torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu")) |
| self.output_dir = Path(cfg.get("output_dir", "./outputs/finetune")) |
| self.output_dir.mkdir(parents=True, exist_ok=True) |
|
|
| |
| seed = cfg.get("seed", 42) |
| random.seed(seed) |
| np.random.seed(seed) |
| torch.manual_seed(seed) |
| if torch.cuda.is_available(): |
| torch.cuda.manual_seed_all(seed) |
|
|
| |
| self.use_ddp = bool(cfg.get("use_ddp", False)) |
| if self.use_ddp: |
| |
| if not torch.distributed.is_initialized(): |
| init_method = os.environ.get("INIT_METHOD", "env://") |
| torch.distributed.init_process_group(backend="nccl", init_method=init_method) |
| self.rank = torch.distributed.get_rank() |
| self.world_size = torch.distributed.get_world_size() |
| torch.cuda.set_device(self.rank % torch.cuda.device_count()) |
| self.device = torch.device(f"cuda:{torch.cuda.current_device()}") |
| else: |
| self.rank = 0 |
| self.world_size = 1 |
|
|
| |
| model_name = cfg["model_name_or_path"] |
| task = cfg.get("task", "sentence_pair") |
| num_labels = cfg.get("num_labels", None) |
|
|
| print('-----------------------------123123123123123123123') |
| print(model_name,'----------------------------------------------------------') |
| log(f"[rank {self.rank}] Loading tokenizer & model: {model_name}, task={task}") |
| self.tokenizer = AutoTokenizer.from_pretrained(model_name, use_fast=True) |
| |
| if getattr(self.tokenizer, "mask_token", None) is None: |
| |
| self.tokenizer.add_special_tokens({"mask_token": "[MASK]"}) |
| |
| if task == "next_sentence_prediction": |
| self.model = BertForNextSentencePrediction.from_pretrained(model_name) |
| else: |
| |
| |
| cfg_model = AutoConfig.from_pretrained(model_name) |
| if num_labels is None and hasattr(cfg_model, "num_labels"): |
| num_labels = getattr(cfg_model, "num_labels", None) |
| if num_labels is None: |
| |
| num_labels = 2 |
| self.model = AutoModelForSequenceClassification.from_pretrained(model_name, num_labels=num_labels) |
|
|
| |
| self.model.to(self.device) |
|
|
| |
| self.fp16 = bool(cfg.get("fp16", True)) |
| self.scaler = GradScaler() if self.fp16 and torch.cuda.is_available() else None |
|
|
| |
| self.batch_size = int(cfg.get("batch_size", 16)) |
| self.eval_batch_size = int(cfg.get("eval_batch_size", max(32, self.batch_size))) |
| self.num_epochs = int(cfg.get("num_epochs", 3)) |
| self.learning_rate = float(cfg.get("lr", 2e-5)) |
| self.weight_decay = float(cfg.get("weight_decay", 0.01)) |
| self.gradient_accumulation_steps = int(cfg.get("gradient_accumulation_steps", 1)) |
| self.max_grad_norm = float(cfg.get("max_grad_norm", 1.0)) |
| self.max_length = int(cfg.get("max_length", 128)) |
| self.num_workers = int(cfg.get("num_workers", 4)) |
| self.logging_steps = int(cfg.get("logging_steps", 100)) |
| self.eval_steps = int(cfg.get("eval_steps", 500)) |
| self.save_steps = int(cfg.get("save_steps", 1000)) |
| self.warmup_steps = int(cfg.get("warmup_steps", 0)) |
| self.max_train_samples = cfg.get("max_train_samples", None) |
| self.max_eval_samples = cfg.get("max_eval_samples", None) |
| self.nsp_negatives_ratio = int(cfg.get("nsp_negatives_ratio", 1)) |
|
|
| |
| self.dataset_name = cfg.get("dataset") |
| self.dataset_config_name = cfg.get("dataset_config_name", None) |
|
|
| |
| self.best_metric = -1.0 |
| self.global_step = 0 |
| self.total_steps = 0 |
|
|
| |
| |
| |
| def _load_hf_dataset(self): |
| |
| |
| |
| |
| ds_name = self.dataset_name |
| cfg_name = self.dataset_config_name |
|
|
| if ds_name is None: |
| raise ValueError("Please specify cfg['dataset'] (Hugging Face dataset id).") |
|
|
| if "/" in ds_name and not ds_name.startswith("glue/"): |
| |
| parts = ds_name.split("/", 1) |
| ds = load_dataset(parts[0], parts[1]) |
| elif ds_name.startswith("glue/"): |
| |
| parts = ds_name.split("/", 1) |
| ds = load_dataset(parts[0], parts[1]) |
| else: |
| |
| if cfg_name: |
| ds = load_dataset(ds_name, cfg_name) |
| else: |
| ds = load_dataset(ds_name) |
|
|
| |
| if isinstance(ds, Dataset): |
| ds = DatasetDict({"train": ds}) |
| if isinstance(ds, dict) and not isinstance(ds, DatasetDict): |
| ds = DatasetDict(ds) |
|
|
| return ds |
|
|
| def _prepare_sentence_pair(self, dataset: Dataset, text1_key: str, text2_key: str): |
| |
| tokenizer = self.tokenizer |
| max_length = self.max_length |
|
|
| def fn_examples(examples): |
| texts1 = examples[text1_key] |
| texts2 = examples[text2_key] |
| |
| enc = tokenizer(texts1, texts2, truncation=True, padding="max_length", max_length=max_length) |
| |
| out = {"input_ids": enc["input_ids"], "attention_mask": enc["attention_mask"]} |
| if "label" in examples: |
| out["labels"] = examples["label"] |
| return out |
|
|
| return dataset.map(fn_examples, batched=True, remove_columns=[c for c in dataset.column_names if c not in (text1_key, text2_key, "label")], num_proc=1) |
|
|
| def _prepare_classification(self, dataset: Dataset, text_key: str): |
| tokenizer = self.tokenizer |
| max_length = self.max_length |
|
|
| def fn_examples(examples): |
| texts = examples[text_key] |
| enc = tokenizer(texts, truncation=True, padding="max_length", max_length=max_length) |
| out = {"input_ids": enc["input_ids"], "attention_mask": enc["attention_mask"]} |
| if "label" in examples: |
| out["labels"] = examples["label"] |
| return out |
|
|
| return dataset.map(fn_examples, batched=True, remove_columns=[c for c in dataset.column_names if c not in (text_key, "label")], num_proc=1) |
|
|
| def _synthesize_nsp_dataset(self, ds: Dataset): |
| """ |
| Build next-sentence pairs from a text dataset: |
| - consecutive sentences -> label=1 (is_next) |
| - random sentence pair -> label=0 |
| This is a simple heuristic synthesizer. |
| """ |
| tokenizer = self.tokenizer |
| max_length = self.max_length |
| neg_ratio = self.nsp_negatives_ratio |
|
|
| texts = [] |
| |
| for ex in ds: |
| |
| if "text" in ex: |
| t = ex["text"] |
| elif "content" in ex: |
| t = ex["content"] |
| else: |
| |
| |
| t = " ".join(str(v) for k, v in ex.items() if isinstance(v, str)) |
| if not t: |
| continue |
| |
| sents = [s.strip() for s in t.replace("\n", " ").split(". ") if s.strip()] |
| for i in range(len(sents)-1): |
| texts.append((sents[i], sents[i+1], 1)) |
| |
| n_pos = len(texts) |
| if n_pos == 0: |
| raise RuntimeError("No sentence pairs extracted for NSP. Use a dataset with 'text' or 'content' fields.") |
| n_neg = n_pos * neg_ratio |
| rng = random.Random(42) |
| all_sents = [s for pair in texts for s in pair[:2]] |
| for _ in range(n_neg): |
| a = rng.choice(all_sents) |
| b = rng.choice(all_sents) |
| texts.append((a, b, 0)) |
|
|
| |
| rows = {"sentence1": [], "sentence2": [], "label": []} |
| for a,b,l in texts: |
| rows["sentence1"].append(a) |
| rows["sentence2"].append(b) |
| rows["label"].append(int(l)) |
| nrows = len(rows["label"]) |
| ds_new = Dataset.from_dict(rows) |
| |
| def fn(examples): |
| enc = tokenizer(examples["sentence1"], examples["sentence2"], truncation=True, padding="max_length", max_length=max_length) |
| return {"input_ids": enc["input_ids"], "attention_mask": enc["attention_mask"], "labels": examples["label"]} |
| ds_new = ds_new.map(fn, batched=True, remove_columns=["sentence1", "sentence2", "label"]) |
| return ds_new |
|
|
| def _build_datasets_and_loaders(self): |
| ds = self._load_hf_dataset() |
|
|
| |
| train_key = "train" if "train" in ds else list(ds.keys())[0] |
| valid_key = "validation" if "validation" in ds else ("validation_matched" if "validation_matched" in ds else None) |
| test_key = "test" if "test" in ds else None |
|
|
| |
| if self.max_train_samples: |
| ds[train_key] = ds[train_key].select(range(min(len(ds[train_key]), int(self.max_train_samples)))) |
| if valid_key and self.max_eval_samples: |
| ds[valid_key] = ds[valid_key].select(range(min(len(ds[valid_key]), int(self.max_eval_samples)))) |
|
|
| task = self.cfg.get("task", "sentence_pair") |
| |
| |
| train_ds = ds[train_key] |
| valid_ds = ds[valid_key] if valid_key else None |
|
|
| if task == "sentence_pair": |
| |
| candidates = [("sentence1","sentence2"), ("premise","hypothesis"), ("text_a","text_b"), ("sentence_a","sentence_b"), ("question","sentence")] |
| found = None |
| for a,b in candidates: |
| if a in train_ds.column_names and b in train_ds.column_names: |
| found = (a,b); break |
| if found is None: |
| |
| a = "sentence1" if "sentence1" in train_ds.column_names else None |
| b = "sentence2" if "sentence2" in train_ds.column_names else None |
| if a is None or b is None: |
| raise RuntimeError(f"Could not find sentence-pair fields in dataset columns: {train_ds.column_names}") |
| found = (a,b) |
| text1_key, text2_key = found |
| log(f"[rank {self.rank}] Using fields {text1_key}/{text2_key} for sentence_pair task.") |
| |
| train_tok = self._prepare_sentence_pair(train_ds, text1_key, text2_key) |
| valid_tok = self._prepare_sentence_pair(valid_ds, text1_key, text2_key) if valid_ds is not None else None |
|
|
| elif task == "classification": |
| |
| possible_text = [k for k in train_ds.column_names if k in ("text", "sentence", "content", "review")] |
| text_key = possible_text[0] if possible_text else train_ds.column_names[0] |
| log(f"[rank {self.rank}] Using field {text_key} for classification task.") |
| train_tok = self._prepare_classification(train_ds, text_key) |
| valid_tok = self._prepare_classification(valid_ds, text_key) if valid_ds is not None else None |
|
|
| elif task == "next_sentence_prediction": |
| |
| |
| log(f"[rank {self.rank}] Synthesizing NSP dataset for next_sentence_prediction.") |
| train_tok = self._synthesize_nsp_dataset(train_ds) |
| valid_tok = None |
| else: |
| raise ValueError(f"Unsupported task: {task}") |
|
|
| |
| def collate_fn_train(batch): |
| |
| return {k: torch.tensor([x[k] for x in batch]) for k in batch[0].keys()} |
|
|
| |
| |
| train_sampler = DistributedSampler(train_tok) if self.use_ddp else None |
| train_loader = DataLoader(train_tok, batch_size=self.batch_size, shuffle=(train_sampler is None), sampler=train_sampler, |
| num_workers=self.num_workers, pin_memory=True, drop_last=True) |
| eval_loader = None |
| if valid_ds is not None and valid_tok is not None: |
| eval_sampler = DistributedSampler(valid_tok) if self.use_ddp else None |
| eval_loader = DataLoader(valid_tok, batch_size=self.eval_batch_size, shuffle=False, sampler=eval_sampler, |
| num_workers=self.num_workers, pin_memory=True, drop_last=False) |
|
|
| |
| log(f"[rank {self.rank}] Train samples: {len(train_tok)}; Eval samples: {len(valid_tok) if valid_tok is not None else 0}") |
| return train_loader, eval_loader |
|
|
| |
| |
| |
| def _setup_optimizer_and_scheduler(self, total_training_steps): |
| no_decay = ["bias", "LayerNorm.weight"] |
| params = [ |
| {"params": [p for n, p in self.model.named_parameters() if not any(nd in n for nd in no_decay)], "weight_decay": self.weight_decay}, |
| {"params": [p for n, p in self.model.named_parameters() if any(nd in n for nd in no_decay)], "weight_decay": 0.0}, |
| ] |
| optimizer = AdamW(params, lr=self.learning_rate) |
| scheduler = get_linear_schedule_with_warmup(optimizer, num_warmup_steps=self.warmup_steps, num_training_steps=total_training_steps) |
| return optimizer, scheduler |
|
|
| |
| |
| |
| def _forward(self, batch): |
| |
| input_ids = batch["input_ids"].to(self.device) |
| attention_mask = batch["attention_mask"].to(self.device) |
| labels = batch.get("labels", None) |
| if labels is not None: |
| labels = torch.tensor(labels).to(self.device) |
| out = self.model(input_ids=input_ids, attention_mask=attention_mask, labels=labels) |
| loss = out.loss |
| logits = out.logits |
| else: |
| out = self.model(input_ids=input_ids, attention_mask=attention_mask) |
| logits = out.logits |
| loss = None |
| return loss, logits, labels |
|
|
| def _evaluate(self, eval_loader): |
| self.model.eval() |
| all_preds = [] |
| all_labels = [] |
| total_loss = 0.0 |
| n_batches = 0 |
| with torch.no_grad(): |
| for batch in tqdm(eval_loader, desc=f"Eval (rank {self.rank})", disable=(self.rank != 0)): |
| |
| if isinstance(batch, dict) and isinstance(batch.get("labels"), list): |
| |
| batch = {k: torch.tensor(v) if isinstance(v, list) else v for k,v in batch.items()} |
| loss, logits, labels = self._forward(batch) |
| if loss is not None: |
| total_loss += loss.item() |
| if logits is not None: |
| preds = torch.argmax(logits, dim=-1).cpu().tolist() |
| all_preds.extend(preds) |
| if labels is not None: |
| all_labels.extend(labels.cpu().tolist()) |
| n_batches += 1 |
|
|
| |
| if len(all_labels) == 0: |
| |
| return {"loss": total_loss / (n_batches or 1), "accuracy": None} |
| acc = accuracy_score(all_labels, all_preds) |
| precision, recall, f1, _ = precision_recall_fscore_support(all_labels, all_preds, average="weighted", zero_division=0) |
| return {"loss": total_loss / (n_batches or 1), "accuracy": acc, "precision": precision, "recall": recall, "f1": f1} |
|
|
| |
| |
| |
| def _save_checkpoint(self, step_or_epoch): |
| ckpt_dir = self.output_dir / f"ckpt-{step_or_epoch}" |
| ckpt_dir.mkdir(parents=True, exist_ok=True) |
| model_to_save = self.model.module if hasattr(self.model, "module") else self.model |
| model_to_save.save_pretrained(ckpt_dir) |
| self.tokenizer.save_pretrained(ckpt_dir) |
| |
| state = { |
| "global_step": self.global_step, |
| "best_metric": self.best_metric, |
| "cfg": self.cfg |
| } |
| with open(ckpt_dir / "train_state.json", "w") as f: |
| json.dump(state, f) |
| log(f"[rank {self.rank}] Saved checkpoint -> {ckpt_dir}") |
|
|
| |
| |
| |
| def train(self): |
| log(f"[rank {self.rank}] Starting finetune. device={self.device} fp16={self.fp16} ddp={self.use_ddp}") |
| train_loader, eval_loader = self._build_datasets_and_loaders() |
|
|
| |
| steps_per_epoch = math.ceil(len(train_loader) / (1.0 * self.gradient_accumulation_steps)) |
| total_training_steps = int(steps_per_epoch * self.num_epochs) |
| self.total_steps = total_training_steps |
| log(f"[rank {self.rank}] Steps per epoch: {steps_per_epoch}, total training steps: {total_training_steps}") |
|
|
| optimizer, scheduler = self._setup_optimizer_and_scheduler(total_training_steps) |
|
|
| |
| if self.use_ddp: |
| |
| self.model = torch.nn.parallel.DistributedDataParallel(self.model, device_ids=[torch.cuda.current_device()], output_device=torch.cuda.current_device(), find_unused_parameters=False) |
|
|
| |
| self.model.train() |
| optimizer.zero_grad() |
| self.global_step = 0 |
| best_metric = -1.0 |
|
|
| for epoch in range(self.num_epochs): |
| if self.use_ddp: |
| train_loader.sampler.set_epoch(epoch) |
| epoch_loss = 0.0 |
| pbar = tqdm(train_loader, desc=f"Train Epoch {epoch} (rank {self.rank})", disable=(self.rank != 0)) |
| for step, batch in enumerate(pbar): |
| |
| if isinstance(batch.get("labels"), list): |
| batch["labels"] = torch.tensor(batch["labels"]) |
|
|
| |
| with autocast(enabled=(self.fp16 and torch.cuda.is_available())): |
| loss, logits, labels = self._forward(batch) |
| if loss is None: |
| |
| loss = torch.tensor(0.0, device=self.device) |
|
|
| loss = loss / self.gradient_accumulation_steps |
|
|
| |
| if self.scaler is not None: |
| self.scaler.scale(loss).backward() |
| else: |
| loss.backward() |
| epoch_loss += loss.item() * self.gradient_accumulation_steps |
|
|
| |
| if (step + 1) % self.gradient_accumulation_steps == 0: |
| |
| if self.scaler is not None: |
| self.scaler.unscale_(optimizer) |
| torch.nn.utils.clip_grad_norm_(self.model.parameters(), self.max_grad_norm) |
| self.scaler.step(optimizer) |
| self.scaler.update() |
| else: |
| torch.nn.utils.clip_grad_norm_(self.model.parameters(), self.max_grad_norm) |
| optimizer.step() |
| scheduler.step() |
| optimizer.zero_grad() |
| self.global_step += 1 |
|
|
| if self.rank == 0 and (self.global_step % self.logging_steps == 0): |
| pbar.set_postfix({"loss": f"{epoch_loss/((step+1) or 1):.4f}", "step": self.global_step}) |
|
|
| |
| if self.rank == 0 and self.eval_steps and (self.global_step % self.eval_steps == 0): |
| if eval_loader is not None: |
| metrics = self._evaluate(eval_loader) |
| log(f"[rank {self.rank}] Eval at step {self.global_step}: {metrics}") |
| |
| metric_val = metrics.get("f1") or metrics.get("accuracy") or 0.0 |
| if metric_val > best_metric: |
| best_metric = metric_val |
| self.best_metric = best_metric |
| |
| self._save_checkpoint(f"best-step-{self.global_step}") |
|
|
| if self.rank == 0 and self.save_steps and (self.global_step % self.save_steps == 0): |
| self._save_checkpoint(f"step-{self.global_step}") |
|
|
| |
| log(f"[rank {self.rank}] Epoch {epoch} completed. avg_loss={(epoch_loss/len(train_loader)):.4f}") |
|
|
| |
| if eval_loader is not None and self.rank == 0: |
| metrics = self._evaluate(eval_loader) |
| log(f"[rank {self.rank}] Epoch {epoch} eval: {metrics}") |
| metric_val = metrics.get("f1") or metrics.get("accuracy") or 0.0 |
| if metric_val > best_metric: |
| best_metric = metric_val |
| self.best_metric = best_metric |
| self._save_checkpoint(f"best-epoch-{epoch}") |
|
|
| log(f"[rank {self.rank}] Training complete. Best metric: {self.best_metric}") |
| |
| if self.rank == 0: |
| self._save_checkpoint("final") |
| |
| if self.use_ddp and torch.distributed.is_initialized(): |
| torch.distributed.barrier() |
| torch.distributed.destroy_process_group() |
|
|
|
|