diff --git a/src/axolotl/core/trainer_builder.py b/src/axolotl/core/trainer_builder.py index cc7275184e..8d08d60b36 100644 --- a/src/axolotl/core/trainer_builder.py +++ b/src/axolotl/core/trainer_builder.py @@ -23,6 +23,7 @@ from torch.utils.data import BatchSampler, DataLoader, RandomSampler, SequentialSampler from transformers import ( EarlyStoppingCallback, + PreTrainedModel, Trainer, TrainerCallback, TrainingArguments, @@ -802,6 +803,15 @@ def push_to_hub(self, *args, **kwargs) -> str: return super().push_to_hub(*args, **kwargs) + def tokenize_row( + self, feature, model: Optional[Union[PreTrainedModel, torch.nn.Module]] = None + ) -> Dict: + res = super().tokenize_row(feature, model=model) + if self.tokenizer.bos_token_id is None and res["prompt_input_ids"][0] is None: + for key in res.keys(): + res[key] = res[key][1:] + return res + class TrainerBuilderBase(abc.ABC): """