mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-23 17:01:04 +02:00
args per_device_train_max_batch_len for multipack_sampler
This commit is contained in:
parent
24e942a399
commit
4eaf0b5377
2 changed files with 9 additions and 1 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."},
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue