From 24e942a399e04e1fec6394da5eb1d08c32d9cc62 Mon Sep 17 00:00:00 2001 From: 51616 Date: Mon, 27 Jan 2025 10:13:27 +0000 Subject: [PATCH] add multipack_sampler --- intx_sft.py | 20 ++++++ src/ctx_to_lora/training_utils.py | 112 ++++++++++++++++++++++++++++-- 2 files changed, 128 insertions(+), 4 deletions(-) diff --git a/intx_sft.py b/intx_sft.py index 33903cc..501ea22 100755 --- a/intx_sft.py +++ b/intx_sft.py @@ -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, diff --git a/src/ctx_to_lora/training_utils.py b/src/ctx_to_lora/training_utils.py index f56c593..135a3ea 100644 --- a/src/ctx_to_lora/training_utils.py +++ b/src/ctx_to_lora/training_utils.py @@ -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