| import os |
| from transformers import DistilBertForSequenceClassification, Trainer, TrainingArguments |
| from data import prepare_data |
| from utils import MyDataset, compute_metrics |
|
|
| def train_model(): |
| |
| train_enc, train_lab, test_enc, test_lab, label2id = prepare_data() |
| train_dataset = MyDataset(train_enc, train_lab) |
| test_dataset = MyDataset(test_enc, test_lab) |
|
|
| model = DistilBertForSequenceClassification.from_pretrained( |
| 'distilbert-base-cased', |
| num_labels=len(label2id) |
| ) |
|
|
| |
| training_args = TrainingArguments( |
| output_dir='mlops-assignment2-model', |
| num_train_epochs=1, |
| per_device_train_batch_size=16, |
| logging_dir='./logs', |
| logging_steps=10, |
| eval_strategy='steps', |
| report_to='wandb', |
| push_to_hub=True, |
| hub_model_id='mlops-assignment2-distilbert', |
| hub_strategy='every_save' |
| ) |
|
|
| trainer = Trainer( |
| model=model, |
| args=training_args, |
| train_dataset=train_dataset, |
| eval_dataset=test_dataset, |
| compute_metrics=compute_metrics |
| ) |
| |
| trainer.train() |
| trainer.push_to_hub() |
| return trainer, test_dataset |
|
|
| if __name__ == '__main__': |
| train_model() |
|
|