{ "cells": [ { "cell_type": "markdown", "metadata": {}, "source": [ "# テキスト分類タスク\n", "\n", "前述の通り、今回は**AG_NEWS**データセットを使用したシンプルなテキスト分類タスクに焦点を当てます。このタスクでは、ニュースの見出しを「世界」「スポーツ」「ビジネス」「科学/技術」の4つのカテゴリのいずれかに分類します。\n", "\n", "## データセット\n", "\n", "このデータセットは[`torchtext`](https://github.com/pytorch/text)モジュールに組み込まれているため、簡単にアクセスできます。\n" ] }, { "cell_type": "code", "execution_count": 1, "metadata": {}, "outputs": [], "source": [ "import torch\n", "import torchtext\n", "import os\n", "import collections\n", "os.makedirs('./data',exist_ok=True)\n", "train_dataset, test_dataset = torchtext.datasets.AG_NEWS(root='./data')\n", "classes = ['World', 'Sports', 'Business', 'Sci/Tech']" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "ここで、`train_dataset` と `test_dataset` は、それぞれクラス番号(ラベル)とテキストのペアを返すコレクションを含んでいます。例えば:\n" ] }, { "cell_type": "code", "execution_count": 2, "metadata": {}, "outputs": [ { "data": { "text/plain": [ "(3,\n", " \"Wall St. Bears Claw Back Into the Black (Reuters) Reuters - Short-sellers, Wall Street's dwindling\\\\band of ultra-cynics, are seeing green again.\")" ] }, "execution_count": 2, "metadata": {}, "output_type": "execute_result" } ], "source": [ "list(train_dataset)[0]" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "では、データセットから最初の10件の新しい見出しを出力しましょう。\n" ] }, { "cell_type": "code", "execution_count": 5, "metadata": {}, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "**Sci/Tech** -> Wall St. Bears Claw Back Into the Black (Reuters) Reuters - Short-sellers, Wall Street's dwindling\\band of ultra-cynics, are seeing green again.\n", "**Sci/Tech** -> 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.\n", "**Sci/Tech** -> 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.\n", "**Sci/Tech** -> 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.\n", "**Sci/Tech** -> 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.\n" ] } ], "source": [ "for i,x in zip(range(5),train_dataset):\n", " print(f\"**{classes[x[0]]}** -> {x[1]}\")\n" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "データセットはイテレーターであるため、データを複数回使用したい場合はリストに変換する必要があります。\n" ] }, { "cell_type": "code", "execution_count": 3, "metadata": {}, "outputs": [], "source": [ "train_dataset, test_dataset = torchtext.datasets.AG_NEWS(root='./data')\n", "train_dataset = list(train_dataset)\n", "test_dataset = list(test_dataset)" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## トークン化\n", "\n", "次に、テキストを**数値**に変換し、それをテンソルとして表現できるようにする必要があります。単語レベルの表現を求める場合、以下の2つの作業が必要です:\n", "* **トークナイザー**を使用してテキストを**トークン**に分割する\n", "* それらのトークンの**語彙**を構築する\n" ] }, { "cell_type": "code", "execution_count": 4, "metadata": {}, "outputs": [ { "data": { "text/plain": [ "['he', 'said', 'hello']" ] }, "execution_count": 4, "metadata": {}, "output_type": "execute_result" } ], "source": [ "tokenizer = torchtext.data.utils.get_tokenizer('basic_english')\n", "tokenizer('He said: hello')" ] }, { "cell_type": "code", "execution_count": 5, "metadata": {}, "outputs": [], "source": [ "counter = collections.Counter()\n", "for (label, line) in train_dataset:\n", " counter.update(tokenizer(line))\n", "vocab = torchtext.vocab.vocab(counter, min_freq=1)" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "語彙を使用することで、トークン化された文字列を簡単に数値のセットにエンコードできます。\n" ] }, { "cell_type": "code", "execution_count": 19, "metadata": {}, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "Vocab size if 95810\n" ] }, { "data": { "text/plain": [ "[599, 3279, 97, 1220, 329, 225, 7368]" ] }, "execution_count": 19, "metadata": {}, "output_type": "execute_result" } ], "source": [ "vocab_size = len(vocab)\n", "print(f\"Vocab size if {vocab_size}\")\n", "\n", "stoi = vocab.get_stoi() # dict to convert tokens to indices\n", "\n", "def encode(x):\n", " return [stoi[s] for s in tokenizer(x)]\n", "\n", "encode('I love to play with my words')" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## Bag of Words テキスト表現\n", "\n", "言葉は意味を表すため、時には文中の順序に関係なく、個々の単語を見るだけでテキストの意味を推測できることがあります。例えば、ニュースを分類する際に、*weather* や *snow* といった単語は *天気予報* を示す可能性が高く、*stocks* や *dollar* といった単語は *金融ニュース* に関連すると考えられます。\n", "\n", "**Bag of Words** (BoW) ベクトル表現は、最も一般的に使用される伝統的なベクトル表現です。各単語はベクトルのインデックスにリンクされ、ベクトル要素には特定の文書内での単語の出現回数が含まれます。\n", "\n", "![Bag of Words ベクトル表現がメモリ内でどのように表現されるかを示す画像。](../../../../../translated_images/ja/bag-of-words-example.606fc1738f1d7ba9.webp) \n", "\n", "> **Note**: BoW は、テキスト内の個々の単語に対するすべてのワンホットエンコードされたベクトルの合計として考えることもできます。\n", "\n", "以下は、Scikit Learn の Python ライブラリを使用して Bag of Words 表現を生成する例です:\n" ] }, { "cell_type": "code", "execution_count": 7, "metadata": {}, "outputs": [ { "data": { "text/plain": [ "array([[1, 1, 0, 2, 0, 0, 0, 0, 0]], dtype=int64)" ] }, "execution_count": 7, "metadata": {}, "output_type": "execute_result" } ], "source": [ "from sklearn.feature_extraction.text import CountVectorizer\n", "vectorizer = CountVectorizer()\n", "corpus = [\n", " 'I like hot dogs.',\n", " 'The dog ran fast.',\n", " 'Its hot outside.',\n", " ]\n", "vectorizer.fit_transform(corpus)\n", "vectorizer.transform(['My dog likes hot dogs on a hot day.']).toarray()" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "AG_NEWSデータセットのベクトル表現からbag-of-wordsベクトルを計算するには、次の関数を使用できます:\n" ] }, { "cell_type": "code", "execution_count": 20, "metadata": {}, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "tensor([2., 1., 2., ..., 0., 0., 0.])\n" ] } ], "source": [ "vocab_size = len(vocab)\n", "\n", "def to_bow(text,bow_vocab_size=vocab_size):\n", " res = torch.zeros(bow_vocab_size,dtype=torch.float32)\n", " for i in encode(text):\n", " if i **注意:** ここでは、デフォルトの語彙サイズを指定するためにグローバル変数 `vocab_size` を使用しています。語彙サイズが非常に大きい場合が多いため、最も頻繁に使用される単語に語彙サイズを制限することができます。`vocab_size` の値を下げて以下のコードを実行し、その影響を確認してください。精度が多少低下することが予想されますが、劇的な変化はなく、性能が向上するはずです。\n" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## BoW分類器のトレーニング\n", "\n", "テキストのBag-of-Words表現を構築する方法を学んだので、その上で分類器をトレーニングしてみましょう。まず、トレーニング用のデータセットを変換する必要があります。すべての位置ベクトル表現をBag-of-Words表現に変換する方法です。これを実現するには、標準的なtorchの`DataLoader`に`collate_fn`パラメータとして`bowify`関数を渡します。\n" ] }, { "cell_type": "code", "execution_count": 21, "metadata": {}, "outputs": [], "source": [ "from torch.utils.data import DataLoader\n", "import numpy as np \n", "\n", "# this collate function gets list of batch_size tuples, and needs to \n", "# return a pair of label-feature tensors for the whole minibatch\n", "def bowify(b):\n", " return (\n", " torch.LongTensor([t[0]-1 for t in b]),\n", " torch.stack([to_bow(t[1]) for t in b])\n", " )\n", "\n", "train_loader = DataLoader(train_dataset, batch_size=16, collate_fn=bowify, shuffle=True)\n", "test_loader = DataLoader(test_dataset, batch_size=16, collate_fn=bowify, shuffle=True)" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "では、1つの線形層を含むシンプルな分類器ニューラルネットワークを定義しましょう。入力ベクトルのサイズは `vocab_size` に等しく、出力サイズはクラス数(4)に対応します。分類タスクを解いているため、最終的な活性化関数は `LogSoftmax()` です。\n" ] }, { "cell_type": "code", "execution_count": 22, "metadata": {}, "outputs": [], "source": [ "net = torch.nn.Sequential(torch.nn.Linear(vocab_size,4),torch.nn.LogSoftmax(dim=1))" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "これから標準的なPyTorchのトレーニングループを定義します。データセットが非常に大きいため、学習目的のために1エポックのみ、場合によっては1エポック未満でトレーニングを行います(`epoch_size`パラメータを指定することでトレーニングを制限できます)。また、トレーニング中に累積されたトレーニング精度を報告します。報告の頻度は`report_freq`パラメータを使用して指定します。\n" ] }, { "cell_type": "code", "execution_count": 24, "metadata": {}, "outputs": [], "source": [ "def train_epoch(net,dataloader,lr=0.01,optimizer=None,loss_fn = torch.nn.NLLLoss(),epoch_size=None, report_freq=200):\n", " optimizer = optimizer or torch.optim.Adam(net.parameters(),lr=lr)\n", " net.train()\n", " total_loss,acc,count,i = 0,0,0,0\n", " for labels,features in dataloader:\n", " optimizer.zero_grad()\n", " out = net(features)\n", " loss = loss_fn(out,labels) #cross_entropy(out,labels)\n", " loss.backward()\n", " optimizer.step()\n", " total_loss+=loss\n", " _,predicted = torch.max(out,1)\n", " acc+=(predicted==labels).sum()\n", " count+=len(labels)\n", " i+=1\n", " if i%report_freq==0:\n", " print(f\"{count}: acc={acc.item()/count}\")\n", " if epoch_size and count>epoch_size:\n", " break\n", " return total_loss.item()/count, acc.item()/count" ] }, { "cell_type": "code", "execution_count": 25, "metadata": {}, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "3200: acc=0.8028125\n", "6400: acc=0.8371875\n", "9600: acc=0.8534375\n", "12800: acc=0.85765625\n" ] }, { "data": { "text/plain": [ "(0.026090790722161722, 0.8620069296375267)" ] }, "execution_count": 25, "metadata": {}, "output_type": "execute_result" } ], "source": [ "train_epoch(net,train_loader,epoch_size=15000)" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## バイグラム、トライグラム、Nグラム\n", "\n", "バッグオブワーズアプローチの1つの制約は、一部の単語が複数の単語で構成される表現の一部である場合があることです。例えば、「hot dog」という単語は、他の文脈での「hot」や「dog」という単語とは全く異なる意味を持ちます。「hot」と「dog」を常に同じベクトルで表現すると、モデルが混乱する可能性があります。\n", "\n", "これに対処するために、**Nグラム表現**が文書分類の手法でよく使用されます。ここでは、各単語、2単語(バイグラム)、または3単語(トライグラム)の頻度が分類器を訓練するための有用な特徴となります。例えば、バイグラム表現では、元の単語に加えて、すべての単語ペアを語彙に追加します。\n", "\n", "以下は、Scikit Learnを使用してバイグラムのバッグオブワード表現を生成する方法の例です:\n" ] }, { "cell_type": "code", "execution_count": 26, "metadata": {}, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "Vocabulary:\n", " {'i': 7, 'like': 11, 'hot': 4, 'dogs': 2, 'i like': 8, 'like hot': 12, 'hot dogs': 5, 'the': 16, 'dog': 0, 'ran': 14, 'fast': 3, 'the dog': 17, 'dog ran': 1, 'ran fast': 15, 'its': 9, 'outside': 13, 'its hot': 10, 'hot outside': 6}\n" ] }, { "data": { "text/plain": [ "array([[1, 0, 1, 0, 2, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0]],\n", " dtype=int64)" ] }, "execution_count": 26, "metadata": {}, "output_type": "execute_result" } ], "source": [ "bigram_vectorizer = CountVectorizer(ngram_range=(1, 2), token_pattern=r'\\b\\w+\\b', min_df=1)\n", "corpus = [\n", " 'I like hot dogs.',\n", " 'The dog ran fast.',\n", " 'Its hot outside.',\n", " ]\n", "bigram_vectorizer.fit_transform(corpus)\n", "print(\"Vocabulary:\\n\",bigram_vectorizer.vocabulary_)\n", "bigram_vectorizer.transform(['My dog likes hot dogs on a hot day.']).toarray()\n" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "N-gramアプローチの主な欠点は、語彙サイズが非常に速く増加し始めることです。実際には、N-gram表現を*埋め込み*などの次元削減技術と組み合わせる必要があります。この点については次のユニットで説明します。\n", "\n", "**AG News**データセットでN-gram表現を使用するには、特別なN-gram語彙を構築する必要があります。\n" ] }, { "cell_type": "code", "execution_count": 27, "metadata": {}, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "Bigram vocabulary length = 1308842\n" ] } ], "source": [ "counter = collections.Counter()\n", "for (label, line) in train_dataset:\n", " l = tokenizer(line)\n", " counter.update(torchtext.data.utils.ngrams_iterator(l,ngrams=2))\n", " \n", "bi_vocab = torchtext.vocab.vocab(counter, min_freq=1)\n", "\n", "print(\"Bigram vocabulary length = \",len(bi_vocab))" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "次のコードを使って分類器を訓練することもできますが、それでは非常にメモリ効率が悪くなります。次のユニットでは、埋め込みを使用してバイグラム分類器を訓練します。\n", "\n", "> **Note:** テキスト内で指定された回数以上出現するn-gramのみを残すことができます。これにより、頻度の低いバイグラムが省略され、次元数が大幅に減少します。そのためには、`min_freq`パラメータを高い値に設定し、語彙の長さの変化を観察してください。\n" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## Term Frequency Inverse Document Frequency TF-IDF\n", "\n", "BoW表現では、単語の出現頻度は単語そのものに関係なく均等に重み付けされます。しかし、*a* や *in* などの頻繁に使われる単語は、専門用語に比べて分類において重要性が低いことは明らかです。実際、ほとんどのNLPタスクでは、ある単語が他の単語よりも重要である場合があります。\n", "\n", "**TF-IDF**は、**term frequency–inverse document frequency(単語頻度–逆文書頻度)**の略です。これはBoWの変形で、単語が文書に出現するかどうかを示す二値の0/1ではなく、コーパス内での単語の出現頻度に関連する浮動小数点値を使用します。\n", "\n", "より正式には、文書$j$内の単語$i$の重み$w_{ij}$は以下のように定義されます:\n", "$$\n", "w_{ij} = tf_{ij}\\times\\log({N\\over df_i})\n", "$$\n", "ここで、\n", "* $tf_{ij}$は文書$j$内で単語$i$が出現する回数、つまり以前に見たBoW値\n", "* $N$はコレクション内の文書数\n", "* $df_i$はコレクション全体で単語$i$を含む文書数\n", "\n", "TF-IDF値$w_{ij}$は、単語が文書内で出現する回数に比例して増加し、コーパス内でその単語を含む文書数によって調整されます。これにより、ある単語が他の単語よりも頻繁に出現する事実を補正することができます。例えば、単語がコレクション内の*すべて*の文書に出現する場合、$df_i=N$となり、$w_{ij}=0$となります。この場合、その単語は完全に無視されます。\n", "\n", "Scikit Learnを使用すれば、簡単にテキストのTF-IDFベクトル化を作成できます:\n" ] }, { "cell_type": "code", "execution_count": 28, "metadata": {}, "outputs": [ { "data": { "text/plain": [ "array([[0.43381609, 0. , 0.43381609, 0. , 0.65985664,\n", " 0.43381609, 0. , 0. , 0. , 0. ,\n", " 0. , 0. , 0. , 0. , 0. ,\n", " 0. ]])" ] }, "execution_count": 28, "metadata": {}, "output_type": "execute_result" } ], "source": [ "from sklearn.feature_extraction.text import TfidfVectorizer\n", "vectorizer = TfidfVectorizer(ngram_range=(1,2))\n", "vectorizer.fit_transform(corpus)\n", "vectorizer.transform(['My dog likes hot dogs on a hot day.']).toarray()" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "## 結論\n", "\n", "しかし、TF-IDF表現は異なる単語に頻度の重みを与えるものの、意味や順序を表現することはできません。有名な言語学者J.R.ファースが1935年に述べたように、「単語の完全な意味は常に文脈に依存しており、文脈を無視した意味の研究は真剣に受け止められるべきではない」のです。このコースの後半では、言語モデルを使用してテキストから文脈情報を捉える方法を学びます。\n" ] }, { "cell_type": "markdown", "metadata": {}, "source": [ "\n---\n\n**免責事項**: \nこの文書は、AI翻訳サービス [Co-op Translator](https://github.com/Azure/co-op-translator) を使用して翻訳されています。正確性を追求しておりますが、自動翻訳には誤りや不正確な部分が含まれる可能性があることをご承知ください。元の言語で記載された文書が正式な情報源とみなされるべきです。重要な情報については、専門の人間による翻訳を推奨します。この翻訳の使用に起因する誤解や誤認について、当方は一切の責任を負いません。\n" ] } ], "metadata": { "interpreter": { "hash": "16af2a8bbb083ea23e5e41c7f5787656b2ce26968575d8763f2c4b17f9cd711f" }, "kernelspec": { "display_name": "Python 3.8.12 ('py38')", "language": "python", "name": "python3" }, "language_info": { "codemirror_mode": { "name": "ipython", "version": 3 }, "file_extension": ".py", "mimetype": "text/x-python", "name": "python", "nbconvert_exporter": "python", "pygments_lexer": "ipython3", "version": "3.8.12" }, "coopTranslator": { "original_hash": "7b9040985e748e4e2d4c689892456ad7", "translation_date": "2025-08-30T10:40:32+00:00", "source_file": "lessons/5-NLP/13-TextRep/TextRepresentationPyTorch.ipynb", "language_code": "ja" } }, "nbformat": 4, "nbformat_minor": 2 }