[feature] working trainer code

This commit is contained in:
theblackcat102
2022-12-31 03:02:10 +00:00
parent ad98a28241
commit bcd5c52b3b
6 changed files with 174 additions and 24 deletions
+102 -2
View File
@@ -1,2 +1,102 @@
import wandb
from accelerate import Accelerator
from typing import Callable, List, Optional, Tuple, Union, Dict
import torch
from torch import nn
import numpy as np
import evaluate
from dataclasses import dataclass
from torch.utils.data import Dataset
from transformers import AutoModelForSequenceClassification, AutoModelForMultipleChoice
from transformers import Trainer, PreTrainedModel, TrainingArguments, DataCollator, EvalPrediction, TrainerCallback, PreTrainedTokenizerBase
from rank_datasets import DataCollatorForPairRank, WebGPT
from utils import get_tokenizer, train_val_dataset
accuracy = evaluate.load("accuracy")
@dataclass
class CustomTrainingArguments(TrainingArguments):
loss_function: str='rank'
def compute_metrics(eval_pred):
predictions, _ = eval_pred
predictions = np.argmax(predictions, axis=1)
return accuracy.compute(predictions=predictions, references=[0]*predictions.shape[0])
class RankLoss(nn.Module):
def __init__(self, eps=1e-8) -> None:
super().__init__()
self.eps = eps
self.log_sigmoid = nn.LogSigmoid()
def forward(self, pos, neg):
return -self.log_sigmoid(pos - neg + self.eps).mean()
class RankTrainer(Trainer):
def __init__(self, model: Union[PreTrainedModel, nn.Module] = None,
args: TrainingArguments = None,
data_collator: Optional[DataCollator] = None,
train_dataset: Optional[Dataset] = None,
eval_dataset: Optional[Dataset] = None,
tokenizer: Optional[PreTrainedTokenizerBase] = None,
model_init: Callable[[], PreTrainedModel] = None,
compute_metrics: Optional[Callable[[EvalPrediction], Dict]] = None,
callbacks: Optional[List[TrainerCallback]] = None,
optimizers: Tuple[torch.optim.Optimizer, torch.optim.lr_scheduler.LambdaLR] = (None, None),
preprocess_logits_for_metrics: Callable[[torch.Tensor, torch.Tensor], torch.Tensor] = None):
super().__init__(model, args, data_collator, train_dataset, eval_dataset, tokenizer,
model_init, compute_metrics, callbacks, optimizers, preprocess_logits_for_metrics)
self.loss_fct = RankLoss() if args.loss_function == 'rank' else nn.CrossEntropyLoss()
self.loss_function = args.loss_function
def compute_loss(self, model, inputs, return_outputs=False):
# forward pass
outputs = model(**inputs)
logits = outputs.get("logits").view(-1, 2)
if self.loss_function == 'rank':
loss = self.loss_fct(logits[:, 0], logits[:, 1])
else:
loss = self.loss_fct(logits, torch.zeros(logits.shape[0], device=logits.device, dtype=torch.long))
return (loss, outputs) if return_outputs else loss
if __name__ == "__main__":
model_name = 'bigscience/bloomz-560m'
model_name = 'google/electra-base-discriminator'
model = AutoModelForSequenceClassification.from_pretrained(model_name, num_labels=1, problem_type='regression')
tokenizer = get_tokenizer(model_name)
args = CustomTrainingArguments(
output_dir=f"outputs/{model_name}-finetuned",
fp16=True,
num_train_epochs=4,
warmup_steps=500,
learning_rate=3e-5,
# half_precision_backend="apex",
gradient_checkpointing=False,
gradient_accumulation_steps=6,
per_device_train_batch_size=12,
per_device_eval_batch_size=5,
weight_decay=0.01,
max_grad_norm=2.0,
logging_steps=10,
save_total_limit=4,
evaluation_strategy='steps',
loss_function='rank',
eval_steps=500,
save_steps=1000,
report_to="wandb",
run_name='reward-model'
)
dataset = WebGPT()
train, eval = train_val_dataset(dataset)
collate_fn = DataCollatorForPairRank(tokenizer, max_length=400)
trainer = RankTrainer(
model,
args,
train_dataset=train,
eval_dataset=eval,
data_collator=collate_fn,
tokenizer=tokenizer
)
trainer.train()