Caffin/bert-dllm
0
1"""Fine-tune a RoBERTa masked LM as the dLLM checkpoint used by this Space."""2 3from __future__ import annotations4 5import argparse6import random7from dataclasses import dataclass8from typing import Any9 10import torch11from datasets import Dataset, load_dataset12from transformers import (13 AutoModelForMaskedLM,14 AutoTokenizer,15 Trainer,16 TrainingArguments,17 set_seed,18)19 20from diffusion import DiffusionConfig21 22 23@dataclass24class DiffusionDataCollator:25 """Apply a randomly selected dLLM masking rate to each training batch."""26 27 tokenizer: Any28 config: DiffusionConfig29 30 def __post_init__(self) -> None:31 self.special_ids = set(self.tokenizer.all_special_ids)32 33 def __call__(self, features: list[dict[str, str]]) -> dict[str, torch.Tensor]:34 texts = [feature["text"] for feature in features]35 batch = self.tokenizer(36 texts,37 max_length=self.config.canvas_length,38 truncation=True,39 padding="max_length",40 return_tensors="pt",41 )42 original_ids = batch["input_ids"].clone()43 candidate_positions = batch["attention_mask"].bool()44 candidate_positions[:, : self.config.prefix_length] = False45 for special_id in self.special_ids:46 candidate_positions &= original_ids.ne(special_id)47 48 mask_probability = random.choice(self.config.mask_probabilities)49 selected = torch.rand_like(original_ids, dtype=torch.float).lt(mask_probability)50 selected &= candidate_positions51 for row in range(selected.shape[0]):52 if candidate_positions[row].any() and not selected[row].any():53 candidates = candidate_positions[row].nonzero(as_tuple=False).flatten()54 selected[row, candidates[torch.randint(len(candidates), (1,))]] = True55 56 batch["input_ids"][selected] = self.tokenizer.mask_token_id57 labels = original_ids58 labels[~selected] = -10059 batch["labels"] = labels60 return batch61 62 63def parse_args() -> argparse.Namespace:64 parser = argparse.ArgumentParser(description=__doc__)65 parser.add_argument("--model-id", default="roberta-base")66 parser.add_argument("--dataset", default="Salesforce/wikitext")67 parser.add_argument("--dataset-config", default="wikitext-2-raw-v1")68 parser.add_argument("--split", default="train")69 parser.add_argument("--output-dir", default="bert-dllm-wikitext2")70 parser.add_argument("--hub-model-id", default=None)71 parser.add_argument("--max-steps", type=int, default=1_000)72 parser.add_argument("--batch-size", type=int, default=16)73 parser.add_argument("--learning-rate", type=float, default=5e-5)74 parser.add_argument("--seed", type=int, default=42)75 return parser.parse_args()76 77 78def main() -> None:79 """Train and optionally publish the fixed-prefix, variable-mask checkpoint."""80 args = parse_args()81 set_seed(args.seed)82 config = DiffusionConfig()83 84 tokenizer = AutoTokenizer.from_pretrained(args.model_id, use_fast=True)85 if tokenizer.mask_token_id is None:86 raise ValueError(f"{args.model_id} is not a masked-language model tokenizer")87 model = AutoModelForMaskedLM.from_pretrained(args.model_id)88 89 dataset: Dataset = load_dataset(args.dataset, args.dataset_config, split=args.split)90 dataset = dataset.filter(91 lambda row: len(92 tokenizer(row["text"], add_special_tokens=True, truncation=True)["input_ids"]93 )94 > config.prefix_length + 195 )96 collator = DiffusionDataCollator(tokenizer=tokenizer, config=config)97 98 training_args = TrainingArguments(99 output_dir=args.output_dir,100 max_steps=args.max_steps,101 per_device_train_batch_size=args.batch_size,102 learning_rate=args.learning_rate,103 optim="adamw_torch",104 save_strategy="steps",105 save_steps=max(100, args.max_steps // 2),106 save_total_limit=1,107 logging_steps=25,108 remove_unused_columns=False,109 dataloader_pin_memory=False,110 report_to="none",111 seed=args.seed,112 )113 trainer = Trainer(114 model=model,115 args=training_args,116 train_dataset=dataset,117 data_collator=collator,118 )119 trainer.train()120 trainer.save_model(args.output_dir)121 tokenizer.save_pretrained(args.output_dir)122 123 if args.hub_model_id:124 model.push_to_hub(args.hub_model_id, commit_message="Train RoBERTa dLLM checkpoint")125 tokenizer.push_to_hub(args.hub_model_id, commit_message="Upload tokenizer")126 127 128if __name__ == "__main__":129 main()130 