File size: 1,626 Bytes
b557902 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 | from __future__ import annotations
from dataclasses import dataclass
from typing import Dict, Optional
from datasets import DatasetDict
from transformers import (
DataCollatorWithPadding,
RobertaForSequenceClassification,
RobertaTokenizerFast,
)
@dataclass(frozen=True)
class ModelConfig:
model_name: str = "roberta-base"
max_length: int = 256
text_field: str = "statement"
def get_tokenizer(model_name: str) -> RobertaTokenizerFast:
return RobertaTokenizerFast.from_pretrained(model_name)
def get_model(
model_name: str,
num_labels: int,
id2label: Dict[int, str],
label2id: Dict[str, int],
) -> RobertaForSequenceClassification:
return RobertaForSequenceClassification.from_pretrained(
model_name,
num_labels=num_labels,
id2label=id2label,
label2id=label2id,
problem_type="single_label_classification",
)
def get_data_collator(tokenizer: RobertaTokenizerFast) -> DataCollatorWithPadding:
return DataCollatorWithPadding(tokenizer=tokenizer)
def tokenize_dataset(
dataset: DatasetDict,
tokenizer: RobertaTokenizerFast,
*,
max_length: int,
text_field: str,
remove_columns: Optional[list] = None,
) -> DatasetDict:
def _tokenize(batch):
return tokenizer(
batch[text_field],
truncation=True,
max_length=max_length,
)
if remove_columns is None:
remove_columns = [
col for col in dataset["train"].column_names if col != "labels"
]
return dataset.map(_tokenize, batched=True, remove_columns=remove_columns)
|