UD03 · Notebook 3 — Modelos de lenguaje, DistilBERT¶
En esta sección se realiza un análisis de sentimiento de tweets, como un ejemplo del uso de modelos de lenguaje previamente entrenado. En este caso, se utiliza el modelo Distilbert para el análisis de sentimientos. Este modelo es una versión más ligera del modelo BERT, que es un modelo de lenguaje previo a la entrada que se ha utilizado con muy buenos resultados en diferentes tareas de procesamiento del lenguaje natural, como Análisis de sentimientos, Clasificación de texto o Extracción de información.
En este caso, se utiliza el modelo previamente entrenado para el análisis de sentimientos en inglés.
Carga del conjunto de datos`¶
Usaremos la librería datasets para cargar el dataset de los tweets. Esta La librería te permite cargar datasets de diferentes fuentes como Hugging Face Hub, Amazon AWS o Google Cloud. En este caso, cargaremos el dataset de tuits desde Hugging Face Hub.
# Instalamos las librerias que vamos a usar
import os
os.environ["WANDB_DISABLED"] = "true"
%pip install -U transformers datasets evaluate accelerate scikit-learn accuracy
Requirement already satisfied: transformers in /usr/local/lib/python3.12/dist-packages (4.57.3) Requirement already satisfied: datasets in /usr/local/lib/python3.12/dist-packages (4.0.0) Collecting datasets Downloading datasets-4.4.1-py3-none-any.whl.metadata (19 kB) Collecting evaluate Downloading evaluate-0.4.6-py3-none-any.whl.metadata (9.5 kB) Requirement already satisfied: accelerate in /usr/local/lib/python3.12/dist-packages (1.12.0) Requirement already satisfied: scikit-learn in /usr/local/lib/python3.12/dist-packages (1.6.1) Collecting scikit-learn Downloading scikit_learn-1.8.0-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl.metadata (11 kB) Collecting accuracy Downloading accuracy-0.1.1-py2.py3-none-any.whl.metadata (4.8 kB) Requirement already satisfied: filelock in /usr/local/lib/python3.12/dist-packages (from transformers) (3.20.0) Requirement already satisfied: huggingface-hub<1.0,>=0.34.0 in /usr/local/lib/python3.12/dist-packages (from transformers) (0.36.0) Requirement already satisfied: numpy>=1.17 in /usr/local/lib/python3.12/dist-packages (from transformers) (2.0.2) Requirement already satisfied: packaging>=20.0 in /usr/local/lib/python3.12/dist-packages (from transformers) (25.0) Requirement already satisfied: pyyaml>=5.1 in /usr/local/lib/python3.12/dist-packages (from transformers) (6.0.3) Requirement already satisfied: regex!=2019.12.17 in /usr/local/lib/python3.12/dist-packages (from transformers) (2025.11.3) Requirement already satisfied: requests in /usr/local/lib/python3.12/dist-packages (from transformers) (2.32.4) Requirement already satisfied: tokenizers<=0.23.0,>=0.22.0 in /usr/local/lib/python3.12/dist-packages (from transformers) (0.22.1) Requirement already satisfied: safetensors>=0.4.3 in /usr/local/lib/python3.12/dist-packages (from transformers) (0.7.0) Requirement already satisfied: tqdm>=4.27 in /usr/local/lib/python3.12/dist-packages (from transformers) (4.67.1) Collecting pyarrow>=21.0.0 (from datasets) Downloading pyarrow-22.0.0-cp312-cp312-manylinux_2_28_x86_64.whl.metadata (3.2 kB) Requirement already satisfied: dill<0.4.1,>=0.3.0 in /usr/local/lib/python3.12/dist-packages (from datasets) (0.3.8) Requirement already satisfied: pandas in /usr/local/lib/python3.12/dist-packages (from datasets) (2.2.2) Requirement already satisfied: httpx<1.0.0 in /usr/local/lib/python3.12/dist-packages (from datasets) (0.28.1) Requirement already satisfied: xxhash in /usr/local/lib/python3.12/dist-packages (from datasets) (3.6.0) Requirement already satisfied: multiprocess<0.70.19 in /usr/local/lib/python3.12/dist-packages (from datasets) (0.70.16) Requirement already satisfied: fsspec<=2025.10.0,>=2023.1.0 in /usr/local/lib/python3.12/dist-packages (from fsspec[http]<=2025.10.0,>=2023.1.0->datasets) (2025.3.0) Requirement already satisfied: psutil in /usr/local/lib/python3.12/dist-packages (from accelerate) (5.9.5) Requirement already satisfied: torch>=2.0.0 in /usr/local/lib/python3.12/dist-packages (from accelerate) (2.9.0+cu126) Requirement already satisfied: scipy>=1.10.0 in /usr/local/lib/python3.12/dist-packages (from scikit-learn) (1.16.3) Requirement already satisfied: joblib>=1.3.0 in /usr/local/lib/python3.12/dist-packages (from scikit-learn) (1.5.2) Requirement already satisfied: threadpoolctl>=3.2.0 in /usr/local/lib/python3.12/dist-packages (from scikit-learn) (3.6.0) Requirement already satisfied: Jinja2>=3.1.1 in /usr/local/lib/python3.12/dist-packages (from accuracy) (3.1.6) Requirement already satisfied: altair>=4.2.0 in /usr/local/lib/python3.12/dist-packages (from accuracy) (5.5.0) Collecting clumper>=0.2.15 (from accuracy) Downloading clumper-0.2.15-py2.py3-none-any.whl.metadata (1.2 kB) Requirement already satisfied: rich>=10.3.0 in /usr/local/lib/python3.12/dist-packages (from accuracy) (13.9.4) Requirement already satisfied: spacy>=3.0.0 in /usr/local/lib/python3.12/dist-packages (from accuracy) (3.8.11) Requirement already satisfied: typer>=0.3.0 in /usr/local/lib/python3.12/dist-packages (from accuracy) (0.20.0) Requirement already satisfied: jsonschema>=3.0 in /usr/local/lib/python3.12/dist-packages (from altair>=4.2.0->accuracy) (4.25.1) Requirement already satisfied: narwhals>=1.14.2 in /usr/local/lib/python3.12/dist-packages (from altair>=4.2.0->accuracy) (2.13.0) Requirement already satisfied: typing-extensions>=4.10.0 in /usr/local/lib/python3.12/dist-packages (from altair>=4.2.0->accuracy) (4.15.0) Requirement already satisfied: aiohttp!=4.0.0a0,!=4.0.0a1 in /usr/local/lib/python3.12/dist-packages (from fsspec[http]<=2025.10.0,>=2023.1.0->datasets) (3.13.2) Requirement already satisfied: anyio in /usr/local/lib/python3.12/dist-packages (from httpx<1.0.0->datasets) (4.12.0) Requirement already satisfied: certifi in /usr/local/lib/python3.12/dist-packages (from httpx<1.0.0->datasets) (2025.11.12) Requirement already satisfied: httpcore==1.* in /usr/local/lib/python3.12/dist-packages (from httpx<1.0.0->datasets) (1.0.9) Requirement already satisfied: idna in /usr/local/lib/python3.12/dist-packages (from httpx<1.0.0->datasets) (3.11) Requirement already satisfied: h11>=0.16 in /usr/local/lib/python3.12/dist-packages (from httpcore==1.*->httpx<1.0.0->datasets) (0.16.0) Requirement already satisfied: hf-xet<2.0.0,>=1.1.3 in /usr/local/lib/python3.12/dist-packages (from huggingface-hub<1.0,>=0.34.0->transformers) (1.2.0) Requirement already satisfied: MarkupSafe>=2.0 in /usr/local/lib/python3.12/dist-packages (from Jinja2>=3.1.1->accuracy) (3.0.3) Requirement already satisfied: charset_normalizer<4,>=2 in /usr/local/lib/python3.12/dist-packages (from requests->transformers) (3.4.4) Requirement already satisfied: urllib3<3,>=1.21.1 in /usr/local/lib/python3.12/dist-packages (from requests->transformers) (2.5.0) Requirement already satisfied: markdown-it-py>=2.2.0 in /usr/local/lib/python3.12/dist-packages (from rich>=10.3.0->accuracy) (4.0.0) Requirement already satisfied: pygments<3.0.0,>=2.13.0 in /usr/local/lib/python3.12/dist-packages (from rich>=10.3.0->accuracy) (2.19.2) Requirement already satisfied: spacy-legacy<3.1.0,>=3.0.11 in /usr/local/lib/python3.12/dist-packages (from spacy>=3.0.0->accuracy) (3.0.12) Requirement already satisfied: spacy-loggers<2.0.0,>=1.0.0 in /usr/local/lib/python3.12/dist-packages (from spacy>=3.0.0->accuracy) (1.0.5) Requirement already satisfied: murmurhash<1.1.0,>=0.28.0 in /usr/local/lib/python3.12/dist-packages (from spacy>=3.0.0->accuracy) (1.0.15) Requirement already satisfied: cymem<2.1.0,>=2.0.2 in /usr/local/lib/python3.12/dist-packages (from spacy>=3.0.0->accuracy) (2.0.13) Requirement already satisfied: preshed<3.1.0,>=3.0.2 in /usr/local/lib/python3.12/dist-packages (from spacy>=3.0.0->accuracy) (3.0.12) Requirement already satisfied: thinc<8.4.0,>=8.3.4 in /usr/local/lib/python3.12/dist-packages (from spacy>=3.0.0->accuracy) (8.3.10) Requirement already satisfied: wasabi<1.2.0,>=0.9.1 in /usr/local/lib/python3.12/dist-packages (from spacy>=3.0.0->accuracy) (1.1.3) Requirement already satisfied: srsly<3.0.0,>=2.4.3 in /usr/local/lib/python3.12/dist-packages (from spacy>=3.0.0->accuracy) (2.5.2) Requirement already satisfied: catalogue<2.1.0,>=2.0.6 in /usr/local/lib/python3.12/dist-packages (from spacy>=3.0.0->accuracy) (2.0.10) Requirement already satisfied: weasel<0.5.0,>=0.4.2 in /usr/local/lib/python3.12/dist-packages (from spacy>=3.0.0->accuracy) (0.4.3) Requirement already satisfied: typer-slim<1.0.0,>=0.3.0 in /usr/local/lib/python3.12/dist-packages (from spacy>=3.0.0->accuracy) (0.20.0) Requirement already satisfied: pydantic!=1.8,!=1.8.1,<3.0.0,>=1.7.4 in /usr/local/lib/python3.12/dist-packages (from spacy>=3.0.0->accuracy) (2.12.3) Requirement already satisfied: setuptools in /usr/local/lib/python3.12/dist-packages (from spacy>=3.0.0->accuracy) (75.2.0) Requirement already satisfied: sympy>=1.13.3 in /usr/local/lib/python3.12/dist-packages (from torch>=2.0.0->accelerate) (1.14.0) Requirement already satisfied: networkx>=2.5.1 in /usr/local/lib/python3.12/dist-packages (from torch>=2.0.0->accelerate) (3.6.1) Requirement already satisfied: nvidia-cuda-nvrtc-cu12==12.6.77 in /usr/local/lib/python3.12/dist-packages (from torch>=2.0.0->accelerate) (12.6.77) Requirement already satisfied: nvidia-cuda-runtime-cu12==12.6.77 in /usr/local/lib/python3.12/dist-packages (from torch>=2.0.0->accelerate) (12.6.77) Requirement already satisfied: nvidia-cuda-cupti-cu12==12.6.80 in /usr/local/lib/python3.12/dist-packages (from torch>=2.0.0->accelerate) (12.6.80) Requirement already satisfied: nvidia-cudnn-cu12==9.10.2.21 in /usr/local/lib/python3.12/dist-packages (from torch>=2.0.0->accelerate) (9.10.2.21) Requirement already satisfied: nvidia-cublas-cu12==12.6.4.1 in /usr/local/lib/python3.12/dist-packages (from torch>=2.0.0->accelerate) (12.6.4.1) Requirement already satisfied: nvidia-cufft-cu12==11.3.0.4 in /usr/local/lib/python3.12/dist-packages (from torch>=2.0.0->accelerate) (11.3.0.4) Requirement already satisfied: nvidia-curand-cu12==10.3.7.77 in /usr/local/lib/python3.12/dist-packages (from torch>=2.0.0->accelerate) (10.3.7.77) Requirement already satisfied: nvidia-cusolver-cu12==11.7.1.2 in /usr/local/lib/python3.12/dist-packages (from torch>=2.0.0->accelerate) (11.7.1.2) Requirement already satisfied: nvidia-cusparse-cu12==12.5.4.2 in /usr/local/lib/python3.12/dist-packages (from torch>=2.0.0->accelerate) (12.5.4.2) Requirement already satisfied: nvidia-cusparselt-cu12==0.7.1 in /usr/local/lib/python3.12/dist-packages (from torch>=2.0.0->accelerate) (0.7.1) Requirement already satisfied: nvidia-nccl-cu12==2.27.5 in /usr/local/lib/python3.12/dist-packages (from torch>=2.0.0->accelerate) (2.27.5) Requirement already satisfied: nvidia-nvshmem-cu12==3.3.20 in /usr/local/lib/python3.12/dist-packages (from torch>=2.0.0->accelerate) (3.3.20) Requirement already satisfied: nvidia-nvtx-cu12==12.6.77 in /usr/local/lib/python3.12/dist-packages (from torch>=2.0.0->accelerate) (12.6.77) Requirement already satisfied: nvidia-nvjitlink-cu12==12.6.85 in /usr/local/lib/python3.12/dist-packages (from torch>=2.0.0->accelerate) (12.6.85) Requirement already satisfied: nvidia-cufile-cu12==1.11.1.6 in /usr/local/lib/python3.12/dist-packages (from torch>=2.0.0->accelerate) (1.11.1.6) Requirement already satisfied: triton==3.5.0 in /usr/local/lib/python3.12/dist-packages (from torch>=2.0.0->accelerate) (3.5.0) Requirement already satisfied: click>=8.0.0 in /usr/local/lib/python3.12/dist-packages (from typer>=0.3.0->accuracy) (8.3.1) Requirement already satisfied: shellingham>=1.3.0 in /usr/local/lib/python3.12/dist-packages (from typer>=0.3.0->accuracy) (1.5.4) Requirement already satisfied: python-dateutil>=2.8.2 in /usr/local/lib/python3.12/dist-packages (from pandas->datasets) (2.9.0.post0) Requirement already satisfied: pytz>=2020.1 in /usr/local/lib/python3.12/dist-packages (from pandas->datasets) (2025.2) Requirement already satisfied: tzdata>=2022.7 in /usr/local/lib/python3.12/dist-packages (from pandas->datasets) (2025.2) Requirement already satisfied: aiohappyeyeballs>=2.5.0 in /usr/local/lib/python3.12/dist-packages (from aiohttp!=4.0.0a0,!=4.0.0a1->fsspec[http]<=2025.10.0,>=2023.1.0->datasets) (2.6.1) Requirement already satisfied: aiosignal>=1.4.0 in /usr/local/lib/python3.12/dist-packages (from aiohttp!=4.0.0a0,!=4.0.0a1->fsspec[http]<=2025.10.0,>=2023.1.0->datasets) (1.4.0) Requirement already satisfied: attrs>=17.3.0 in /usr/local/lib/python3.12/dist-packages (from aiohttp!=4.0.0a0,!=4.0.0a1->fsspec[http]<=2025.10.0,>=2023.1.0->datasets) (25.4.0) Requirement already satisfied: frozenlist>=1.1.1 in /usr/local/lib/python3.12/dist-packages (from aiohttp!=4.0.0a0,!=4.0.0a1->fsspec[http]<=2025.10.0,>=2023.1.0->datasets) (1.8.0) Requirement already satisfied: multidict<7.0,>=4.5 in /usr/local/lib/python3.12/dist-packages (from aiohttp!=4.0.0a0,!=4.0.0a1->fsspec[http]<=2025.10.0,>=2023.1.0->datasets) (6.7.0) Requirement already satisfied: propcache>=0.2.0 in /usr/local/lib/python3.12/dist-packages (from aiohttp!=4.0.0a0,!=4.0.0a1->fsspec[http]<=2025.10.0,>=2023.1.0->datasets) (0.4.1) Requirement already satisfied: yarl<2.0,>=1.17.0 in /usr/local/lib/python3.12/dist-packages (from aiohttp!=4.0.0a0,!=4.0.0a1->fsspec[http]<=2025.10.0,>=2023.1.0->datasets) (1.22.0) Requirement already satisfied: jsonschema-specifications>=2023.03.6 in /usr/local/lib/python3.12/dist-packages (from jsonschema>=3.0->altair>=4.2.0->accuracy) (2025.9.1) Requirement already satisfied: referencing>=0.28.4 in /usr/local/lib/python3.12/dist-packages (from jsonschema>=3.0->altair>=4.2.0->accuracy) (0.37.0) Requirement already satisfied: rpds-py>=0.7.1 in /usr/local/lib/python3.12/dist-packages (from jsonschema>=3.0->altair>=4.2.0->accuracy) (0.30.0) Requirement already satisfied: mdurl~=0.1 in /usr/local/lib/python3.12/dist-packages (from markdown-it-py>=2.2.0->rich>=10.3.0->accuracy) (0.1.2) Requirement already satisfied: annotated-types>=0.6.0 in /usr/local/lib/python3.12/dist-packages (from pydantic!=1.8,!=1.8.1,<3.0.0,>=1.7.4->spacy>=3.0.0->accuracy) (0.7.0) Requirement already satisfied: pydantic-core==2.41.4 in /usr/local/lib/python3.12/dist-packages (from pydantic!=1.8,!=1.8.1,<3.0.0,>=1.7.4->spacy>=3.0.0->accuracy) (2.41.4) Requirement already satisfied: typing-inspection>=0.4.2 in /usr/local/lib/python3.12/dist-packages (from pydantic!=1.8,!=1.8.1,<3.0.0,>=1.7.4->spacy>=3.0.0->accuracy) (0.4.2) Requirement already satisfied: six>=1.5 in /usr/local/lib/python3.12/dist-packages (from python-dateutil>=2.8.2->pandas->datasets) (1.17.0) Requirement already satisfied: mpmath<1.4,>=1.1.0 in /usr/local/lib/python3.12/dist-packages (from sympy>=1.13.3->torch>=2.0.0->accelerate) (1.3.0) Requirement already satisfied: blis<1.4.0,>=1.3.0 in /usr/local/lib/python3.12/dist-packages (from thinc<8.4.0,>=8.3.4->spacy>=3.0.0->accuracy) (1.3.3) Requirement already satisfied: confection<1.0.0,>=0.0.1 in /usr/local/lib/python3.12/dist-packages (from thinc<8.4.0,>=8.3.4->spacy>=3.0.0->accuracy) (0.1.5) Requirement already satisfied: cloudpathlib<1.0.0,>=0.7.0 in /usr/local/lib/python3.12/dist-packages (from weasel<0.5.0,>=0.4.2->spacy>=3.0.0->accuracy) (0.23.0) Requirement already satisfied: smart-open<8.0.0,>=5.2.1 in /usr/local/lib/python3.12/dist-packages (from weasel<0.5.0,>=0.4.2->spacy>=3.0.0->accuracy) (7.5.0) Requirement already satisfied: wrapt in /usr/local/lib/python3.12/dist-packages (from smart-open<8.0.0,>=5.2.1->weasel<0.5.0,>=0.4.2->spacy>=3.0.0->accuracy) (2.0.1) Downloading datasets-4.4.1-py3-none-any.whl (511 kB) ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 511.6/511.6 kB 21.7 MB/s eta 0:00:00 Downloading evaluate-0.4.6-py3-none-any.whl (84 kB) ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 84.1/84.1 kB 6.4 MB/s eta 0:00:00 Downloading scikit_learn-1.8.0-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl (8.9 MB) ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 8.9/8.9 MB 62.6 MB/s eta 0:00:00 Downloading accuracy-0.1.1-py2.py3-none-any.whl (7.8 kB) Downloading clumper-0.2.15-py2.py3-none-any.whl (18 kB) Downloading pyarrow-22.0.0-cp312-cp312-manylinux_2_28_x86_64.whl (47.7 MB) ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 47.7/47.7 MB 12.1 MB/s eta 0:00:00 Installing collected packages: clumper, pyarrow, scikit-learn, datasets, evaluate, accuracy Attempting uninstall: pyarrow Found existing installation: pyarrow 18.1.0 Uninstalling pyarrow-18.1.0: Successfully uninstalled pyarrow-18.1.0 Attempting uninstall: scikit-learn Found existing installation: scikit-learn 1.6.1 Uninstalling scikit-learn-1.6.1: Successfully uninstalled scikit-learn-1.6.1 Attempting uninstall: datasets Found existing installation: datasets 4.0.0 Uninstalling datasets-4.0.0: Successfully uninstalled datasets-4.0.0 Successfully installed accuracy-0.1.1 clumper-0.2.15 datasets-4.4.1 evaluate-0.4.6 pyarrow-22.0.0 scikit-learn-1.8.0
import datasets
# Cargamos el Dataset
dataset = datasets.load_dataset('dair-ai/emotion')
# Mostramos los datos de ejemplo
dataset['train'][0]
/usr/local/lib/python3.12/dist-packages/huggingface_hub/utils/_auth.py:94: UserWarning: The secret `HF_TOKEN` does not exist in your Colab secrets. To authenticate with the Hugging Face Hub, create a token in your settings tab (https://huggingface.co/settings/tokens), set it as secret in your Google Colab and restart your session. You will be able to reuse this secret in all of your notebooks. Please note that authentication is recommended but still optional to access public models or datasets. warnings.warn(
README.md: 0.00B [00:00, ?B/s]
split/train-00000-of-00001.parquet: 0%| | 0.00/1.03M [00:00<?, ?B/s]
split/validation-00000-of-00001.parquet: 0%| | 0.00/127k [00:00<?, ?B/s]
split/test-00000-of-00001.parquet: 0%| | 0.00/129k [00:00<?, ?B/s]
Generating train split: 0%| | 0/16000 [00:00<?, ? examples/s]
Generating validation split: 0%| | 0/2000 [00:00<?, ? examples/s]
Generating test split: 0%| | 0/2000 [00:00<?, ? examples/s]
{'text': 'i didnt feel humiliated', 'label': 0} Podemos ver que cada registro del dataset contiene el texto del tweet y el sentimiento asociado. En este caso, el sentimiento está codificado con un entero entre 0 y 5, donde 0 corresponde a la tristeza, 1 a Alegria, 2 para amor, 3 a la ira, 4 para miedo y 5 para sorpresa.
Preparación del dataset¶
En este caso, el dataset ya está dividido en conjuntos de entrenamiento, prueba y validación. El siguiente paso es preparar el dataset para el entrenamiento del modelo. En este caso, el modelo que utilizaremos es el modelo BERT. Este modelo requiere que el texto sea tokenizado y los tokens estén codificados con sus identificadores numéricos correspondientes. Para hacer esto, utilizaremos un tokenizador DistilBERT previamente entrenado.
# importamos el tokenizador de DistilBERT
from transformers import AutoTokenizer
# Cargamos el tokenizador
tokenizer = AutoTokenizer.from_pretrained('distilbert/distilbert-base-uncased')
# Mostramos un ejemplo de tokenización
tokenizer.tokenize('FC Barcelona is fucked this year')
tokenizer_config.json: 0%| | 0.00/48.0 [00:00<?, ?B/s]
config.json: 0%| | 0.00/483 [00:00<?, ?B/s]
vocab.txt: 0.00B [00:00, ?B/s]
tokenizer.json: 0.00B [00:00, ?B/s]
['fc', 'barcelona', 'is', 'fucked', 'this', 'year']
# Definimos una función para preprocesar el texto.
# Truncamos los textos para asegurarnos de que no excedan el máximo tamaño de entrada deDistilBert
def tokenize(examples):
return tokenizer(examples["text"], padding='max_length', truncation=True)
Para aplicar la tokenización, usaremos la función map de datasets. Esta función permite aplicar una función a cada registro del dataset. En este caso, la función que aplicaremos es la función tokenize que hemos definido anteriormente. Nosotros también usaremos batched=True para indicar que la función se aplicará a todo dataset en bloques.
dades_tokenitzades = dataset.map(tokenize, batched=True)
Map: 0%| | 0/16000 [00:00<?, ? examples/s]
Map: 0%| | 0/2000 [00:00<?, ? examples/s]
Map: 0%| | 0/2000 [00:00<?, ? examples/s]
Evaluación¶
Para evaluar el modelo, debemos cargar el método que utilizaremos para la evaluación. En este caso usaremos la métrica accuracy del módulo evaluate de HuggingFace.
También definiremos una función para calcular las métricas del modelo. Esta función se utilizará para evaluar el modelo después de cada época.
import evaluate
accuracy = evaluate.load('accuracy')
Downloading builder script: 0.00B [00:00, ?B/s]
# Definimos una función para calcular la precisión del modelo
def compute_metrics(eval_pred):
predictions, labels = eval_pred
predictions = predictions.argmax(axis=1)
return accuracy.compute(predictions=predictions, references=labels)
Definición de etiquetas¶
Antes de entrenar el modelo, debemos crear un diccionario que traduzca los identificadores numéricos del sentimiento a sus etiquetas correspondientes y viceversa.
id_a_etiqueta = {
0: "SADNESS",
1: "JOY",
2: "LOVE",
3: "ANGER",
4: "FEAR",
5: "SUPRISE"
}
etiqueta_a_id = {
"SADNESS": 0,
"JOY": 1,
"LOVE": 2,
"ANGER": 3,
"FEAR": 4,
"SUPRISE": 5
}
Fine tuning del modelo¶
El proceso de ajuste fino del modelo es entrenar el modelo con nuestro dataset. Esto permite que el modelo se adapte mejor a nuestros datos y mejore su rendimiento.
Necesitamos definir la función de optimización, el tamaño de los bloques y el número de épocas.
BATCH_SIZE = 16
NUM_EPOCHS = 1
Ahora podemos cargar el modelo previamente entrenado y hacer el fine tuning. Usaremos AutoModelForSequenceClassification y agregaremos las etiquetas que hemos definido previamente.
from transformers import AutoModelForSequenceClassification
model = AutoModelForSequenceClassification.from_pretrained(
'distilbert/distilbert-base-uncased',
num_labels=len(etiqueta_a_id),
id2label=id_a_etiqueta,
label2id=etiqueta_a_id
)
model.safetensors: 0%| | 0.00/268M [00:00<?, ?B/s]
Some weights of DistilBertForSequenceClassification were not initialized from the model checkpoint at distilbert/distilbert-base-uncased and are newly initialized: ['classifier.bias', 'classifier.weight', 'pre_classifier.bias', 'pre_classifier.weight'] You should probably TRAIN this model on a down-stream task to be able to use it for predictions and inference.
Para evaluar el modelo debemos definir un objeto TrainingArguments con los parámetros de entrenamiento. Podemos incluir el número de épocas, el tamaño de los bloques, el tamaño del lote, la tasa de aprendizaje, etc.
from transformers import TrainingArguments, Trainer
training_args = TrainingArguments(
output_dir="test_trainer",
eval_strategy="epoch",
per_device_train_batch_size=BATCH_SIZE,
per_device_eval_batch_size=BATCH_SIZE,
num_train_epochs=NUM_EPOCHS,
)
Using the `WANDB_DISABLED` environment variable is deprecated and will be removed in v5. Use the --report_to flag to control the integrations used for logging result (for instance --report_to none).
Ahora podemos entrenar el modelo, utilizando el trainer.
trainer = Trainer(
model=model,
args=training_args,
train_dataset=dades_tokenitzades['train'],
eval_dataset=dades_tokenitzades['validation'],
compute_metrics=compute_metrics,
)
trainer.train()
| Epoch | Training Loss | Validation Loss |
|---|
| Epoch | Training Loss | Validation Loss | Accuracy |
|---|---|---|---|
| 1 | 0.205400 | 0.179499 | 0.927000 |
TrainOutput(global_step=1000, training_loss=0.379160888671875, metrics={'train_runtime': 728.165, 'train_samples_per_second': 21.973, 'train_steps_per_second': 1.373, 'total_flos': 2119629570048000.0, 'train_loss': 0.379160888671875, 'epoch': 1.0}) Inferencia¶
Para hacer inferencia usando el modelo, crearemos una pipeline Huggingface. Esta pipeline utilizará el modelo y el tokenizador que importamos anteriormente.
Luego usaremos la pipeline para hacer una inferencia con un texto de ejemplo.
from transformers import pipeline
classifier = pipeline("sentiment-analysis", model=model, tokenizer=tokenizer)
print(classifier("Suddenly, I'm not half the man I used to be, There's a shadow hanging over me, Oh, yesterday came suddenly"))
print(classifier("Don't stop me now. I'm havin' such a good time, I'm havin' a ball. If you wanna have a good time, just give me a call"))
print(classifier("Remember those who win the game. Lose the love they sought to gain. In debentures of quality. And dubious integrity. Their small town eyes will gape at you. In dull surprise when payment due. Exceeds accounts received. At seventeen"))
print(classifier("You got your bitches with the silicone injections. Crystal meth and yeast infections. Bleached blond hair, collagen lip injections. Who are you to criticize my intentions?. Got your subtle, manipulative devices. Just like you, I got my vices. I got a thought that would be nice. I'd like to crush your head, tight in my vice. Pain"))
Device set to use cuda:0
[{'label': 'ANGER', 'score': 0.34769198298454285}]
[{'label': 'JOY', 'score': 0.9859300255775452}]
[{'label': 'ANGER', 'score': 0.48408079147338867}]
[{'label': 'ANGER', 'score': 0.9862304329872131}]