remove multipack_sampler + fix name collision

This commit is contained in:
51616 2025-02-25 11:34:39 +00:00
parent c9b20890cb
commit 7866a43fa4
3 changed files with 2 additions and 117 deletions

View file

@ -256,8 +256,9 @@ def main():
# should be the same across processes
# still possible to have a name crash though
# logging_dir is just "runs/DATE_TIME_HOSTNAME"
slurm_job_id = f"_{os.getenv('SLURM_JOB_ID')}" if os.getenv("SLURM_JOB_ID") else ""
run_name = (
get_run_name(seed_str=training_args.logging_dir.strip("runs/"))
get_run_name(seed_str=training_args.logging_dir.strip("runs/") + slurm_job_id)
if not checkpoint_dir
else checkpoint_dir.strip("/").split("/")[-2]
)
@ -506,26 +507,6 @@ def main():
else partial(train_collator, tokenizer=tokenizer)
)
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=ctx_args.per_device_train_max_batch_len,
lengths=lengths,
seed=training_args.seed,
)
# TODO: use SFTTrainer instead? https://huggingface.co/docs/trl/en/sft_trainer
# TODO: use packing with SFTTrainer
@ -569,7 +550,6 @@ def main():
val_ds,
test_ds,
train_collator,
train_batch_sampler=batch_sampler,
# partial(generation_collator, tokenizer=tokenizer),
compute_metrics=partial(
compute_metrics,

View file

@ -292,10 +292,6 @@ class CtxTrainingArguments:
default=2**13,
metadata={"help": "Maximum context length for training."},
)
use_multipack_sampler: bool = field(
default=False,
metadata={"help": "Whether to use multipack sampler."},
)
use_sequence_packing: bool = field(
default=False,
metadata={"help": "Whether to use sequence packing."},

View file

@ -31,92 +31,6 @@ 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()
@ -135,7 +49,6 @@ def train_model(
# 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,
@ -154,10 +67,6 @@ def train_model(
data_collator=train_collator,
compute_metrics=compute_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)