# main.py (updated fragment to add bert_finetune mode) import os import re import shutil import csv import sys import json from pathlib import Path import torch import torch.distributed as dist from dotenv import load_dotenv import hydra from omegaconf import OmegaConf # existing imports from your repo from lmr.config import initialize_config from lmr.tokenizer import Tokenizer from lmr.models import get_model from lmr.data import get_dataset_splits from lmr.checkpointing import Checkpointing from lmr.utils.seed import set_seed from lmr.training import Bert_Trainer, Trainer from lmr.generation import Generator from lmr.benchmark import Benchmark from lmr.ddp import unwrap_model # NEW: import the BertFineTuneTrainer implementation you created earlier. # Adjust this import to where you saved the class. Example: lmr.training.bert_finetune_trainer from lmr.training.bert_finetune_trainer import BertFineTuneTrainer # transformers from transformers import AutoTokenizer, AutoConfig, AutoModelForSequenceClassification, BertForNextSentencePrediction DATASET_DIR = Path("datasets") CHECKPOINT_DIR = Path("/work/jf381/checkpoints") BENCHMARK_DIR = Path("output") # ------------------------- # Helper: create model+tokenizer for finetune mode # ------------------------- def create_tokenizer_and_model_for_finetune(model_name_or_path: str, task: str, num_labels: int = None): """ Returns (tokenizer, model). - task == "sentence_pair" -> AutoModelForSequenceClassification (num_labels required or inferred) - task == "next_sentence_prediction" -> BertForNextSentencePrediction """ tokenizer = AutoTokenizer.from_pretrained(model_name_or_path, use_fast=True) if task == "sentence_pair": # If user didn't pass num_labels, default to 2 (binary) nlab = num_labels if num_labels is not None else 2 model = AutoModelForSequenceClassification.from_pretrained(model_name_or_path, num_labels=nlab) elif task == "next_sentence_prediction": # load BERT NSP model which exposes next sentence prediction head model = BertForNextSentencePrediction.from_pretrained(model_name_or_path) else: raise ValueError(f"Unsupported finetune task: {task}") return tokenizer, model # ============================================================================= # MAIN # ============================================================================= @hydra.main(config_path="config", config_name="config", version_base="1.3") def main(config): # common init load_dotenv() set_seed(config) initialize_config(config) mode = config.mode maybe_print = print # keep simple; you may swap for Logger if desired # existing modes (train/generate/...) if mode == "train": # unchanged: legacy pretrain flow if torch.cuda.is_available() and torch.cuda.device_count() > 1: if not dist.is_initialized(): dist.init_process_group(backend="nccl") torch.cuda.set_device(dist.get_rank() % torch.cuda.device_count()) tokenizer, model = setup_model_and_tokenizer(config) tokenized_dataset_dir = DATASET_DIR / config.tokenizer_base splits = get_dataset_splits(config.dataset, 1024, tokenized_dataset_dir) checkpointing = Checkpointing(model, CHECKPOINT_DIR / config.checkpoint_name) if "bert" in str(config.model).lower(): from lmr.training import Bert_Trainer trainer = Bert_Trainer(config.training, model, tokenizer, splits, checkpointing, None) else: trainer = Trainer(config.training, model, tokenizer, splits, checkpointing) trainer.train() elif mode == "generate": generate(config) elif mode == "generate_all": generate_all(config) elif mode == "benchmark": tokenizer, model = setup_model_and_tokenizer(config) checkpointing = Checkpointing(model, CHECKPOINT_DIR / config.checkpoint_name) benchmarking = Benchmark(config.benchmark, model, tokenizer, checkpointing, BENCHMARK_DIR / config.checkpoint_name) benchmarking.run_benchmarks() # --------------------------- # NEW: bert fine-tune mode using HF datasets directly # --------------------------- elif mode == "bert_finetune": # config should contain fields: # model_name_or_path: pretrained model (e.g., bert-base-uncased or local folder) # task: "sentence_pair" | "next_sentence_prediction" # dataset: huggingface dataset id (e.g., "glue/mrpc" or "glue", config "mrpc") # batch_size, num_epochs, lr, save_dir, use_ddp, etc maybe_print("=== BERT FINETUNE MODE ===") maybe_print(f"Config: task={config.finetune.task}, dataset={config.finetune.dataset}, model={config.finetune.model_name_or_path}") # create tokenizer + model appropriate for finetune task tokenizer, model = create_tokenizer_and_model_for_finetune( model_name_or_path=config.finetune.model_name_or_path, task=config.finetune.task, num_labels=getattr(config.finetune, "num_labels", None) ) # device use_ddp = getattr(config.finetune, "use_ddp", False) if use_ddp and torch.cuda.is_available(): # torchrun should provide local_rank; hydra may not, so rely on env var for rank local_rank = int(os.environ.get("LOCAL_RANK", os.environ.get("RANK", 0))) world_size = int(os.environ.get("WORLD_SIZE", 1)) device = torch.device(f"cuda:{local_rank}") maybe_print(f"[finetune] DDP mode: local_rank={local_rank} world_size={world_size} device={device}") cfg_namespace = argparse.Namespace(**config.finetune) # convert to simple namespace # instantiate trainer with DDP flags finetune_trainer = BertFineTuneTrainer(model=model, tokenizer=tokenizer, cfg=cfg_namespace, device=device) # caller should run via torchrun --nproc_per_node=N python main.py mode=bert_finetune ... finetune_trainer.train() else: # single-GPU / CPU device = torch.device("cuda" if torch.cuda.is_available() else "cpu") maybe_print(f"[finetune] Single-device mode: device={device}") cfg_namespace = argparse.Namespace(**config.finetune) finetune_trainer = BertFineTuneTrainer(model=model, tokenizer=tokenizer, cfg=cfg_namespace, device=device) finetune_trainer.train() else: print(f"❌ Unknown mode: {mode}") if __name__ == "__main__": main()