{ "cells": [ { "cell_type": "markdown", "metadata": { "id": "iKENeWEEwKIb" }, "source": [ "\n", "Zero-shot learning en un problema de clasificación\n", "==================================================" ] }, { "cell_type": "markdown", "source": [ "Los grandes modelos de lenguaje exhiben grandes habilidades en zero-shot learning. Sin embargo, los resultados dependen mucho de la capacidad del modelo, y de la ténica que utilicemos para resolver el problema.\n", "\n", "En este ejemplo, utizaremos un modelo de lenguaje para resolver el problema de clasificación de tweets sin entrenar ningún modelo (zero-shot)." ], "metadata": { "id": "MJSIpyXadLcR" } }, { "cell_type": "markdown", "metadata": { "id": "LV5DTexIybKw" }, "source": [ "Introducción\n", "------------" ] }, { "cell_type": "markdown", "metadata": { "id": "sT_t9OxYwKIc" }, "source": [ "Los grandes modelos de lenguaje son capaces de resolver problemas de clasificación al utilizar determinadas estructuras del idioma." ] }, { "cell_type": "markdown", "metadata": { "id": "Dcyc_TQ6dis7" }, "source": [ "### Para ejecutar este notebook" ] }, { "cell_type": "markdown", "metadata": { "id": "YWGfRkiHybK1" }, "source": [ "Para ejecutar este notebook, instale las siguientes librerias:" ] }, { "cell_type": "code", "execution_count": null, "metadata": { "id": "YXE4ZBrlybK1", "outputId": "90a8b895-c869-4de5-f9e8-c46a0c2d53aa", "colab": { "base_uri": "https://localhost:8080/" } }, "outputs": [ { "output_type": "stream", "name": "stdout", "text": [ "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m7.2/7.2 MB\u001b[0m \u001b[31m30.2 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m268.8/268.8 kB\u001b[0m \u001b[31m24.0 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m1.3/1.3 MB\u001b[0m \u001b[31m46.9 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m45.9/45.9 kB\u001b[0m \u001b[31m6.2 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m486.2/486.2 kB\u001b[0m \u001b[31m24.6 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m7.8/7.8 MB\u001b[0m \u001b[31m72.2 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m1.3/1.3 MB\u001b[0m \u001b[31m41.5 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m244.2/244.2 kB\u001b[0m \u001b[31m24.1 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m86.0/86.0 kB\u001b[0m \u001b[31m8.7 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", "\u001b[?25h Preparing metadata (setup.py) ... \u001b[?25l\u001b[?25hdone\n", "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m81.4/81.4 kB\u001b[0m \u001b[31m9.5 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m110.5/110.5 kB\u001b[0m \u001b[31m16.1 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m212.5/212.5 kB\u001b[0m \u001b[31m22.5 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", "\u001b[2K \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m134.3/134.3 kB\u001b[0m \u001b[31m17.5 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n", "\u001b[?25h Building wheel for sentence-transformers (setup.py) ... \u001b[?25l\u001b[?25hdone\n" ] } ], "source": [ "!wget https://raw.githubusercontent.com/santiagxf/M72109/master/NLP/Datasets/mascorpus/tweets_marketing.csv \\\n", " --quiet --no-clobber --directory-prefix ./Datasets/mascorpus/\n", "\n", "!wget https://raw.githubusercontent.com/santiagxf/M72109/master/docs/nlp/neural/zero_shot_classification.txt \\\n", " --quiet --no-clobber\n", "\n", "!pip install -r zero_shot_classification.txt --quiet" ] }, { "cell_type": "code", "execution_count": null, "metadata": { "id": "Ntcs1AlpfckX" }, "outputs": [], "source": [ "import warnings\n", "warnings.filterwarnings('ignore')" ] }, { "cell_type": "markdown", "metadata": { "id": "_gBXNzwYwKIu" }, "source": [ "Cargamos el set de datos" ] }, { "cell_type": "code", "execution_count": null, "metadata": { "id": "I8vqJD9JwKIv" }, "outputs": [], "source": [ "import pandas as pd\n", "\n", "tweets = pd.read_csv('Datasets/mascorpus/tweets_marketing.csv')" ] }, { "cell_type": "code", "execution_count": null, "metadata": { "id": "q-9nkTp9wKIy" }, "outputs": [], "source": [ "from sklearn.model_selection import train_test_split\n", "\n", "X_train, X_test, y_train, y_test = train_test_split(tweets['TEXTO'], tweets['SECTOR'],\n", " test_size=0.33,\n", " stratify=tweets['SECTOR'])" ] }, { "cell_type": "markdown", "metadata": { "id": "ExV73PU4Ak1I" }, "source": [ "### Verificando el hardware disponible" ] }, { "cell_type": "code", "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/" }, "gather": { "logged": 1604083146665 }, "id": "ULqO3R6_Am2I", "outputId": "465079ce-98b1-4ac0-fed7-131554cb7819" }, "outputs": [ { "output_type": "stream", "name": "stdout", "text": [ "Este notebook se está ejecutando en cuda\n" ] } ], "source": [ "import torch\n", "device = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')\n", "\n", "print(\"Este notebook se está ejecutando en\", device)" ] }, { "cell_type": "markdown", "metadata": { "id": "L1u3qFaSwKI9" }, "source": [ "## Creando un modelo de clasificación utilizando zero-shot learning" ] }, { "cell_type": "markdown", "metadata": { "id": "YgmeXBEMwKI9" }, "source": [ "Trataremos de resolver entonces el mismo problema de clasificación con el que veniamos trabajando: clasificar los tweets dependiendo del sector al que pertenecen. Recordemos que tenemos 7 categorias distintas:" ] }, { "cell_type": "code", "execution_count": null, "metadata": { "colab": { "base_uri": "https://localhost:8080/" }, "id": "PdSw_zU_wKI-", "outputId": "e6efc384-a685-4455-c860-982ac993e0ae" }, "outputs": [ { "output_type": "execute_result", "data": { "text/plain": [ "['RETAIL',\n", " 'TELCO',\n", " 'ALIMENTACION',\n", " 'AUTOMOCION',\n", " 'BANCA',\n", " 'BEBIDAS',\n", " 'DEPORTES']" ] }, "metadata": {}, "execution_count": 5 } ], "source": [ "labels = tweets['SECTOR'].unique().tolist()\n", "labels" ] }, { "cell_type": "markdown", "source": [ "El modelo base que utizaremos es BART el cual es multi-lenguaje y puede manejar texto en multiples idiomas:" ], "metadata": { "id": "bpRYEuaSeomx" } }, { "cell_type": "code", "source": [ "model_name = \"facebook/bart-large-mnli\"" ], "metadata": { "id": "vIA59JYVh_nZ" }, "execution_count": null, "outputs": [] }, { "cell_type": "markdown", "source": [ "En la libraría `transformers`, podemos utilizar un pipeline del tipo `zero-shot-classification`, el cual nos permite resolver una tarea de clasificación al modelarla como un problema de *text completion*.\n", "\n", "Este pipeline construye tantos *prompts* como diferentes clases querramos predecir. Luego, aplica una plantilla que combina el texto de entrada y la etiqueta y busca qué etiqueta genera un texto al cual el modelo de lenguaje le asigna la mayor probabilidad:" ], "metadata": { "id": "Dhy1znAkeuOj" } }, { "cell_type": "code", "source": [ "from transformers import pipeline\n", "\n", "classifier = pipeline(task=\"zero-shot-classification\", model=model_name, device=0)" ], "metadata": { "id": "BidKbvennpdB", "outputId": "e929f033-1ac4-466e-b16c-7ceb690c6a7d", "colab": { "base_uri": "https://localhost:8080/", "height": 209, "referenced_widgets": [ "46ae5516e1c844d6a7e84e201458fc1a", "69ec97064b05411d99da76093a2d2e24", "8602f96b7a1c4a8e94cf6bdf733fa840", "43f977a99dbc47559094f31233a28fe1", "0f7164809a7642ed9613dd1212714074", "ae5bcdd63ac64d19bcc7920466928f99", "915bad7cddab47c2815f3e9bc9a3c482", "a9cbdac574eb4d47ab413e7cf7378bc3", "c3116c6aab8e45fb9aaf214d975de642", "134a3a574acb43e6a6a8dba1d9d12ee2", "049d342f75544860ad64cb4747e90d25", "2aa83447b4854280a206057bfc3ab952", "60ce20857fb8494c98511af8d3b40c96", "16c5ca95718a415789152e0b750874e9", "50195876a4724801a7a46249b43a820e", "1dced5f52c774e68b27feb2ad05b500a", "2ba567eff9b245469346b8438ef9fd3d", "4e06b55c9a5242cd8b74e420dbe5615d", "a7795c3229ea40e88b7574d347758df8", "72889350b3a448bb95732c2fb0071d88", "b7d11f741c3f43b3b302ca0aa20a9deb", "f5e02e5f62de47e286f6818ad4626d02", "d92c8257caa54904b7f1b8be7cd9a815", "d5b625d4175f42cfb1a49f3ca2aaeb30", "e67608d66b3248c9b4124b2b50b6a2bc", "1c12ab1d723549038ce89d4985aee0d6", "5cd5ba5276074e80b31a4c02b416d8a4", "7dfe1be2f4f346929fc725c3dbde0478", "c5d4f084502d46bf8cb689aaab85a8c0", "80a282ebd15a447a8f98b9fb85dfdec8", "98c16a98e3f94a80a0725846c0ac4b9f", "0f8666c179f34e5db54df2e9cef90a04", "bfde162f9c044682adccc28a2b6af6e0", "d5771c74aea34dbdb8538e710960999c", "fe257d3d0878497d89f4a061802df1cc", "dfaed72bfc3a4dc28be86ad3838f4d36", "380ab6a8b04f4ad0b63d72a8c4449053", "bd5450e64cf749cbbcc0b30305dbc7af", "5d3b6c283a9d4b05a45111fa5ede882f", "6d245e8e616d44d7b8932a33adf0b0a9", "fc12746c22db41d3b572747d0dcaa0cc", "cf8e30c860564f029f6411e3fb788f76", "57066fea8bf740d79f0a9a3b67b79a80", "3b27aeb40dfa4a8d99fde3892b7761a4", "67fa874ba6d04143babb2272a2423d3b", "b2f261863fd1443cb19b3389eaf7ff58", "a0e9f7858ea046c1a4783d3369b26ad4", "8e67f989f32c4cfd82abde0501523227", "ff3cf42d8d4a43b28dec9713dbe8d2ca", "66d1293b2b4d4b4084113a96afa41e36", "412128c15f1a41b086b5a0c91ab84a2f", "0b6968bfdb4e4a7e9fbd7e04aef48044", "4c6ec7310a2843c6a322280f12489da7", "f9bf329b22fd4486ae4d36b7d4160618", "30f02583473340fab0169230539d69e8", "17571b3ed02d41e1a0521529f41ca512", "65fcf03aa2c442f19f0a395f6b655b15", "c8fd46728a4c44c0b2b04f3a5ac6286a", "731a8f4a66174da7b728f1c9f517a92f", "821f4ed97f4b4ded8aedd31679434666", "99ac5ffb6d1d480e998755604fea6b8a", "38df2483e7cc46e9b2673204d781718d", "2d9ee02295bf4cdd9c523e1ce6f656ed", "50ea79d630e6452ba2ae3ca1c5c9c77b", "72cdd941e6e24e37bc8930ef2f5fd723", "0dda713aa70040f688695d22eb47cbb3" ] } }, "execution_count": null, "outputs": [ { "output_type": "display_data", "data": { "text/plain": [ "Downloading (…)lve/main/config.json: 0%| | 0.00/1.15k [00:00