From 4eaf0b5377155a79d14b5ff3b8260b0aa1650d06 Mon Sep 17 00:00:00 2001 From: 51616 Date: Mon, 27 Jan 2025 10:26:24 +0000 Subject: [PATCH] args per_device_train_max_batch_len for multipack_sampler --- intx_sft.py | 4 +++- src/ctx_to_lora/configs.py | 6 ++++++ 2 files changed, 9 insertions(+), 1 deletion(-) diff --git a/intx_sft.py b/intx_sft.py index 501ea22..7afd91f 100755 --- a/intx_sft.py +++ b/intx_sft.py @@ -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 diff --git a/src/ctx_to_lora/configs.py b/src/ctx_to_lora/configs.py index cab89b6..a5ff563 100644 --- a/src/ctx_to_lora/configs.py +++ b/src/ctx_to_lora/configs.py @@ -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."},