mirror of
https://github.com/wassname/Open-Assistant.git
synced 2026-09-10 11:41:04 +08:00
[feature] working trainer code
This commit is contained in:
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user