CoolFace
Apppublic

Caffin/bert-dllm

sourceHugging Faceupdated 21d agoView on Hugging Face
0likes
train.py130 linesDownload Raw Back to root
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