fix SNAC training crash on variable-length sequences with DataCollatorForSeq2Seq

This commit is contained in:
Manan17 2026-03-05 07:03:38 +00:00
commit c723f8d4da
2 changed files with 7 additions and 1 deletions

@ -0,0 +1 @@
Subproject commit 59d896747aa0a6a207837e7da2d6921805eae684

View file

@ -2087,12 +2087,17 @@ class UnslothTrainer:
elif self._audio_type == 'snac':
# Orpheus: language model with SNAC codec tokens — plain HF Trainer
from transformers import Trainer as HFTrainer, TrainingArguments
# DataCollatorForSeq2Seq dynamically pads variable-length sequences per batch
# (text + audio codes vary in length) and pads labels with -100.
from transformers import Trainer as HFTrainer, TrainingArguments, DataCollatorForSeq2Seq
config = self._build_audio_training_args(training_args, output_dir)
self.trainer = HFTrainer(
model=self.model, train_dataset=dataset,
args=TrainingArguments(**config),
data_collator=DataCollatorForSeq2Seq(
tokenizer=self.tokenizer, padding=True, pad_to_multiple_of=8,
),
)
self.trainer.add_callback(self._create_progress_callback())