Views
No views yet
bert-base-uncased ajustado con el dataset glue de MRPC.
Este es el código utilizado para realizar el ajuste:1# importamos las librerias
2from datasets import load_dataset
3
4from transformers import (AutoTokenizer,
5AutoModelForSequenceClassification,
6DataCollatorWithPadding,
7Trainer,
8TrainingArguments
9)
10
11from sklearn.metrics import accuracy_score, f1_score
12
13
14
15# cargamos el dataset
16glue_dataset = load_dataset('glue', 'mrpc')
17
18
19# definimos el modelo
20modelo = 'bert-base-uncased'
21
22
23
24# iniciamos el tokenizador, con este objeto se vectorizan las palabras
25tokenizador = AutoTokenizer.from_pretrained(modelo)
26
27
28# iniciamos el modelo BERT
29modelo_bert = AutoModelForSequenceClassification.from_pretrained(modelo, num_labels=2)
30
31
32# tokenizar el dataset
33token_dataset = glue_dataset.map(lambda x: tokenizador(x['sentence1'], x['sentence2'], truncation=True), batched=True)
34
35
36
37# iniciamos el data collator
38data_collator = DataCollatorWithPadding(tokenizer=tokenizador)
39
40
41# función evaluación
42def evaluacion(modelo_preds):
43
44 """
45 Función para obtener la métricas de evaluación.
46
47 Params:
48 + modelo_preds: transformers.trainer_utils.PredictionOutput, predicciones del modelo y etiquetas.
49
50 Return:
51 dict: diccionario con keys accuracy y f1-score y sus valores respectivos.
52 """
53
54 preds, etiquetas = modelo_preds
55
56 preds = np.argmax(preds, axis=-1)
57
58 return {'accuracy': accuracy_score(preds, etiquetas),
59 'f1': f1_score(preds, etiquetas)}
60
61
62
63# iniciamos los argumentos del entrenador
64args_entrenamiento = TrainingArguments(output_dir='../../training/glue-trainer',
65 evaluation_strategy='steps',
66 logging_steps=100,
67 )
68
69
70# iniciamos el entrenador
71entrenador = Trainer(model=modelo_bert,
72 args=args_entrenamiento,
73 train_dataset=token_dataset['train'],
74 eval_dataset=token_dataset['validation'],
75 data_collator=data_collator,
76 tokenizer=tokenizador,
77 compute_metrics=evaluacion
78 )
79
80# entrenamiento
81entrenador.train()
82
83
84# evaluación desde el entrenador
85print(entrenador.evaluate())