[fix] Custom collate_fn for training

This commit is contained in:
theblackcat102
2023-02-03 06:08:01 +00:00
parent 1041564db7
commit 8b2080559c
3 changed files with 131 additions and 6 deletions
+3 -3
View File
@@ -8,7 +8,7 @@ import evaluate
import transformers
import yaml
from custom_datasets import get_one_dataset
from custom_datasets.dialogue_collator import DialogueDataCollator
from custom_datasets.dialogue_collator import DialogueDataCollator, TrainDialogueDataCollator
from custom_datasets.qa_datasets import QA_SPECIAL_TOKENS
from losses import CrossEntropyLoss, PolyLoss
from models import freeze_top_n_layers, get_specific_model
@@ -126,8 +126,8 @@ def get_dataset(conf, tokenizer):
train = ConcatDataset(train_datasets)
collate_fn = DialogueDataCollator(tokenizer, max_length=conf.max_length)
return train, evals, collate_fn
train_collate_fn = TrainDialogueDataCollator(tokenizer, max_length=conf.max_length)
return train, evals, collate_fn, train_collate_fn
def get_loss(loss, poly_eps):