args per_device_train_max_batch_len for multipack_sampler

This commit is contained in:
51616 2025-01-27 10:26:24 +00:00
parent 24e942a399
commit 4eaf0b5377
2 changed files with 9 additions and 1 deletions

View file

@ -526,7 +526,9 @@ def main():
"""
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
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

View file

@ -284,6 +284,12 @@ class CtxTrainingArguments:
default=False,
metadata={"help": "Whether to use multipack sampler."},
)
per_device_train_max_batch_len: Optional[int] = field(
default=2**12,
metadata={
"help": "Maximum batch length for training. Only used with multipack sampler."
},
)
max_new_tokens: Optional[int] = field(
default=2**10,
metadata={"help": "Maximum new tokens for generation-based evaluation."},