add multipack_sampler

This commit is contained in:
51616 2025-01-27 10:13:27 +00:00
parent 1665d89eee
commit 24e942a399
2 changed files with 128 additions and 4 deletions

View file

@ -319,6 +319,7 @@ def main():
ctx_tokenizer = tokenizer
if ctx_args.exp_setup == ExperimentSetup.HYPER_LORA:
# TODO: handle only extra_modules case (no target_modules)
logger.info("Using HyperLoRA")
if not ctx_args.from_pretrained_checkpoint:
hypernet_config = get_hypernet_config(
@ -510,6 +511,24 @@ def main():
out["chat_labels"] = chat_labels
return out
batch_sampler = None
if ctx_args.use_multipack_sampler:
from multipack_sampler.multipack_sampler import MultipackDistributedBatchSampler
"""
sampler = MultipackDistributedBatchSampler(
batch_max_length=batch_max_len,
lengths=lengths,
seed=0
)
dataloader = DataLoader(data, batch_sampler=sampler)
"""
lengths = np.array([len(x["ctx_ids"]) for x in train_ds])
batch_sampler = MultipackDistributedBatchSampler(
batch_max_length=2048, lengths=lengths, seed=training_args.seed
)
# TODO: use SFTTrainer instead? https://huggingface.co/docs/trl/en/sft_trainer
# TODO: use packing with SFTTrainer
@ -551,6 +570,7 @@ def main():
val_ds,
test_ds,
partial(train_collator, tokenizer=tokenizer),
train_batch_sampler=batch_sampler,
# partial(generation_collator, tokenizer=tokenizer),
compute_metrics=partial(
compute_metrics,

View file

@ -4,23 +4,119 @@ import logging
from collections import defaultdict
from dataclasses import fields
from enum import Enum
from typing import Optional
import datasets
import torch
import numpy as np
from rouge_score import rouge_scorer
import torch
from torch.utils.data import RandomSampler, DataLoader
from transformers import (
GenerationConfig,
Seq2SeqTrainer,
Seq2SeqTrainingArguments,
Trainer,
)
from transformers.trainer_utils import get_last_checkpoint
from transformers.trainer_utils import (
get_last_checkpoint,
has_length,
seed_worker,
)
from transformers.trainer_pt_utils import LengthGroupedSampler
from transformers.utils import is_datasets_available
TRAINING_TASK = Enum("TRAINING_TASK", ["CAUSAL_LM", "COMPLETION"])
logger = logging.getLogger()
# TODO: refactor to make this Trainer optional
class TrainerWithCustomSampler(Trainer):
def __init__(self, *args, **kwargs):
self.train_sampler = kwargs.pop("train_sampler", None)
self.train_batch_sampler = kwargs.pop("train_batch_sampler", None)
assert not (
self.train_sampler and self.train_batch_sampler
), "train_sampler and train_batch_sampler cannot be both provided"
super().__init__(*args, **kwargs)
# overriding to use custom sampler
def _get_train_sampler(self) -> Optional[torch.utils.data.Sampler]:
if self.train_dataset is None or not has_length(self.train_dataset):
return None
# Build the sampler.
if self.args.group_by_length:
if is_datasets_available() and isinstance(
self.train_dataset, datasets.Dataset
):
lengths = (
self.train_dataset[self.args.length_column_name]
if self.args.length_column_name in self.train_dataset.column_names
else None
)
else:
lengths = None
model_input_name = (
self.processing_class.model_input_names[0]
if self.processing_class is not None
else None
)
return LengthGroupedSampler(
self.args.train_batch_size * self.args.gradient_accumulation_steps,
dataset=self.train_dataset,
lengths=lengths,
model_input_name=model_input_name,
)
elif self.train_sampler:
return self.train_sampler
elif self.train_batch_sampler:
return self.train_batch_sampler
else:
return RandomSampler(self.train_dataset)
def get_train_dataloader(self) -> DataLoader:
"""
Returns the training [`~torch.utils.data.DataLoader`].
Will use no sampler if `train_dataset` does not implement `__len__`, a random sampler (adapted to distributed
training if necessary) otherwise.
Subclass and override this method if you want to inject some custom behavior.
"""
if self.train_dataset is None:
raise ValueError("Trainer: training requires a train_dataset.")
train_dataset = self.train_dataset
data_collator = self.data_collator
if is_datasets_available() and isinstance(train_dataset, datasets.Dataset):
train_dataset = self._remove_unused_columns(
train_dataset, description="training"
)
else:
data_collator = self._get_collator_with_removed_columns(
data_collator, description="training"
)
dataloader_params = {
# "batch_size": self._train_batch_size,
"collate_fn": data_collator,
"num_workers": self.args.dataloader_num_workers,
"pin_memory": self.args.dataloader_pin_memory,
"persistent_workers": self.args.dataloader_persistent_workers,
}
if not isinstance(train_dataset, torch.utils.data.IterableDataset):
# change "sampler" to "batch_sampler"
dataloader_params["batch_sampler"] = self._get_train_sampler()
# dataloader_params["drop_last"] = self.args.dataloader_drop_last
dataloader_params["worker_init_fn"] = seed_worker
dataloader_params["prefetch_factor"] = self.args.dataloader_prefetch_factor
return self.accelerator.prepare(DataLoader(train_dataset, **dataloader_params))
def clear_gpu():
gc.collect()
torch.cuda.empty_cache()
@ -38,6 +134,8 @@ def train_model(
train_collator=None,
# generation_collator=None,
compute_metrics=None,
train_sampler=None,
train_batch_sampler=None,
# preprocess_logits_for_metrics=None,
# max_new_tokens=2**13,
# gen_per_device_eval_batch_size=1,
@ -47,15 +145,21 @@ def train_model(
checkpoint = training_args.resume_from_checkpoint
logger.info(f"Resuming from the checkpoint: {checkpoint}")
trainer = Trainer(
trainer_cls = Trainer
trainer_kwargs = dict(
model=model,
args=training_args,
train_dataset=train_dataset,
eval_dataset=val_dataset,
data_collator=train_collator,
compute_metrics=compute_metrics,
# preprocess_logits_for_metrics=preprocess_logits_for_metrics,
)
if train_batch_sampler or train_sampler:
trainer_kwargs["train_sampler"] = train_sampler
trainer_kwargs["train_batch_sampler"] = train_batch_sampler
trainer_cls = TrainerWithCustomSampler
trainer = trainer_cls(**trainer_kwargs)
# Trainer loads the best model after training
# is done when load_best_model_at_end=True