UD03 · Notebook 2 — Clasificador de noticias¶
En esta práctica, crearemos un clasificador de noticias utilizando las técnicas de procesamiento del lenguaje natural que hemos visto en clase, centrándose en la representación del texto.
Usaremos el dataset AG News que contiene 1,000,000 noticias de 4 categorías diferentes.
Dataset¶
Para cargar el conjunto de datos, usaremos la librería datasets. Esta librería nos permitirá cargar muchos conjuntos de datos diferentes de una manera simple.En este caso, cargaremos el conjunto de datos AG News.
Preparación del dataset¶
Para instalar las librerías necesarias, ejecutaremos la siguiente celda.
Usaremos pytorch (una libreria de deep learning), pipeline (una libreria de tratamiento de datos), scikit-learn (una libreria de machine learning) y transformers (una libreria de modelos de lenguaje).
# Instalamos las librerías necesarias en las versiones correctas
%pip install --upgrade torch datasets scikit-learn transformers gensim
Requirement already satisfied: torch in /usr/local/lib/python3.12/dist-packages (2.9.0+cpu) Collecting torch Downloading torch-2.9.1-cp312-cp312-manylinux_2_28_x86_64.whl.metadata (30 kB) 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) 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) Requirement already satisfied: transformers in /usr/local/lib/python3.12/dist-packages (4.57.3) Collecting gensim Downloading gensim-4.4.0-cp312-cp312-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl.metadata (8.4 kB) Requirement already satisfied: filelock in /usr/local/lib/python3.12/dist-packages (from torch) (3.20.0) Requirement already satisfied: typing-extensions>=4.10.0 in /usr/local/lib/python3.12/dist-packages (from torch) (4.15.0) Requirement already satisfied: setuptools in /usr/local/lib/python3.12/dist-packages (from torch) (75.2.0) Requirement already satisfied: sympy>=1.13.3 in /usr/local/lib/python3.12/dist-packages (from torch) (1.14.0) Requirement already satisfied: networkx>=2.5.1 in /usr/local/lib/python3.12/dist-packages (from torch) (3.6.1) Requirement already satisfied: jinja2 in /usr/local/lib/python3.12/dist-packages (from torch) (3.1.6) Requirement already satisfied: fsspec>=0.8.5 in /usr/local/lib/python3.12/dist-packages (from torch) (2025.3.0) Collecting nvidia-cuda-nvrtc-cu12==12.8.93 (from torch) Downloading nvidia_cuda_nvrtc_cu12-12.8.93-py3-none-manylinux2010_x86_64.manylinux_2_12_x86_64.whl.metadata (1.7 kB) Collecting nvidia-cuda-runtime-cu12==12.8.90 (from torch) Downloading nvidia_cuda_runtime_cu12-12.8.90-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl.metadata (1.7 kB) Collecting nvidia-cuda-cupti-cu12==12.8.90 (from torch) Downloading nvidia_cuda_cupti_cu12-12.8.90-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl.metadata (1.7 kB) Collecting nvidia-cudnn-cu12==9.10.2.21 (from torch) Downloading nvidia_cudnn_cu12-9.10.2.21-py3-none-manylinux_2_27_x86_64.whl.metadata (1.8 kB) Collecting nvidia-cublas-cu12==12.8.4.1 (from torch) Downloading nvidia_cublas_cu12-12.8.4.1-py3-none-manylinux_2_27_x86_64.whl.metadata (1.7 kB) Collecting nvidia-cufft-cu12==11.3.3.83 (from torch) Downloading nvidia_cufft_cu12-11.3.3.83-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl.metadata (1.7 kB) Collecting nvidia-curand-cu12==10.3.9.90 (from torch) Downloading nvidia_curand_cu12-10.3.9.90-py3-none-manylinux_2_27_x86_64.whl.metadata (1.7 kB) Collecting nvidia-cusolver-cu12==11.7.3.90 (from torch) Downloading nvidia_cusolver_cu12-11.7.3.90-py3-none-manylinux_2_27_x86_64.whl.metadata (1.8 kB) Collecting nvidia-cusparse-cu12==12.5.8.93 (from torch) Downloading nvidia_cusparse_cu12-12.5.8.93-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl.metadata (1.8 kB) Collecting nvidia-cusparselt-cu12==0.7.1 (from torch) Downloading nvidia_cusparselt_cu12-0.7.1-py3-none-manylinux2014_x86_64.whl.metadata (7.0 kB) Collecting nvidia-nccl-cu12==2.27.5 (from torch) Downloading nvidia_nccl_cu12-2.27.5-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl.metadata (2.0 kB) Collecting nvidia-nvshmem-cu12==3.3.20 (from torch) Downloading nvidia_nvshmem_cu12-3.3.20-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl.metadata (2.1 kB) Collecting nvidia-nvtx-cu12==12.8.90 (from torch) Downloading nvidia_nvtx_cu12-12.8.90-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl.metadata (1.8 kB) Collecting nvidia-nvjitlink-cu12==12.8.93 (from torch) Downloading nvidia_nvjitlink_cu12-12.8.93-py3-none-manylinux2010_x86_64.manylinux_2_12_x86_64.whl.metadata (1.7 kB) Collecting nvidia-cufile-cu12==1.13.1.3 (from torch) Downloading nvidia_cufile_cu12-1.13.1.3-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl.metadata (1.7 kB) Collecting triton==3.5.1 (from torch) Downloading triton-3.5.1-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl.metadata (1.7 kB) Requirement already satisfied: numpy>=1.17 in /usr/local/lib/python3.12/dist-packages (from datasets) (2.0.2) 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: requests>=2.32.2 in /usr/local/lib/python3.12/dist-packages (from datasets) (2.32.4) Requirement already satisfied: httpx<1.0.0 in /usr/local/lib/python3.12/dist-packages (from datasets) (0.28.1) Requirement already satisfied: tqdm>=4.66.3 in /usr/local/lib/python3.12/dist-packages (from datasets) (4.67.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: huggingface-hub<2.0,>=0.25.0 in /usr/local/lib/python3.12/dist-packages (from datasets) (0.36.0) Requirement already satisfied: packaging in /usr/local/lib/python3.12/dist-packages (from datasets) (25.0) Requirement already satisfied: pyyaml>=5.1 in /usr/local/lib/python3.12/dist-packages (from datasets) (6.0.3) 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: regex!=2019.12.17 in /usr/local/lib/python3.12/dist-packages (from transformers) (2025.11.3) 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: smart_open>=1.8.1 in /usr/local/lib/python3.12/dist-packages (from gensim) (7.5.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<2.0,>=0.25.0->datasets) (1.2.0) Requirement already satisfied: charset_normalizer<4,>=2 in /usr/local/lib/python3.12/dist-packages (from requests>=2.32.2->datasets) (3.4.4) Requirement already satisfied: urllib3<3,>=1.21.1 in /usr/local/lib/python3.12/dist-packages (from requests>=2.32.2->datasets) (2.5.0) Requirement already satisfied: wrapt in /usr/local/lib/python3.12/dist-packages (from smart_open>=1.8.1->gensim) (2.0.1) Requirement already satisfied: mpmath<1.4,>=1.1.0 in /usr/local/lib/python3.12/dist-packages (from sympy>=1.13.3->torch) (1.3.0) Requirement already satisfied: MarkupSafe>=2.0 in /usr/local/lib/python3.12/dist-packages (from jinja2->torch) (3.0.3) 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: six>=1.5 in /usr/local/lib/python3.12/dist-packages (from python-dateutil>=2.8.2->pandas->datasets) (1.17.0) Downloading torch-2.9.1-cp312-cp312-manylinux_2_28_x86_64.whl (899.7 MB) ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 899.7/899.7 MB 1.8 MB/s eta 0:00:00 Downloading nvidia_cublas_cu12-12.8.4.1-py3-none-manylinux_2_27_x86_64.whl (594.3 MB) ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 594.3/594.3 MB 2.1 MB/s eta 0:00:00 Downloading nvidia_cuda_cupti_cu12-12.8.90-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl (10.2 MB) ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 10.2/10.2 MB 31.6 MB/s eta 0:00:00 Downloading nvidia_cuda_nvrtc_cu12-12.8.93-py3-none-manylinux2010_x86_64.manylinux_2_12_x86_64.whl (88.0 MB) ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 88.0/88.0 MB 9.6 MB/s eta 0:00:00 Downloading nvidia_cuda_runtime_cu12-12.8.90-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl (954 kB) ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 954.8/954.8 kB 26.4 MB/s eta 0:00:00 Downloading nvidia_cudnn_cu12-9.10.2.21-py3-none-manylinux_2_27_x86_64.whl (706.8 MB) ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 706.8/706.8 MB 1.3 MB/s eta 0:00:00 Downloading nvidia_cufft_cu12-11.3.3.83-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl (193.1 MB) ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 193.1/193.1 MB 6.9 MB/s eta 0:00:00 Downloading nvidia_cufile_cu12-1.13.1.3-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl (1.2 MB) ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 1.2/1.2 MB 72.4 MB/s eta 0:00:00 Downloading nvidia_curand_cu12-10.3.9.90-py3-none-manylinux_2_27_x86_64.whl (63.6 MB) ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 63.6/63.6 MB 13.0 MB/s eta 0:00:00 Downloading nvidia_cusolver_cu12-11.7.3.90-py3-none-manylinux_2_27_x86_64.whl (267.5 MB) ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 267.5/267.5 MB 5.4 MB/s eta 0:00:00 Downloading nvidia_cusparse_cu12-12.5.8.93-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl (288.2 MB) ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 288.2/288.2 MB 6.0 MB/s eta 0:00:00 Downloading nvidia_cusparselt_cu12-0.7.1-py3-none-manylinux2014_x86_64.whl (287.2 MB) ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 287.2/287.2 MB 5.8 MB/s eta 0:00:00 Downloading nvidia_nccl_cu12-2.27.5-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl (322.3 MB) ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 322.3/322.3 MB 5.0 MB/s eta 0:00:00 Downloading nvidia_nvjitlink_cu12-12.8.93-py3-none-manylinux2010_x86_64.manylinux_2_12_x86_64.whl (39.3 MB) ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 39.3/39.3 MB 20.4 MB/s eta 0:00:00 Downloading nvidia_nvshmem_cu12-3.3.20-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl (124.7 MB) ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 124.7/124.7 MB 9.0 MB/s eta 0:00:00 Downloading nvidia_nvtx_cu12-12.8.90-py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl (89 kB) ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 90.0/90.0 kB 7.2 MB/s eta 0:00:00 Downloading triton-3.5.1-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl (170.5 MB) ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 170.5/170.5 MB 7.1 MB/s eta 0:00:00 Downloading datasets-4.4.1-py3-none-any.whl (511 kB) ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 511.6/511.6 kB 32.2 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 100.3 MB/s eta 0:00:00 Downloading gensim-4.4.0-cp312-cp312-manylinux_2_24_x86_64.manylinux_2_28_x86_64.whl (27.9 MB) ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 27.9/27.9 MB 69.7 MB/s eta 0:00:00 Downloading pyarrow-22.0.0-cp312-cp312-manylinux_2_28_x86_64.whl (47.7 MB) ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 47.7/47.7 MB 14.1 MB/s eta 0:00:00 Installing collected packages: nvidia-cusparselt-cu12, triton, pyarrow, nvidia-nvtx-cu12, nvidia-nvshmem-cu12, nvidia-nvjitlink-cu12, nvidia-nccl-cu12, nvidia-curand-cu12, nvidia-cufile-cu12, nvidia-cuda-runtime-cu12, nvidia-cuda-nvrtc-cu12, nvidia-cuda-cupti-cu12, nvidia-cublas-cu12, scikit-learn, nvidia-cusparse-cu12, nvidia-cufft-cu12, nvidia-cudnn-cu12, gensim, nvidia-cusolver-cu12, torch, datasets Attempting uninstall: pyarrow Found existing installation: pyarrow 18.1.0 Uninstalling pyarrow-18.1.0: Successfully uninstalled pyarrow-18.1.0 Attempting uninstall: nvidia-nccl-cu12 Found existing installation: nvidia-nccl-cu12 2.28.9 Uninstalling nvidia-nccl-cu12-2.28.9: Successfully uninstalled nvidia-nccl-cu12-2.28.9 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: torch Found existing installation: torch 2.9.0+cpu Uninstalling torch-2.9.0+cpu: Successfully uninstalled torch-2.9.0+cpu Attempting uninstall: datasets Found existing installation: datasets 4.0.0 Uninstalling datasets-4.0.0: Successfully uninstalled datasets-4.0.0 ERROR: pip's dependency resolver does not currently take into account all the packages that are installed. This behaviour is the source of the following dependency conflicts. torchvision 0.24.0+cpu requires torch==2.9.0, but you have torch 2.9.1 which is incompatible. torchaudio 2.9.0+cpu requires torch==2.9.0, but you have torch 2.9.1 which is incompatible. Successfully installed datasets-4.4.1 gensim-4.4.0 nvidia-cublas-cu12-12.8.4.1 nvidia-cuda-cupti-cu12-12.8.90 nvidia-cuda-nvrtc-cu12-12.8.93 nvidia-cuda-runtime-cu12-12.8.90 nvidia-cudnn-cu12-9.10.2.21 nvidia-cufft-cu12-11.3.3.83 nvidia-cufile-cu12-1.13.1.3 nvidia-curand-cu12-10.3.9.90 nvidia-cusolver-cu12-11.7.3.90 nvidia-cusparse-cu12-12.5.8.93 nvidia-cusparselt-cu12-0.7.1 nvidia-nccl-cu12-2.27.5 nvidia-nvjitlink-cu12-12.8.93 nvidia-nvshmem-cu12-3.3.20 nvidia-nvtx-cu12-12.8.90 pyarrow-22.0.0 scikit-learn-1.8.0 torch-2.9.1 triton-3.5.1
from datasets import load_dataset
# Cargamos el conjunto de datos. Se descargará y almacenará automáticamente en local.
# Este conjunto de datos contiene noticias de diferentes categorías. En este caso
# usaremos las categorías de mundo, deportes, negocios y ciencia ficción/tecnología.
# 'ag_news' a secas ya no resuelve: el Hub exige el formato namespace/name
dataset = load_dataset('fancyzhx/ag_news')
dataset
/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]
data/train-00000-of-00001.parquet: 0%| | 0.00/18.6M [00:00<?, ?B/s]
data/test-00000-of-00001.parquet: 0%| | 0.00/1.23M [00:00<?, ?B/s]
Generating train split: 0%| | 0/120000 [00:00<?, ? examples/s]
Generating test split: 0%| | 0/7600 [00:00<?, ? examples/s]
DatasetDict({
train: Dataset({
features: ['text', 'label'],
num_rows: 120000
})
test: Dataset({
features: ['text', 'label'],
num_rows: 7600
})
}) print(dataset['train'][0])
print(dataset['train'].features)
classes = dataset['train'].features["label"].names
classes
{'text': "Wall St. Bears Claw Back Into the Black (Reuters) Reuters - Short-sellers, Wall Street's dwindling\\band of ultra-cynics, are seeing green again.", 'label': 2}
{'text': Value('string'), 'label': ClassLabel(names=['World', 'Sports', 'Business', 'Sci/Tech'])}
['World', 'Sports', 'Business', 'Sci/Tech']
Automáticamente, la función load ha dividido el conjunto de datos en dos conjuntos: uno de train y uno para test. Para acceder a estos conjuntos, usaremos los atributos train i test del objeto dataset. Estos atributos son objetos data.Dataset que contiene los ejemplos y etiquetas del conjunto de capacitación y prueba. Para acceder a ejemplos y etiquetas, utilizaremos los atributos data y label del objeto data.Dataset.
# Separar el conjunto de datos en entrenamiento y test
ds_train = dataset['train']
ds_test = dataset['test']
# Veamos cuántos ejemplos hay en cada set
print('Número de ejemplos de train:', len(ds_train))
print('Número de ejemplos de test:', len(ds_test))
Número de ejemplos de train: 120000 Número de ejemplos de test: 7600
Imprimimos los primeros 5 ejemplos del conjunto de entrenamiento.Como podemos ver, cada ejemplo es una noticia y su etiqueta.
# Imprimimos los primeros 5 ejemplos del conjunto de entrenamiento
for w in ds_train.take(5):
print(f"{w['label']} ({classes[w['label']]}) -> {w['text']}")
2 (Business) -> Wall St. Bears Claw Back Into the Black (Reuters) Reuters - Short-sellers, Wall Street's dwindling\band of ultra-cynics, are seeing green again. 2 (Business) -> Carlyle Looks Toward Commercial Aerospace (Reuters) Reuters - Private investment firm Carlyle Group,\which has a reputation for making well-timed and occasionally\controversial plays in the defense industry, has quietly placed\its bets on another part of the market. 2 (Business) -> Oil and Economy Cloud Stocks' Outlook (Reuters) Reuters - Soaring crude prices plus worries\about the economy and the outlook for earnings are expected to\hang over the stock market next week during the depth of the\summer doldrums. 2 (Business) -> Iraq Halts Oil Exports from Main Southern Pipeline (Reuters) Reuters - Authorities have halted oil export\flows from the main pipeline in southern Iraq after\intelligence showed a rebel militia could strike\infrastructure, an oil official said on Saturday. 2 (Business) -> Oil prices soar to all-time record, posing new menace to US economy (AFP) AFP - Tearaway world oil prices, toppling records and straining wallets, present a new economic menace barely three months before the US presidential elections.
Tokenización¶
La representación del texto en un modelo de idioma requiere que el texto se convierta en números. Si queremos una representación de nivel de palabra, necesitamos hacer dos cosas:
- Utilizar un tokenizador para dividir el texto en tokens.
- Construir un vocabulario con estos tokens.
# Utilizamos el tokenizador de Bert (uno de los primeros modelos de lenguaje basados en transformación) para tokenizar las oraciones
from transformers import BertTokenizer
tokenizer = BertTokenizer.from_pretrained("google-bert/bert-base-uncased")
print(tokenizer.tokenize("-- Hello, how are you doing today?"))
# Podiamos ver el vocabulario de tokenización
vocab = tokenizer.get_vocab()
print(len(vocab))
tokenizer_config.json: 0%| | 0.00/48.0 [00:00<?, ?B/s]
vocab.txt: 0.00B [00:00, ?B/s]
tokenizer.json: 0.00B [00:00, ?B/s]
config.json: 0%| | 0.00/570 [00:00<?, ?B/s]
['-', '-', 'hello', ',', 'how', 'are', 'you', 'doing', 'today', '?'] 30522
Usando el tokenizer, también podemos convertir nuestra cadena tokenizada en un conjunto de números:
tokenitzada = tokenizer.tokenize("-- Hello, how are you doing today?")
def encode(text):
tk = tokenizer.tokenize(text)
return tokenizer.convert_tokens_to_ids(tk)
print(encode("-- Hello, how are you doing today?"))
[1011, 1011, 7592, 1010, 2129, 2024, 2017, 2725, 2651, 1029]
Representación del texto¶
Para entrenar un modelo de redes neuronales, necesitamos representar el texto como números. En esta práctica, usaremos la representación de la representación de la bolsa de palabras (BoW) que consiste en representar cada palabra como un número. Esta representación es muy simple y no tiene en cuenta el orden de las palabras o su semántica. Pero es una representación que funciona lo suficientemente bien en muchos casos.
Representación de la bolsa de las palabras¶
Aunque el significado de las palabras no es fácil de deducir sin poder acceder al contexto, en algunos casos, la representación de la bolsa de palabras puede ser útil.Por ejemplo, en el texto de una noticia, la palabra covid puede ser un buen indicador de que las noticias hablan sobre Covid-19 y la palabra snow puede ser un buen indicador de que las noticias hablan sobre el tiempo atmosférico.
De las técnicas clásicas de vectorización de texto, la más simple es la representación de la bolsa de las palabras (BoW). En esta representación, cada palabra se representa como un número. Para convertir un texto en una representación de BoW, primero creamos un vector con tantos ceros como las palabras están en el vocabulario. Luego, para cada palabra del texto, aumentamos el valor de la posición correspondiente al vector por 1. Por ejemplo, si el texto es this sentence is a test sentence, el vector resultante seria [1, 2, 1, 1, 0, 0, 0, 0, 0, 0, ...].
Si recordamos la representación one-hot, veremos que la representación de BoW es muy similar. La diferencia es que la representación one-hot será una serie de vectores con un solo 1 y el resto de los valores en 0. En cambio, la representación de BoW será un vector con tantas veces que aparezca cada palabra. Podemos considerar que la representación de BoW sería la suma de vectores únicos.
Por ejemplo, si el texto es this sentence is a test sentence, el vector one-hot de la primera palabra seria [1, 0, 0, 0, 0, 0, 0, 0, 0, 0, ...] y el vector one-hot de la segunda palabra seria [0, 1, 0, 0, 0, 0, 0, 0, 0, 0, ...]. La representación de BoW sería la suma de estos dos vectores: [1, 1, 0, 0, 0, 0, 0, 0, 0, 0, ...].
Para generar una representación de BoW, usaremos esta técnica para convertir cada palabra en un vector único y luego agregar todos los vectores. Para hacer esto, usaremos la función to_bow que crearemos a continuación. Esta función recibe un texto y devuelve un vector con la representación BoW del texto.
from sklearn.feature_extraction.text import CountVectorizer
vectorizer = CountVectorizer()
corpus = [
'I like hot dogs.',
'The dog ran fast.',
'Its hot outside.',
]
vectorizer.fit_transform(corpus)
vectorizer.transform(['My dog likes hot dogs on a hot day.']).toarray()
array([[1, 1, 0, 2, 0, 0, 0, 0, 0]])
Para calcular el vector BoW de una noticia de nuestro dataset AG_NEWS, podemos usar la siguiente función:
import torch
len_vocab = len(vocab)
def to_bow(text, tamany_vocabulari=len_vocab):
res = torch.zeros(tamany_vocabulari, dtype=torch.float32)
for i in encode(text):
if i<tamany_vocabulari:
res[i] += 1
return res
print(ds_train[0])
print(to_bow(ds_train[0]["text"]))
{'text': "Wall St. Bears Claw Back Into the Black (Reuters) Reuters - Short-sellers, Wall Street's dwindling\\band of ultra-cynics, are seeing green again.", 'label': 2}
tensor([0., 0., 0., ..., 0., 0., 0.])
Entrenamiento de modelos de clasificación BoW¶
Nuestro primer modelo será un clasificador de noticias utilizando la representación de BoW. Para hacer esto, crearemos un modelo de redes neuronales con una capa de entrada con tantas neuronas como las palabras están en nuestro vocabulario y una capa de salida con tantas neuronas como categorías que hay en nuestro conjunto de datos.
Representación BoW¶
Primero necesitamos convertir el texto en la representación de BoW utilizando la función to_bow que hemos creado antes. Esta función recibe un texto y devuelve un vector con la representación BoW del texto.
En pytorch se utilizan los DataLoaders, para cargar los datos en lotes y convertirlos en tensores de Pytorch. Aprovecharemos esta funcionalidad para convertir los datos de BoW en tensores de PyTorch, Usando el parámetro collate_fn del DataLoader y proporcionando una función que convierta los datos textuales en tensores de BoW.
from torch.utils.data import DataLoader
def bowify(batch):
'''
Esta característica recibe una lista de noticias y devuelve un tensor con las etiquetas
(vector de floats) y otro con las noticias codificadas como BoW (matriz de floats donde cada fila
es un vector de BoW).
'''
# Las etiquetas son 0, 1, 2 o 3.
# Usamos Longtensor porque son enteros.
etiquetes = torch.LongTensor([noticia["label"] for noticia in batch])
# La noticias son tensores de BoW
noticies = torch.stack([to_bow(noticia["text"]) for noticia in batch])
return (
etiquetes,
noticies
)
train_loader = DataLoader(ds_train, batch_size=16, collate_fn=bowify)
test_loader = DataLoader(ds_test, batch_size=16, collate_fn=bowify)
Modelo de classificación¶
Ahora definamos una red neuronal clasificadora simple que contiene una capa lineal. El tamaño del vector de entrada es igual a vocab_size, y el tamaño de salida corresponde al número de clases (4). Debido a que estamos resolviendo la tarea de clasificación, la función de activación final es LogSoftmax().
net = torch.nn.Sequential(
torch.nn.Linear(len(vocab), 4),
torch.nn.LogSoftmax(dim=1)
)
Entrenamiento del modelo¶
Ahora definiremos el bucle de entrenamiento estándar de Pytorch. Debido a que nuestro conjunto de datos es bastante grande, con nuestro propósito de enseñanza, entrenaremos solo para una epoca, y a veces incluso por menos de una epoca (especificando el parámetro epoch_size nos permite limitar el entrenamiento). También informaremos la precisión del entrenamiento acumulado durante el entrenamiento; La frecuencia de notificación se especifica utilizando el parámetro report_freq.
Para entrenar el modelo, usaremos el optimizador Adam(ya que es uno de los optimizadores más utilizados) y la función de costo CrossEntropyLoss (ya que tenemos un problema de calificación con más de dos clases).
def train_epoch(
net,
dataloader,
lr=0.01,
optimizer=None,
loss_fn=torch.nn.NLLLoss(),
epoch_size=None,
report_freq=200,
):
# Si no se especifica un optimizador, usamos Adam
optimizer = optimizer or torch.optim.Adam(net.parameters(), lr=lr)
# Ponemos la red en modo de entrenamiento.Esto activa el comportamiento de las capas de DropOut, por ejemplo.
net.train()
# Inicializar las variables que nos servirán para calcular la precisión
total_loss, acc, count, i = 0, 0, 0, 0
# Iteremamos sobre el dataloader
for labels, features in dataloader:
# Ponemos los gradientes a cero
optimizer.zero_grad()
# calculamos la salida de la red
out = net(features)
# Calculamos la pérdida. Esta función ya se aplica a Softmax a la salida.
loss = loss_fn(out, labels) # cross_entropy(out,labels)
# Propagamos la pérdida de regreso. Esto hará que se calculen los gradientes .
loss.backward()
# Actualizamos los pesos de la red. Esto toma un paso de optimización.
optimizer.step()
# Actualizamos variables para calcular la precisión.
total_loss += loss
# Calculamos la precisión. Para hacer esto, debemos convertir la salida de red en etiquetas.
# La clase con la mayor probabilidad es la que predecimos como etiqueta.
_, predicted = torch.max(out, 1)
acc += (predicted == labels).sum()
# Actualizamos el contador de muestras
count += len(labels)
# Mostramos la precisión cada report_freq muestras
i += 1
if i % report_freq == 0:
print(f"{count}: acc={acc.item()/count}")
# Si se especifica epoch_size y ya hemos procesado este número de muestras, dejamos el bucle.
if epoch_size and count > epoch_size:
break
return total_loss.item() / count, acc.item() / count
train_epoch(net, train_loader, epoch_size=15000)
3200: acc=0.7403125 6400: acc=0.80359375 9600: acc=0.8258333333333333 12800: acc=0.841171875
(0.03097911379230556, 0.84761460554371)
El modelo ha logrado una precisión cercano a 0.85 en el conjunto de entrenamiento; Un número suficientemente aceptable considerando que hemos simplificado el problema para reducir el tiempo de ejecución del tutorial. En un caso real, usaríamos todas las noticias del conjunto de entrenamiento y el modelo sería más preciso.
Representación de Word2Vec¶
La representación de Word2Vec es una representación ampliamente utilizada en el procesamiento del lenguaje natural. Esta representación tiene en cuenta el contexto de las palabras y permite operaciones con las palabras. Por ejemplo, si restamos la palabra vector king y sumamos el vector de la palabra woman, obtendremos un vector que será muy similar al vector de la plabra queen.
Para generar representación de Word2Vec, usaremos la librería gensim. Esta librería contiene muchos modelos de representación de palabras. En este caso usaremos el modelo word2vec-google-news-300 que contiene la representación de Word2Vec de 3 millones de palabras y frases.
La primera vez que esta celda se está ejecutando, la función
loadDescargará el modelo de 1.5GB. Esto puede tomar unos minutos. Esta función devuelve un objetoKeyedVectorsque contiene la representación Word2Vec.
import gensim.downloader as api
w2v = api.load('word2vec-google-news-300')
[==================================================] 100.0% 1662.8/1662.8MB downloaded
Ahora podemos acceder a la representación de Word2Vec de cada palabra. Por ejemplo, para acceder a la representación de la palabra king, usaremos la función get_vector del objeto KeyedVectors.
w2v.get_vector('king')
array([ 1.25976562e-01, 2.97851562e-02, 8.60595703e-03, 1.39648438e-01,
-2.56347656e-02, -3.61328125e-02, 1.11816406e-01, -1.98242188e-01,
5.12695312e-02, 3.63281250e-01, -2.42187500e-01, -3.02734375e-01,
-1.77734375e-01, -2.49023438e-02, -1.67968750e-01, -1.69921875e-01,
3.46679688e-02, 5.21850586e-03, 4.63867188e-02, 1.28906250e-01,
1.36718750e-01, 1.12792969e-01, 5.95703125e-02, 1.36718750e-01,
1.01074219e-01, -1.76757812e-01, -2.51953125e-01, 5.98144531e-02,
3.41796875e-01, -3.11279297e-02, 1.04492188e-01, 6.17675781e-02,
1.24511719e-01, 4.00390625e-01, -3.22265625e-01, 8.39843750e-02,
3.90625000e-02, 5.85937500e-03, 7.03125000e-02, 1.72851562e-01,
1.38671875e-01, -2.31445312e-01, 2.83203125e-01, 1.42578125e-01,
3.41796875e-01, -2.39257812e-02, -1.09863281e-01, 3.32031250e-02,
-5.46875000e-02, 1.53198242e-02, -1.62109375e-01, 1.58203125e-01,
-2.59765625e-01, 2.01416016e-02, -1.63085938e-01, 1.35803223e-03,
-1.44531250e-01, -5.68847656e-02, 4.29687500e-02, -2.46582031e-02,
1.85546875e-01, 4.47265625e-01, 9.58251953e-03, 1.31835938e-01,
9.86328125e-02, -1.85546875e-01, -1.00097656e-01, -1.33789062e-01,
-1.25000000e-01, 2.83203125e-01, 1.23046875e-01, 5.32226562e-02,
-1.77734375e-01, 8.59375000e-02, -2.18505859e-02, 2.05078125e-02,
-1.39648438e-01, 2.51464844e-02, 1.38671875e-01, -1.05468750e-01,
1.38671875e-01, 8.88671875e-02, -7.51953125e-02, -2.13623047e-02,
1.72851562e-01, 4.63867188e-02, -2.65625000e-01, 8.91113281e-03,
1.49414062e-01, 3.78417969e-02, 2.38281250e-01, -1.24511719e-01,
-2.17773438e-01, -1.81640625e-01, 2.97851562e-02, 5.71289062e-02,
-2.89306641e-02, 1.24511719e-02, 9.66796875e-02, -2.31445312e-01,
5.81054688e-02, 6.68945312e-02, 7.08007812e-02, -3.08593750e-01,
-2.14843750e-01, 1.45507812e-01, -4.27734375e-01, -9.39941406e-03,
1.54296875e-01, -7.66601562e-02, 2.89062500e-01, 2.77343750e-01,
-4.86373901e-04, -1.36718750e-01, 3.24218750e-01, -2.46093750e-01,
-3.03649902e-03, -2.11914062e-01, 1.25000000e-01, 2.69531250e-01,
2.04101562e-01, 8.25195312e-02, -2.01171875e-01, -1.60156250e-01,
-3.78417969e-02, -1.20117188e-01, 1.15234375e-01, -4.10156250e-02,
-3.95507812e-02, -8.98437500e-02, 6.34765625e-03, 2.03125000e-01,
1.86523438e-01, 2.73437500e-01, 6.29882812e-02, 1.41601562e-01,
-9.81445312e-02, 1.38671875e-01, 1.82617188e-01, 1.73828125e-01,
1.73828125e-01, -2.37304688e-01, 1.78710938e-01, 6.34765625e-02,
2.36328125e-01, -2.08984375e-01, 8.74023438e-02, -1.66015625e-01,
-7.91015625e-02, 2.43164062e-01, -8.88671875e-02, 1.26953125e-01,
-2.16796875e-01, -1.73828125e-01, -3.59375000e-01, -8.25195312e-02,
-6.49414062e-02, 5.07812500e-02, 1.35742188e-01, -7.47070312e-02,
-1.64062500e-01, 1.15356445e-02, 4.45312500e-01, -2.15820312e-01,
-1.11328125e-01, -1.92382812e-01, 1.70898438e-01, -1.25000000e-01,
2.65502930e-03, 1.92382812e-01, -1.74804688e-01, 1.39648438e-01,
2.92968750e-01, 1.13281250e-01, 5.95703125e-02, -6.39648438e-02,
9.96093750e-02, -2.72216797e-02, 1.96533203e-02, 4.27246094e-02,
-2.46093750e-01, 6.39648438e-02, -2.25585938e-01, -1.68945312e-01,
2.89916992e-03, 8.20312500e-02, 3.41796875e-01, 4.32128906e-02,
1.32812500e-01, 1.42578125e-01, 7.61718750e-02, 5.98144531e-02,
-1.19140625e-01, 2.74658203e-03, -6.29882812e-02, -2.72216797e-02,
-4.82177734e-03, -8.20312500e-02, -2.49023438e-02, -4.00390625e-01,
-1.06933594e-01, 4.24804688e-02, 7.76367188e-02, -1.16699219e-01,
7.37304688e-02, -9.22851562e-02, 1.07910156e-01, 1.58203125e-01,
4.24804688e-02, 1.26953125e-01, 3.61328125e-02, 2.67578125e-01,
-1.01074219e-01, -3.02734375e-01, -5.76171875e-02, 5.05371094e-02,
5.26428223e-04, -2.07031250e-01, -1.38671875e-01, -8.97216797e-03,
-2.78320312e-02, -1.41601562e-01, 2.07031250e-01, -1.58203125e-01,
1.27929688e-01, 1.49414062e-01, -2.24609375e-02, -8.44726562e-02,
1.22558594e-01, 2.15820312e-01, -2.13867188e-01, -3.12500000e-01,
-3.73046875e-01, 4.08935547e-03, 1.07421875e-01, 1.06933594e-01,
7.32421875e-02, 8.97216797e-03, -3.88183594e-02, -1.29882812e-01,
1.49414062e-01, -2.14843750e-01, -1.83868408e-03, 9.91210938e-02,
1.57226562e-01, -1.14257812e-01, -2.05078125e-01, 9.91210938e-02,
3.69140625e-01, -1.97265625e-01, 3.54003906e-02, 1.09375000e-01,
1.31835938e-01, 1.66992188e-01, 2.35351562e-01, 1.04980469e-01,
-4.96093750e-01, -1.64062500e-01, -1.56250000e-01, -5.22460938e-02,
1.03027344e-01, 2.43164062e-01, -1.88476562e-01, 5.07812500e-02,
-9.37500000e-02, -6.68945312e-02, 2.27050781e-02, 7.61718750e-02,
2.89062500e-01, 3.10546875e-01, -5.37109375e-02, 2.28515625e-01,
2.51464844e-02, 6.78710938e-02, -1.21093750e-01, -2.15820312e-01,
-2.73437500e-01, -3.07617188e-02, -3.37890625e-01, 1.53320312e-01,
2.33398438e-01, -2.08007812e-01, 3.73046875e-01, 8.20312500e-02,
2.51953125e-01, -7.61718750e-02, -4.66308594e-02, -2.23388672e-02,
2.99072266e-02, -5.93261719e-02, -4.66918945e-03, -2.44140625e-01,
-2.09960938e-01, -2.87109375e-01, -4.54101562e-02, -1.77734375e-01,
-2.79296875e-01, -8.59375000e-02, 9.13085938e-02, 2.51953125e-01],
dtype=float32) También podemos acceder a las palabras más similares a una palabra. Por ejemplo, para acceder a las palabras más similares a la palabra king, usaremos la función most_similar del objeto KeyedVectors.
for w, p in w2v.most_similar('king'):
print(f"{w} -> {p}")
kings -> 0.7138045430183411 queen -> 0.6510956883430481 monarch -> 0.6413194537162781 crown_prince -> 0.6204220056533813 prince -> 0.6159993410110474 sultan -> 0.5864824056625366 ruler -> 0.5797567367553711 princes -> 0.5646552443504333 Prince_Paras -> 0.5432944297790527 throne -> 0.5422105193138123
Lo más interesante de la representación de Word2Vec es que los vectores tienen una estructura matemática que nos permite realizar operaciones con las palabras. Por ejemplo, si restamos el vector de la palabra man al vector de la palabra king y sumamos el vector de la palabra woman, obtendremos un vector que será muy similar al vector de la palabra queen.
$$ KING - MAN + WOMAN = QUEEN $$
Para hacer esta operación, usaremos la función most_similar del objeto KeyedVectors y pasaremos los vectores de las palabras king, woman y man. Esta característica devolverá una lista con las palabras más similares al vector resultante.Como podemos ver, la palabra más similar es queen.
w2v.most_similar(positive=['king', 'woman'], negative=['man'])[0]
('queen', 0.7118193507194519) Clasificador de Word2Vec¶
Ahora crearemos un clasificador de noticias utilizando la representación Word2Vec. Primero tendremos que obtener la representación de cada palabra para convertir el texto en vectores. Luego agregaremos todos los vectores para obtener un vector para cada noticia. Este vector será la representación de las noticias.
Para convertir un texto en un vector, usaremos la función to_w2v que crearemos a continuación. Esta función recibe un texto y devuelve un vector con la representación Word2Vec del texto.
def to_w2v(text):
res = torch.zeros(300, dtype=torch.float32)
for word in text:
if word in w2v:
res += torch.tensor(w2v.get_vector(word))
return res
print(to_w2v(ds_train[0]["text"]))
tensor([-17.0809, 11.0404, -0.9337, 12.4042, -6.2286, 3.0224, -10.0442,
-8.5156, -5.9407, 1.1501, -3.8471, -8.0006, -18.2444, 4.3982,
-14.2061, 11.0110, 11.2352, 14.8521, -2.5686, 2.8961, -22.3914,
-3.2182, 9.7872, 0.3238, -8.6214, 4.2367, -21.9348, 5.7704,
-0.6942, -1.7075, -2.4800, 2.1805, -7.0602, -12.3824, -11.6949,
8.2563, -18.9995, 11.3932, -7.3198, 7.3370, -6.1129, -3.6244,
5.8519, 8.3060, 3.9137, -1.8091, -3.2730, -15.8203, -9.6418,
8.9092, -16.8270, 24.5614, -2.5387, 21.7112, 6.0571, 14.3324,
-17.4978, -12.2693, 1.1129, -15.9192, -12.1886, -9.5650, -19.0873,
-7.7948, -4.9111, -18.4653, -10.2332, 11.3437, -6.0452, 5.4705,
3.7500, -9.5068, 4.4747, -0.2912, -3.9221, 0.3543, 13.0927,
2.3088, 3.5300, -11.2126, -14.8031, -2.9008, -3.4219, -0.3365,
13.8353, 7.0914, -5.2219, 22.0132, 4.2657, 5.8488, -0.5776,
-1.5022, -5.0004, -13.3813, 4.8757, 10.3992, -9.8992, 10.6411,
25.6584, -3.4937, -6.5989, -1.0960, -6.7775, 0.1842, -6.0798,
13.2260, -6.2332, 4.9711, 0.8566, 4.1002, -11.5986, -16.8590,
-6.4362, -3.6979, 4.9203, 15.2933, 7.6364, -5.8566, -1.6903,
-2.3312, 12.3486, 7.5709, -0.6597, 2.7831, 12.6196, -15.9392,
-9.4420, -1.7229, 7.7839, 10.5602, -5.9280, -2.6489, -6.4361,
-3.8383, -16.0124, 8.0287, -3.4375, 2.8186, 22.9197, 13.0072,
20.2472, -6.0054, 4.0575, -4.7046, -6.8406, -12.1006, -4.0645,
18.0959, -4.8794, 1.5283, 8.0677, -28.4229, -6.3982, -4.6095,
-8.2329, -7.8615, 8.5930, 14.3553, -2.6136, -0.7672, 0.5814,
6.9687, 0.7667, 0.1969, -1.1499, -4.6281, 15.1071, 5.1883,
-10.1836, 6.9755, -9.1791, -8.8102, -6.4574, -8.9768, -1.3514,
20.5255, 22.8086, -22.4160, -2.1751, -5.2745, 0.4971, 1.9747,
5.5252, -4.8856, -2.1867, 5.9344, -5.4659, 5.1147, -3.3837,
5.4895, 12.1746, -4.1896, -27.1298, -4.4509, 10.7126, 4.8896,
-6.0110, -0.7719, 7.8879, -10.4668, -9.2913, 2.1059, -19.3102,
-10.8101, 5.4989, -7.4446, -4.7968, 9.8521, -3.9826, 14.8542,
16.3674, 7.4929, -11.3996, 1.8357, -8.1945, 6.2330, 15.2261,
-3.4122, -16.1802, -2.0000, -12.0552, 11.2962, 5.6537, -0.8348,
-0.8463, -6.4080, 5.8111, 2.4668, 1.0925, -14.5064, 1.1021,
-4.3229, -8.5156, 1.3596, 0.2417, 1.4028, 7.4663, 8.9206,
7.3249, 8.0591, 7.7924, 6.9987, 27.3159, -4.7353, 0.7053,
6.7754, -12.8845, 13.8699, 5.8623, -6.8129, 5.8627, 2.7595,
6.3065, 9.9255, -2.8854, -10.1693, -7.0736, 7.9216, 2.5093,
-9.5866, 7.4031, -0.9011, 9.9832, -2.0049, -6.3317, 0.4062,
0.0936, 1.1288, -2.5539, -2.7307, -5.4014, -2.8721, 1.3374,
0.2924, 3.5125, -8.9189, -15.5585, -13.4326, -9.4054, 3.7766,
-12.4168, 12.1424, -1.4987, 0.1738, -0.9734, 6.0570, -0.8122,
-3.2520, -5.7413, -4.4579, 4.8879, -0.8176, -9.8232, 8.6069,
-3.3508, -15.2089, -10.0510, -1.9859, -10.7878, 18.1031])
Como lo hicimos con la representación de BoW, usaremos el DataLoaders de PyTorch para convertir los datos en vectores Word2Vec en tensores Pytorch. Aprovecharemos el parámetro collate_fn del DataLoader para proporcionar una función que convierta los datos textuales en tensores Word2Vec.
def w2vify(batch):
etiquetes = torch.LongTensor([noticia["label"] for noticia in batch])
noticies = torch.stack([to_w2v(tokenizer.tokenize(noticia["text"])) for noticia in batch])
return etiquetes, noticies
train_loader = DataLoader(ds_train, batch_size=16, collate_fn=w2vify)
test_loader = DataLoader(ds_test, batch_size=16, collate_fn=w2vify)
Modelo de clasificación¶
Ahora crearemos el modelo de clasificación usando Pytorch. Definiremos un modelo simple con una capa lineal. El tamaño del vector de entrada será 300 (el tamaño de la representación Word2Vec) y el tamaño de la salida será el número de clases (4). Como estamos resolviendo una tarea de clasificación, la función de activación final será LogSoftmax().
net = torch.nn.Sequential(
torch.nn.Linear(300, 4),
torch.nn.LogSoftmax(dim=1)
)
Finalmente, entrenamos al modelo utilizando el mismo procedimiento que hemos realizado con la representación de BoW.
train_epoch(net, train_loader, epoch_size=15000)
3200: acc=0.73375 6400: acc=0.76828125 9600: acc=0.7792708333333334 12800: acc=0.792578125
(0.08423855258966051, 0.7985740938166311)
El resultado no es muy bueno. Esto se debe a que el modelo Word2Vec que utilizamos no tiene las palabras que aparecen en el conjunto de datos. Por ejemplo, si buscamos la palabra covid, veremos que no aparece en el modelo.
Para resolver este problema, tendremos que usar un modelo Word2Vec entrenado con las palabras del conjunto de datos. Pero esto es muy lento y no lo haremos en este tutorial.