Spaces:
Paused
Paused
File size: 3,934 Bytes
f030d3a 32c0ef6 f030d3a 23e50b1 fd05733 23e50b1 fd05733 f030d3a fd05733 f030d3a 23e50b1 f030d3a 23e50b1 1d2f1ad 23e50b1 1d2f1ad f030d3a 23e50b1 f030d3a | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 | from transformers import TrainingArguments, Trainer
from transformers import DataCollatorForSeq2Seq
from transformers import AutoModelForSeq2SeqLM, AutoTokenizer
from datasets import load_dataset, load_from_disk
from textSummarizer.entity import ModelTrainerConfig
import torch
import os
class ModelTrainer:
def __init__(self, config: ModelTrainerConfig):
self.config = config
def train(self):
device = "cuda" if torch.cuda.is_available() else "cpu"
# Choose model checkpoint: prefer dev_model when dev_run is enabled
model_checkpoint = str(self.config.model_ckpt)
if getattr(self.config, 'dev_run', False) and getattr(self.config, 'dev_model', None):
model_checkpoint = self.config.dev_model
tokenizer = AutoTokenizer.from_pretrained(model_checkpoint)
model_pegasus = AutoModelForSeq2SeqLM.from_pretrained(model_checkpoint).to(device)
seq2seq_data_collator = DataCollatorForSeq2Seq(tokenizer, model=model_pegasus)
dataset_samsum_pt = load_from_disk(str(self.config.data_path))
# If dev_subset is set, use only a small subset of data for quick testing
dev_subset = getattr(self.config, 'dev_subset', 0)
if dev_subset > 0:
if "train" in dataset_samsum_pt:
dataset_samsum_pt["train"] = dataset_samsum_pt["train"].select(range(min(dev_subset, len(dataset_samsum_pt["train"]))))
if "validation" in dataset_samsum_pt:
validation_size = min(dev_subset // 4, len(dataset_samsum_pt["validation"]))
dataset_samsum_pt["validation"] = dataset_samsum_pt["validation"].select(range(validation_size))
# ensure datasets are returned as torch tensors for efficient dataloading
if "train" in dataset_samsum_pt:
try:
dataset_samsum_pt["train"] = dataset_samsum_pt["train"].with_format("torch")
except Exception:
pass
if "validation" in dataset_samsum_pt:
try:
dataset_samsum_pt["validation"] = dataset_samsum_pt["validation"].with_format("torch")
except Exception:
pass
trainer_args = TrainingArguments(
output_dir=str(self.config.root_dir),
num_train_epochs=int(self.config.num_train_epochs),
warmup_steps=int(self.config.warmup_steps),
per_device_train_batch_size=int(self.config.per_device_train_batch_size),
per_device_eval_batch_size=int(self.config.per_device_train_batch_size),
weight_decay=float(self.config.weight_decay),
logging_steps=int(self.config.logging_steps),
eval_strategy=str(self.config.eval_strategy),
eval_steps=int(self.config.eval_steps),
save_steps=int(float(getattr(self.config, "save_steps", 1e6))),
gradient_accumulation_steps=int(self.config.gradient_accumulation_steps),
fp16=torch.cuda.is_available(),
dataloader_num_workers=int(getattr(self.config, "dataloader_num_workers", 0)),
save_total_limit=int(getattr(self.config, "save_total_limit", 1)),
load_best_model_at_end=bool(getattr(self.config, "load_best_model_at_end", False)),
# Adafactor uses ~4x less RAM than Adam — essential for Pegasus on CPU
optim="adafactor",
)
trainer = Trainer(
model=model_pegasus,
args=trainer_args,
data_collator=seq2seq_data_collator,
train_dataset=dataset_samsum_pt["train"],
eval_dataset=dataset_samsum_pt["validation"]
)
trainer.train()
model_pegasus.save_pretrained(os.path.join(self.config.root_dir, "pegasus-samsum-model"))
tokenizer.save_pretrained(os.path.join(self.config.root_dir, "tokenizer")) |