| |
| """ |
| Domain-Adaptive MLM Pretrainer |
| |
| Continues pretraining BERT-family models (BERT, RoBERTa, DeBERTa, ELECTRA, RemBERT, etc.) |
| on an unlabelled domain corpus using Masked Language Modelling (MLM). |
| |
| Outputs a saved encoder checkpoint per model that can be referenced directly as |
| `model_name` in ensemble_config.json for downstream fine-tuning. |
| |
| Usage: |
| python mlm_pretrainer.py \ |
| --data_path listings.parquet \ |
| --text_col text \ |
| --pretrain_config_path pretrain_config.json \ |
| --output_base_dir pretrained_checkpoints |
| """ |
|
|
| import signal |
| import atexit |
| import argparse |
| import gc |
| import json |
| import warnings |
| import sys |
| import copy |
| import shutil |
| import re |
| from datetime import datetime |
| from pathlib import Path |
|
|
| import numpy as np |
| import torch |
| from datasets import Dataset |
| from transformers import ( |
| AutoModelForMaskedLM, |
| AutoTokenizer, |
| DataCollatorForLanguageModeling, |
| EarlyStoppingCallback, |
| Trainer, |
| TrainerCallback, |
| TrainingArguments, |
| ) |
|
|
| class DataCollatorForWholeWordMasking(DataCollatorForLanguageModeling): |
| """ |
| Whole-word masking collator. When a token is selected for masking, all |
| sub-tokens belonging to the same word are masked together. |
| Falls back to standard token masking for tokenisers that don't set word_ids. |
| """ |
| def torch_call(self, examples): |
| batch = self.tokenizer.pad(examples, return_tensors="pt", |
| pad_to_multiple_of=self.pad_to_multiple_of) |
| input_ids = batch["input_ids"].clone() |
| labels = batch["input_ids"].clone() |
|
|
| for i, (ids, encoding) in enumerate(zip(input_ids, examples)): |
| word_ids = None |
| if hasattr(encoding, "word_ids"): |
| word_ids = encoding.word_ids() |
| elif "word_ids" in encoding: |
| word_ids = encoding["word_ids"] |
|
|
| if word_ids is None: |
| |
| probability_matrix = torch.full(ids.shape, self.mlm_probability) |
| special_tokens_mask = self.tokenizer.get_special_tokens_mask( |
| ids.tolist(), already_has_special_tokens=True) |
| probability_matrix[torch.tensor(special_tokens_mask, dtype=torch.bool)] = 0.0 |
| masked_indices = torch.bernoulli(probability_matrix).bool() |
| labels[i][~masked_indices] = -100 |
| input_ids[i][masked_indices] = self.tokenizer.mask_token_id |
| continue |
|
|
| |
| word_to_tokens: dict = {} |
| for tok_idx, word_idx in enumerate(word_ids): |
| if word_idx is None: |
| continue |
| word_to_tokens.setdefault(word_idx, []).append(tok_idx) |
|
|
| unique_words = list(word_to_tokens.keys()) |
| num_to_mask = max(1, int(round(len(unique_words) * self.mlm_probability))) |
| words_to_mask = np.random.choice(unique_words, size=num_to_mask, replace=False) |
| masked_indices = torch.zeros(ids.shape, dtype=torch.bool) |
| for w in words_to_mask: |
| for tok_idx in word_to_tokens[w]: |
| if tok_idx < len(masked_indices): |
| masked_indices[tok_idx] = True |
|
|
| labels[i][~masked_indices] = -100 |
|
|
| |
| replace_with_mask = torch.bernoulli(torch.full(ids.shape, 0.8)).bool() & masked_indices |
| replace_with_random = (torch.bernoulli(torch.full(ids.shape, 0.5)).bool() |
| & masked_indices & ~replace_with_mask) |
| input_ids[i][replace_with_mask] = self.tokenizer.mask_token_id |
| input_ids[i][replace_with_random] = torch.randint( |
| len(self.tokenizer), ids.shape, dtype=torch.long)[replace_with_random] |
|
|
| batch["input_ids"] = input_ids |
| batch["labels"] = labels |
| return batch |
|
|
| warnings.filterwarnings("ignore") |
|
|
|
|
| |
| |
| |
| class LoggerTee(object): |
| def __init__(self, filename, mode="a"): |
| self.terminal = sys.stdout |
| self.log = open(filename, mode) |
| self.line_buffer = "" |
| self.ansi_escape = re.compile(r'\x1B(?:[@-Z\\-_]|\[[0-?]*[ -/]*[@-~])') |
| self.tqdm_pattern = re.compile(r'\b\d+%\s*\|') |
|
|
| def write(self, message): |
| self.terminal.write(message) |
| for char in message: |
| if char == '\r': |
| self.line_buffer = "" |
| elif char == '\n': |
| clean_line = self.ansi_escape.sub('', self.line_buffer) |
| if self.tqdm_pattern.search(clean_line): |
| if "100%|" not in clean_line.replace(" ", ""): |
| self.line_buffer = "" |
| continue |
| self.log.write(clean_line + '\n') |
| self.log.flush() |
| self.line_buffer = "" |
| else: |
| self.line_buffer += char |
|
|
| def flush(self): |
| self.terminal.flush() |
| self.log.flush() |
|
|
| def isatty(self): |
| return self.terminal.isatty() |
|
|
|
|
| timestamp = datetime.now().strftime("%Y%m%d_%H%M") |
| log_filename = f"pretrain_log_{timestamp}.txt" |
| sys.stdout = LoggerTee(filename=log_filename) |
| sys.stderr = sys.stdout |
|
|
|
|
| |
| |
| |
| def parse_args(): |
| parser = argparse.ArgumentParser( |
| description="Domain-Adaptive MLM Pretrainer for BERT-family models." |
| ) |
|
|
| |
| parser.add_argument("--data_path", type=str, required=True, |
| help="Path to unlabelled corpus (.parquet or .csv).") |
| parser.add_argument("--text_col", type=str, required=True, |
| help="Column name containing raw text.") |
| parser.add_argument("--eval_split", type=float, default=0.05, |
| help="Fraction of corpus held out for perplexity evaluation.") |
|
|
| |
| parser.add_argument("--pretrain_config_path", type=str, required=True, |
| help="Path to pretrain_config.json.") |
| parser.add_argument("--output_base_dir", type=str, default="pretrained_checkpoints", |
| help="Root directory; each model gets its own subdirectory here.") |
| parser.add_argument("--metadata_path", type=str, default="pretrain_metadata.json", |
| help="JSON file recording completed models and their final perplexity.") |
|
|
| |
| parser.add_argument("--max_length", type=int, default=256) |
| parser.add_argument("--batch_size", type=int, default=16) |
| parser.add_argument("--gradient_accumulation_steps", type=int, default=1) |
| parser.add_argument("--max_steps", type=int, default=10000) |
| parser.add_argument("--eval_steps", type=int, default=500) |
| parser.add_argument("--early_stopping_patience", type=int, default=5) |
| parser.add_argument("--learning_rate", type=float, default=5e-5) |
| parser.add_argument("--warmup_steps", type=int, default=500) |
| parser.add_argument("--weight_decay", type=float, default=0.01) |
| parser.add_argument("--adam_epsilon", type=float, default=1e-8) |
| parser.add_argument("--max_grad_norm", type=float, default=1.0) |
| parser.add_argument("--lr_scheduler_type", type=str, default="cosine", |
| choices=["linear", "cosine", "cosine_with_restarts", |
| "polynomial", "constant", "constant_with_warmup", |
| "inverse_sqrt"]) |
| parser.add_argument("--mlm_probability", type=float, default=0.15, |
| help="Fraction of tokens masked per sequence.") |
| parser.add_argument("--whole_word_masking", action="store_true", |
| help="Use whole-word masking instead of sub-token masking. " |
| "Recommended for syllable-level tokenisers (PhoBERT, ViDeBERTa).") |
| parser.add_argument("--freeze_layers", type=int, default=0, |
| help="Freeze the bottom N encoder layers during pretraining.") |
| parser.add_argument("--use_bf16", action="store_true") |
| parser.add_argument("--seed", type=int, default=42) |
|
|
| |
| parser.add_argument("--retrain", action="store_true", |
| help="Retrain ALL active models from their local checkpoint.") |
| parser.add_argument("--retrain_models", type=str, nargs="+", default=None, |
| metavar="MODEL_NAME", |
| help="Retrain specific models by name (space-separated). " |
| "Example: --retrain_models microsoft/mdeberta-v3-base google/rembert") |
|
|
| return parser.parse_args() |
|
|
|
|
| |
| |
| |
| _PRETRAIN_ARG_SPECS = [ |
| |
| ("max_length", "max_length", 256), |
| |
| ("batch_size", "batch_size", 16), |
| ("gradient_accumulation_steps", "gradient_accumulation_steps", 1), |
| |
| ("learning_rate", "learning_rate", 5e-5), |
| ("weight_decay", "weight_decay", 0.01), |
| ("adam_epsilon", "adam_epsilon", 1e-8), |
| ("max_grad_norm", "max_grad_norm", 1.0), |
| |
| ("warmup_steps", "warmup_steps", 500), |
| ("lr_scheduler_type", "lr_scheduler_type", "cosine"), |
| |
| ("max_steps", "max_steps", 10000), |
| ("eval_steps", "eval_steps", 500), |
| ("early_stopping_patience", "early_stopping_patience", 5), |
| |
| ("mlm_probability", "mlm_probability", 0.15), |
| ("whole_word_masking", "whole_word_masking", False), |
| |
| ("freeze_layers", "freeze_layers", 0), |
| |
| ("use_bf16", "use_bf16", False), |
| ("drop_token_type_ids", "drop_token_type_ids", None), |
| |
| ("tokenizer_name", "tokenizer_name", None), |
| ] |
|
|
|
|
| def resolve_model_args(global_args, config: dict): |
| """ |
| Returns a namespace where every training arg is resolved with priority: |
| per-model JSON config > global CLI args > hardcoded default |
| """ |
| resolved = copy.deepcopy(global_args) |
| for json_key, attr, default in _PRETRAIN_ARG_SPECS: |
| if json_key in config: |
| setattr(resolved, attr, config[json_key]) |
| elif not hasattr(resolved, attr) or getattr(resolved, attr) is None: |
| setattr(resolved, attr, default) |
| if not getattr(resolved, "tokenizer_name", None): |
| resolved.tokenizer_name = config["model_name"] |
| resolved.model_name = config["model_name"] |
| return resolved |
|
|
|
|
| def load_pretrain_configs(config_path: str) -> list: |
| with open(config_path, "r") as f: |
| return json.load(f) |
|
|
|
|
| |
| |
| |
| def get_output_dir(output_base: str, model_name: str) -> Path: |
| safe_name = model_name.replace("/", "__") |
| return Path(output_base) / safe_name |
|
|
|
|
| def pretrain_artifact_exists(output_base: str, model_name: str) -> bool: |
| """A completed pretraining run leaves a config.json written by save_pretrained.""" |
| return (get_output_dir(output_base, model_name) / "config.json").exists() |
|
|
|
|
| def load_pretrain_metadata(output_base: str, model_name: str) -> dict: |
| meta_path = get_output_dir(output_base, model_name) / "pretrain_metadata.json" |
| if meta_path.exists(): |
| with open(meta_path) as f: |
| return json.load(f) |
| return {} |
|
|
|
|
| def save_pretrain_metadata(output_dir: Path, perplexity: float, eval_loss: float, |
| model_name: str, args): |
| meta = { |
| "model_name": model_name, |
| "perplexity": round(perplexity, 4), |
| "eval_loss": round(eval_loss, 6), |
| "mlm_probability": args.mlm_probability, |
| "whole_word_masking": args.whole_word_masking, |
| "max_steps": args.max_steps, |
| "learning_rate": args.learning_rate, |
| "pretrained_at": datetime.now().strftime("%Y-%m-%d %H:%M:%S"), |
| } |
| with open(output_dir / "pretrain_metadata.json", "w") as f: |
| json.dump(meta, f, indent=2) |
|
|
|
|
| def mark_model_pretrained(config_path: str, model_name: str, |
| perplexity: float, output_dir: str): |
| """Atomically update the JSON config to mark a model as done.""" |
| configs = load_pretrain_configs(config_path) |
| for cfg in configs: |
| if cfg["model_name"] == model_name: |
| cfg["pretrained"] = True |
| cfg["pretrained_at"] = datetime.now().strftime("%Y-%m-%d %H:%M:%S") |
| cfg["perplexity"] = round(perplexity, 4) |
| cfg["output_dir"] = str(output_dir) |
| break |
| tmp_path = config_path + ".tmp" |
| with open(tmp_path, "w") as f: |
| json.dump(configs, f, indent=2) |
| Path(tmp_path).replace(config_path) |
| print(f" -> [Checkpoint] Marked '{model_name}' as pretrained " |
| f"(Perplexity={perplexity:.4f})") |
|
|
|
|
| |
| |
| |
| def freeze_encoder_layers(model, num_layers: int): |
| if num_layers <= 0: |
| return |
| encoder = None |
| for attr in ["encoder", "bert", "deberta", "electra", "roberta", "rembert"]: |
| enc = getattr(model, attr, None) |
| if enc is None: |
| |
| |
| enc = getattr(getattr(model, "rembert", None), attr, None) |
| if enc is not None: |
| encoder = getattr(enc, "layer", None) |
| if encoder is not None: |
| break |
| if encoder is None: |
| print(" [WARNING] Could not locate encoder layers to freeze — skipping.") |
| return |
| actual = min(num_layers, len(encoder)) |
| for i in range(actual): |
| for param in encoder[i].parameters(): |
| param.requires_grad = False |
| print(f" -> Froze {actual}/{len(encoder)} encoder layers.") |
|
|
|
|
| |
| |
| |
| class FormattedEvalCallback(TrainerCallback): |
| """Prints a clean evaluation summary after each eval step.""" |
|
|
| def on_evaluate(self, args, state, control, metrics=None, **kwargs): |
| if metrics is None: |
| return |
| loss = metrics.get("eval_loss") |
| ppl = metrics.get("eval_perplexity") |
| step = state.global_step |
| print(f"\n [Eval @ step {step}] loss: {loss:.4f} | " |
| f"perplexity: {ppl:.4f}" if (loss and ppl) else |
| f"\n [Eval @ step {step}] {metrics}") |
| if state.best_metric is not None: |
| print(f" Best so far → {state.best_metric:.4f} (eval_loss)") |
| print("-" * 60) |
|
|
|
|
| |
| |
| |
| def compute_mlm_metrics(eval_pred): |
| """Perplexity from MLM eval loss. The Trainer passes (logits, labels) but |
| for MLM we only need the scalar loss, which the Trainer already computes |
| and logs as eval_loss. We add perplexity here as a derived metric.""" |
| |
| |
| |
| |
| |
| logits = eval_pred.predictions |
| |
| |
| return {} |
|
|
|
|
| class PerplexityCallback(TrainerCallback): |
| """Injects eval_perplexity into the metrics dict after each evaluation.""" |
|
|
| def on_evaluate(self, args, state, control, metrics=None, **kwargs): |
| if metrics and "eval_loss" in metrics: |
| try: |
| metrics["eval_perplexity"] = float( |
| torch.exp(torch.tensor(metrics["eval_loss"])).item() |
| ) |
| except Exception: |
| metrics["eval_perplexity"] = float("inf") |
|
|
|
|
| |
| |
| |
| def build_dataset(hf_dataset, tokenizer, max_length: int, |
| cache_path: str, num_proc: int = 4): |
| """ |
| Tokenise a HuggingFace Dataset for MLM. |
| |
| - hf_dataset : already-split HF Dataset (Arrow-backed, memory-mapped) |
| - cache_path : path for the tokenized Arrow cache file; reused on reruns |
| - num_proc : parallel tokenisation workers |
| |
| The tokenized result is written to disk so it never fully resides in RAM. |
| """ |
| def tokenize_fn(examples): |
| return tokenizer( |
| examples["text"], |
| truncation=True, |
| max_length=max_length, |
| padding=False, |
| return_special_tokens_mask=True, |
| ) |
|
|
| ds = hf_dataset.map( |
| tokenize_fn, |
| batched=True, |
| batch_size=1000, |
| num_proc=num_proc, |
| remove_columns=["text"], |
| keep_in_memory=False, |
| cache_file_name=cache_path, |
| desc="Tokenising", |
| ) |
| return ds |
|
|
|
|
| |
| |
| |
| def pretrain_model(args, config: dict, train_ds_raw, eval_ds_raw, cache_dir: Path): |
| """ |
| Domain-adaptive MLM pretraining for a single model. |
| |
| Saves the full ForMaskedLM checkpoint to output_dir. Downstream fine-tuning |
| scripts can load this with AutoModelForSequenceClassification.from_pretrained( |
| output_dir, ignore_mismatched_sizes=True) and the MLM head is simply ignored. |
| """ |
| model_name = config["model_name"] |
| output_dir = get_output_dir(args.output_base_dir, model_name) |
| output_dir.mkdir(parents=True, exist_ok=True) |
| tmp_ckpt_dir = output_dir / "tmp_checkpoints" |
|
|
| print(f"\n{'='*70}") |
| print(f"PRETRAINING: {model_name}") |
| print(f" -> output_dir: {output_dir}") |
| print(f" -> tokenizer: {args.tokenizer_name}") |
| print(f" -> mlm_probability: {args.mlm_probability}") |
| print(f" -> whole_word_masking: {args.whole_word_masking}") |
| print(f" -> max_length: {args.max_length}") |
| print(f" -> max_steps: {args.max_steps} | eval_steps: {args.eval_steps}") |
| print(f" -> lr: {args.learning_rate} | warmup: {args.warmup_steps} " |
| f"| scheduler: {args.lr_scheduler_type}") |
| print(f" -> batch_size: {args.batch_size} " |
| f"| grad_accum: {args.gradient_accumulation_steps}") |
| print(f" -> freeze_layers: {args.freeze_layers}") |
| print(f"{'='*70}") |
|
|
| |
| |
| |
| tokenizer = AutoTokenizer.from_pretrained(args.tokenizer_name) |
|
|
| |
| |
| |
| safe_name = args.model_name.replace("/", "__") |
| cache_base = cache_dir / safe_name |
| cache_base.mkdir(parents=True, exist_ok=True) |
|
|
| print(f"\n Tokenising {len(train_ds_raw):,} train / {len(eval_ds_raw):,} texts " |
| f"(disk-cached at {cache_base}) ...") |
| train_ds = build_dataset(train_ds_raw, tokenizer, args.max_length, |
| cache_path=str(cache_base / "train.arrow")) |
| eval_ds = build_dataset(eval_ds_raw, tokenizer, args.max_length, |
| cache_path=str(cache_base / "eval.arrow")) |
| print(f" Done. Train examples: {len(train_ds):,} | Eval examples: {len(eval_ds):,}") |
|
|
| |
| |
| |
| if args.whole_word_masking: |
| print(" -> Using DataCollatorForWholeWordMasking") |
| collator = DataCollatorForWholeWordMasking( |
| tokenizer=tokenizer, |
| mlm=True, |
| mlm_probability=args.mlm_probability, |
| ) |
| else: |
| print(" -> Using DataCollatorForLanguageModeling (sub-token masking)") |
| collator = DataCollatorForLanguageModeling( |
| tokenizer=tokenizer, |
| mlm=True, |
| mlm_probability=args.mlm_probability, |
| ) |
|
|
| |
| |
| |
| |
| model_type = "" |
| try: |
| from transformers import AutoConfig |
| cfg = AutoConfig.from_pretrained(model_name) |
| model_type = getattr(cfg, "model_type", "") |
| except Exception: |
| pass |
|
|
| _no_tti_types = ["xlm-roberta", "roberta", "camembert", "deberta-v2", |
| "distilbert", "bart", "longformer"] |
| drop_tti = (args.drop_token_type_ids |
| if args.drop_token_type_ids is not None |
| else model_type in _no_tti_types) |
| if drop_tti: |
| print(f" -> Dropping token_type_ids (model_type='{model_type}')") |
| if "token_type_ids" in train_ds.column_names: |
| train_ds = train_ds.remove_columns(["token_type_ids"]) |
| if "token_type_ids" in eval_ds.column_names: |
| eval_ds = eval_ds.remove_columns(["token_type_ids"]) |
|
|
| model = AutoModelForMaskedLM.from_pretrained( |
| args.load_from, |
| use_safetensors=True, |
| ignore_mismatched_sizes=True, |
| torch_dtype=torch.float32, |
| ).to("cuda") |
| freeze_encoder_layers(model, args.freeze_layers) |
|
|
| total_params = sum(p.numel() for p in model.parameters()) |
| trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad) |
| print(f" -> Parameters: {total_params:,} total | {trainable_params:,} trainable") |
|
|
| |
| |
| |
| training_args = TrainingArguments( |
| output_dir=str(tmp_ckpt_dir), |
| max_steps=args.max_steps, |
| per_device_train_batch_size=args.batch_size, |
| per_device_eval_batch_size=args.batch_size, |
| gradient_accumulation_steps=args.gradient_accumulation_steps, |
| learning_rate=args.learning_rate, |
| weight_decay=args.weight_decay, |
| adam_epsilon=args.adam_epsilon, |
| max_grad_norm=args.max_grad_norm, |
| warmup_steps=args.warmup_steps, |
| lr_scheduler_type=args.lr_scheduler_type, |
| eval_strategy="steps", |
| eval_steps=args.eval_steps, |
| save_strategy="steps", |
| save_steps=args.eval_steps, |
| load_best_model_at_end=True, |
| metric_for_best_model="eval_loss", |
| greater_is_better=False, |
| logging_strategy="steps", |
| logging_steps=args.eval_steps, |
| save_total_limit=1, |
| report_to="none", |
| bf16=args.use_bf16, |
| prediction_loss_only=True, |
| ) |
|
|
| |
| |
| |
| trainer = Trainer( |
| model=model, |
| args=training_args, |
| train_dataset=train_ds, |
| eval_dataset=eval_ds, |
| data_collator=collator, |
| callbacks=[ |
| EarlyStoppingCallback(args.early_stopping_patience), |
| PerplexityCallback(), |
| FormattedEvalCallback(), |
| ], |
| ) |
|
|
| |
| |
| |
| trainer.train() |
|
|
| |
| |
| |
| eval_results = trainer.evaluate() |
| eval_loss = eval_results.get("eval_loss", float("inf")) |
| perplexity = float(torch.exp(torch.tensor(eval_loss)).item()) |
| print(f"\n [Final Eval] loss: {eval_loss:.4f} | perplexity: {perplexity:.4f}") |
|
|
| |
| |
| |
| |
| trainer.save_model(str(output_dir)) |
| tokenizer.save_pretrained(str(output_dir)) |
| save_pretrain_metadata(output_dir, perplexity, eval_loss, model_name, args) |
| print(f" -> Saved model + tokenizer → {output_dir}") |
|
|
| |
| if tmp_ckpt_dir.exists(): |
| shutil.rmtree(tmp_ckpt_dir) |
|
|
| del trainer, model |
| torch.cuda.empty_cache() |
| gc.collect() |
|
|
| return perplexity |
|
|
|
|
| |
| |
| |
| def main(): |
| args = parse_args() |
| torch.manual_seed(args.seed) |
| np.random.seed(args.seed) |
|
|
| Path(args.output_base_dir).mkdir(parents=True, exist_ok=True) |
|
|
| |
| |
| |
| print(f"\nLoading corpus from: {args.data_path}") |
| |
| |
| from datasets import load_dataset as hf_load_dataset |
| ext = args.data_path.rsplit(".", 1)[-1].lower() |
| fmt = "parquet" if ext == "parquet" else "csv" |
| raw_ds = hf_load_dataset(fmt, data_files=args.data_path, split="train") |
|
|
| |
| raw_ds = raw_ds.select_columns([args.text_col]) |
| raw_ds = raw_ds.filter(lambda x: x[args.text_col] is not None and |
| str(x[args.text_col]).strip() != "", |
| num_proc=4) |
| |
| if args.text_col != "text": |
| raw_ds = raw_ds.rename_column(args.text_col, "text") |
|
|
| print(f" -> {len(raw_ds):,} texts loaded (memory-mapped).") |
|
|
| |
| splits = raw_ds.train_test_split(test_size=args.eval_split, seed=args.seed) |
| train_ds_raw = splits["train"] |
| eval_ds_raw = splits["test"] |
|
|
| |
| max_eval = 10_000 |
| if len(eval_ds_raw) > max_eval: |
| eval_ds_raw = eval_ds_raw.select(range(max_eval)) |
|
|
| cache_dir = Path(args.output_base_dir) / "_token_cache" |
| cache_dir.mkdir(parents=True, exist_ok=True) |
|
|
| print(f" -> Train: {len(train_ds_raw):,} | Eval: {len(eval_ds_raw):,} (capped at {max_eval:,})") |
| print(f" -> Token cache dir: {cache_dir}") |
|
|
| |
| |
| |
| configs = load_pretrain_configs(args.pretrain_config_path) |
| active_configs = [c for c in configs if c.get("use", True)] |
| print(f"\n -> {len(active_configs)} model(s) active in config.") |
|
|
| |
| |
| |
| completed_models = [] |
| _interrupted = {"flag": False} |
|
|
| def _emergency_save(signum=None, frame=None): |
| if _interrupted["flag"]: |
| return |
| _interrupted["flag"] = True |
| sig_name = f"signal {signum}" if signum else "exit" |
| print(f"\n\n[INTERRUPT] Caught {sig_name}.") |
| if completed_models: |
| with open(args.metadata_path, "w") as f: |
| json.dump({"completed": completed_models}, f, indent=2) |
| print(f" -> Metadata saved → {args.metadata_path}") |
| print(f" -> Artifacts safe in: {args.output_base_dir}/") |
| print(f" -> Re-run with the same command to resume.") |
| else: |
| print(" -> No models completed yet.") |
| if signum is not None: |
| sys.exit(1) |
|
|
| signal.signal(signal.SIGINT, _emergency_save) |
| signal.signal(signal.SIGTERM, _emergency_save) |
| atexit.register(_emergency_save) |
|
|
| |
| |
| |
| print("\n" + "="*70) |
| print(f"INITIATING DOMAIN-ADAPTIVE MLM PRETRAINING ({len(active_configs)} models)") |
| print(f"Output root: {args.output_base_dir}") |
| print("="*70) |
|
|
| for idx, config in enumerate(active_configs): |
| model_name = config["model_name"] |
| local_dir = get_output_dir(args.output_base_dir, model_name) |
| local_exists = pretrain_artifact_exists(args.output_base_dir, model_name) |
|
|
| |
| is_retrain_all = getattr(args, "retrain", False) |
| retrain_list = getattr(args, "retrain_models", None) or [] |
| is_retrain_specific = model_name in retrain_list |
|
|
| should_retrain = local_exists and (is_retrain_all or is_retrain_specific) |
| should_skip = local_exists and not should_retrain |
|
|
| if should_skip: |
| meta = load_pretrain_metadata(args.output_base_dir, model_name) |
| ppl = meta.get("perplexity", "?") |
| saved_at = meta.get("pretrained_at", "?") |
| print(f"\n\n{'*'*70}") |
| print(f"SKIPPING {idx+1}/{len(active_configs)}: {model_name}") |
| print(f" -> Already pretrained at {saved_at} | Perplexity: {ppl}") |
| print(f"{'*'*70}") |
| completed_models.append({"model_name": model_name, "perplexity": ppl, |
| "output_dir": str(local_dir)}) |
| continue |
|
|
| |
| current_args = resolve_model_args(args, config) |
| if should_retrain: |
| current_args.load_from = str(local_dir) |
| print(f"\n\n{'*'*70}") |
| print(f"RETRAINING MODEL {idx+1}/{len(active_configs)}: {model_name}") |
| print(f" -> Loading from local checkpoint: {local_dir}") |
| print(f"{'*'*70}") |
| else: |
| current_args.load_from = model_name |
| print(f"\n\n{'*'*70}") |
| print(f"PRETRAINING MODEL {idx+1}/{len(active_configs)}: {model_name}") |
| print(f"{'*'*70}") |
|
|
| perplexity = pretrain_model(current_args, config, train_ds_raw, eval_ds_raw, cache_dir) |
| output_dir = get_output_dir(args.output_base_dir, model_name) |
|
|
| completed_models.append({"model_name": model_name, |
| "perplexity": round(perplexity, 4), |
| "output_dir": str(output_dir)}) |
| mark_model_pretrained(args.pretrain_config_path, model_name, |
| perplexity, output_dir) |
|
|
| |
| |
| |
| with open(args.metadata_path, "w") as f: |
| json.dump({"completed": completed_models}, f, indent=2) |
|
|
| print("\n" + "="*70) |
| print("PRETRAINING COMPLETE — SUMMARY") |
| print("="*70) |
| print(f" {'Model':<45} {'Perplexity':>12} {'Output Dir'}") |
| print(" " + "-" * 90) |
| for entry in completed_models: |
| name = entry["model_name"].split("/")[-1] |
| print(f" {name:<45} {str(entry['perplexity']):>12} {entry['output_dir']}") |
| print("="*70) |
| print(f"\nUse the output_dir paths as `model_name` in your ensemble_config.json.") |
| print(f"Metadata written → {args.metadata_path}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |