{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# モジュールインポート","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport seaborn as sns\nimport matplotlib.pyplot as plt","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:25:45.678992Z","iopub.execute_input":"2022-07-29T07:25:45.679517Z","iopub.status.idle":"2022-07-29T07:25:46.315297Z","shell.execute_reply.started":"2022-07-29T07:25:45.679399Z","shell.execute_reply":"2022-07-29T07:25:46.314122Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## GPU設定","metadata":{}},{"cell_type":"code","source":"import torch\nfrom torch import cuda\n\ndevice = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:25:46.317476Z","iopub.execute_input":"2022-07-29T07:25:46.318127Z","iopub.status.idle":"2022-07-29T07:25:48.138761Z","shell.execute_reply.started":"2022-07-29T07:25:46.318085Z","shell.execute_reply":"2022-07-29T07:25:48.137567Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## データ確認","metadata":{}},{"cell_type":"code","source":"train = pd.read_csv(\"../input/nlp-getting-started/train.csv\")\ntest = pd.read_csv(\"../input/nlp-getting-started/test.csv\")\ntrain.head(5)","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:25:48.140443Z","iopub.execute_input":"2022-07-29T07:25:48.141134Z","iopub.status.idle":"2022-07-29T07:25:48.228547Z","shell.execute_reply.started":"2022-07-29T07:25:48.141088Z","shell.execute_reply":"2022-07-29T07:25:48.227562Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train['text'][0]","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:25:48.231948Z","iopub.execute_input":"2022-07-29T07:25:48.232306Z","iopub.status.idle":"2022-07-29T07:25:48.243238Z","shell.execute_reply.started":"2022-07-29T07:25:48.232278Z","shell.execute_reply":"2022-07-29T07:25:48.242065Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x = train.target.value_counts()\nsns.barplot(x.index, x)\nplt.gca().set_ylabel('samples')","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:25:48.244669Z","iopub.execute_input":"2022-07-29T07:25:48.245939Z","iopub.status.idle":"2022-07-29T07:25:48.440688Z","shell.execute_reply.started":"2022-07-29T07:25:48.245896Z","shell.execute_reply":"2022-07-29T07:25:48.439514Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## あるキーワードを含むツイートのラベルごとの数を計測","metadata":{}},{"cell_type":"code","source":"accident = train[train['keyword'].str.match('accident', na=False)]\nax = sns.countplot(x='target', data=accident)\nax.set_title(\"accident\")","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:25:48.442482Z","iopub.execute_input":"2022-07-29T07:25:48.442906Z","iopub.status.idle":"2022-07-29T07:25:48.626693Z","shell.execute_reply.started":"2022-07-29T07:25:48.442866Z","shell.execute_reply":"2022-07-29T07:25:48.625485Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"accident = train[train['keyword'].str.match('ambulance', na=False)]\nax = sns.countplot(x='target', data=accident)\nax.set_title(\"ambulance\")\nax.set_yticklabels([0, '', 5, '', 10, '', 15, '', 20])","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:25:48.628332Z","iopub.execute_input":"2022-07-29T07:25:48.628691Z","iopub.status.idle":"2022-07-29T07:25:48.812613Z","shell.execute_reply.started":"2022-07-29T07:25:48.628663Z","shell.execute_reply":"2022-07-29T07:25:48.811608Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"accident = train[train['keyword'].str.match('collapse', na=False)]\nax = sns.countplot(x='target', data=accident)\nax.set_title(\"collapse\")","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:25:48.814624Z","iopub.execute_input":"2022-07-29T07:25:48.815303Z","iopub.status.idle":"2022-07-29T07:25:49.001654Z","shell.execute_reply.started":"2022-07-29T07:25:48.815262Z","shell.execute_reply":"2022-07-29T07:25:49.000511Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## URL、絵文字、HTML、記号を除去","metadata":{}},{"cell_type":"code","source":"df = pd.concat([train, test])\ndf.shape","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:25:49.003405Z","iopub.execute_input":"2022-07-29T07:25:49.003859Z","iopub.status.idle":"2022-07-29T07:25:49.016736Z","shell.execute_reply.started":"2022-07-29T07:25:49.003818Z","shell.execute_reply":"2022-07-29T07:25:49.015462Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import re\n\ndef remove_URL(text):\n    return re.sub(r'https?://\\S+|www\\.\\S+', r'', text)","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:25:49.021867Z","iopub.execute_input":"2022-07-29T07:25:49.022612Z","iopub.status.idle":"2022-07-29T07:25:49.029562Z","shell.execute_reply.started":"2022-07-29T07:25:49.022566Z","shell.execute_reply":"2022-07-29T07:25:49.028540Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"url = \"https://ohke.hateblo.jp/entry/2019/02/09/141500\"\nprint(\"url:\", remove_URL(url))","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:25:49.031540Z","iopub.execute_input":"2022-07-29T07:25:49.032485Z","iopub.status.idle":"2022-07-29T07:25:49.043610Z","shell.execute_reply.started":"2022-07-29T07:25:49.032442Z","shell.execute_reply":"2022-07-29T07:25:49.042528Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df['text'] = df['text'].apply(lambda x: remove_URL(x))\ndf","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:25:49.046413Z","iopub.execute_input":"2022-07-29T07:25:49.046756Z","iopub.status.idle":"2022-07-29T07:25:49.116523Z","shell.execute_reply.started":"2022-07-29T07:25:49.046716Z","shell.execute_reply":"2022-07-29T07:25:49.115138Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install demoji","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:25:49.118380Z","iopub.execute_input":"2022-07-29T07:25:49.119011Z","iopub.status.idle":"2022-07-29T07:26:01.533930Z","shell.execute_reply.started":"2022-07-29T07:25:49.118968Z","shell.execute_reply":"2022-07-29T07:26:01.532563Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import demoji\n\ndef remove_emoji(text):\n    return demoji.replace(string=text, repl='')","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:26:01.536845Z","iopub.execute_input":"2022-07-29T07:26:01.537644Z","iopub.status.idle":"2022-07-29T07:26:01.549920Z","shell.execute_reply.started":"2022-07-29T07:26:01.537597Z","shell.execute_reply":"2022-07-29T07:26:01.548727Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df['text'] = df['text'].apply(lambda x: remove_emoji(x))\ndf","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:26:01.552927Z","iopub.execute_input":"2022-07-29T07:26:01.554303Z","iopub.status.idle":"2022-07-29T07:26:08.402753Z","shell.execute_reply.started":"2022-07-29T07:26:01.554248Z","shell.execute_reply":"2022-07-29T07:26:08.401538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def remove_html(text):\n    return re.sub(r'<.*?>', r'', text)","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:26:08.404574Z","iopub.execute_input":"2022-07-29T07:26:08.405315Z","iopub.status.idle":"2022-07-29T07:26:08.411425Z","shell.execute_reply.started":"2022-07-29T07:26:08.405254Z","shell.execute_reply":"2022-07-29T07:26:08.410140Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df['text'] = df['text'].apply(lambda x: remove_html(x))\ndf","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:26:08.413389Z","iopub.execute_input":"2022-07-29T07:26:08.414153Z","iopub.status.idle":"2022-07-29T07:26:08.459227Z","shell.execute_reply.started":"2022-07-29T07:26:08.414105Z","shell.execute_reply":"2022-07-29T07:26:08.457998Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import string\n\ndef remove_punctuation(text):\n    table = str.maketrans('', '', string.punctuation)\n    return text.translate(table)","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:26:08.461004Z","iopub.execute_input":"2022-07-29T07:26:08.461725Z","iopub.status.idle":"2022-07-29T07:26:08.468117Z","shell.execute_reply.started":"2022-07-29T07:26:08.461681Z","shell.execute_reply":"2022-07-29T07:26:08.466843Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df['text'] = df['text'].apply(lambda x: remove_punctuation(x))\ndf","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:26:08.469761Z","iopub.execute_input":"2022-07-29T07:26:08.471342Z","iopub.status.idle":"2022-07-29T07:26:08.562777Z","shell.execute_reply.started":"2022-07-29T07:26:08.471281Z","shell.execute_reply":"2022-07-29T07:26:08.561567Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = df[:train.shape[0]]\ntest = df[train.shape[0]:]","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:26:08.564859Z","iopub.execute_input":"2022-07-29T07:26:08.565307Z","iopub.status.idle":"2022-07-29T07:26:08.572577Z","shell.execute_reply.started":"2022-07-29T07:26:08.565261Z","shell.execute_reply":"2022-07-29T07:26:08.570557Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"f_int = lambda x: int(x)\ntrain['target'] = train['target'].map(f_int)\n#test['target'] = test['target'].map(f_int)","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:26:08.574184Z","iopub.execute_input":"2022-07-29T07:26:08.576225Z","iopub.status.idle":"2022-07-29T07:26:08.593070Z","shell.execute_reply.started":"2022-07-29T07:26:08.576192Z","shell.execute_reply":"2022-07-29T07:26:08.591988Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:26:08.594769Z","iopub.execute_input":"2022-07-29T07:26:08.595423Z","iopub.status.idle":"2022-07-29T07:26:08.611568Z","shell.execute_reply.started":"2022-07-29T07:26:08.595383Z","shell.execute_reply":"2022-07-29T07:26:08.610497Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport transformers\nfrom torch.utils.data import Dataset, DataLoader\nfrom transformers import BertTokenizer, BertModel","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:26:08.613034Z","iopub.execute_input":"2022-07-29T07:26:08.613687Z","iopub.status.idle":"2022-07-29T07:26:13.846336Z","shell.execute_reply.started":"2022-07-29T07:26:08.613647Z","shell.execute_reply":"2022-07-29T07:26:13.845162Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tokenizer = BertTokenizer.from_pretrained('bert-base-uncased')","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:26:13.847687Z","iopub.execute_input":"2022-07-29T07:26:13.848056Z","iopub.status.idle":"2022-07-29T07:26:16.998992Z","shell.execute_reply.started":"2022-07-29T07:26:13.848019Z","shell.execute_reply":"2022-07-29T07:26:16.997890Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(tokenizer.tokenize(train['text'][3261]))\nprint(tokenizer.convert_tokens_to_ids(tokenizer.tokenize(train['text'][3261])))","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:26:17.000671Z","iopub.execute_input":"2022-07-29T07:26:17.001079Z","iopub.status.idle":"2022-07-29T07:26:17.012664Z","shell.execute_reply.started":"2022-07-29T07:26:17.001037Z","shell.execute_reply":"2022-07-29T07:26:17.011411Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"max_len = []\nfor sent in df['text']:\n    token_words = tokenizer.tokenize(sent)\n    max_len.append(len(token_words))\nprint('最大単語数: ', max(max_len))\nprint('上記の最大単語数にSpecial token（[CLS], [SEP]）の+2をした値が最大単語数')","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:26:17.014593Z","iopub.execute_input":"2022-07-29T07:26:17.015554Z","iopub.status.idle":"2022-07-29T07:26:23.465817Z","shell.execute_reply.started":"2022-07-29T07:26:17.015509Z","shell.execute_reply":"2022-07-29T07:26:23.464687Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_sent = train.text.values\ntrain_lab = train.target.values\ntest_sent = test.text.values\ntest_lab = test.target.values","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:26:23.467407Z","iopub.execute_input":"2022-07-29T07:26:23.468020Z","iopub.status.idle":"2022-07-29T07:26:23.475355Z","shell.execute_reply.started":"2022-07-29T07:26:23.467979Z","shell.execute_reply":"2022-07-29T07:26:23.474095Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(train_lab)","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:26:23.483177Z","iopub.execute_input":"2022-07-29T07:26:23.483993Z","iopub.status.idle":"2022-07-29T07:26:23.490635Z","shell.execute_reply.started":"2022-07-29T07:26:23.483946Z","shell.execute_reply":"2022-07-29T07:26:23.489472Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_input_ids = []\ntrain_attention_masks = []\n\n# 1文づつ処理\nfor sent in train_sent:\n    encoded_dict = tokenizer.encode_plus(\n                        sent,                      \n                        add_special_tokens = True, # Special Tokenの追加\n                        max_length = 55,           # 文章の長さを固定（Padding/Trancatinating）\n                        pad_to_max_length = True,# PADDINGで埋める\n                        return_attention_mask = True,   # Attention maksの作成\n                        return_tensors = 'pt',     #  Pytorch tensorsで返す\n                   )\n\n    # 単語IDを取得    \n    train_input_ids.append(encoded_dict['input_ids'])\n\n    # Attention　maskの取得\n    train_attention_masks.append(encoded_dict['attention_mask'])\n\n# リストに入ったtensorを縦方向（dim=0）へ結合\ntrain_input_ids = torch.cat(train_input_ids, dim=0)\ntrain_attention_masks = torch.cat(train_attention_masks, dim=0)\n\n# tenosor型に変換\ntrain_lab = torch.tensor(train_lab)\n\n# 確認\nprint('Original: ', train_sent[0])\nprint('Token IDs:', train_input_ids[0])","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:26:23.492107Z","iopub.execute_input":"2022-07-29T07:26:23.492888Z","iopub.status.idle":"2022-07-29T07:26:29.998451Z","shell.execute_reply.started":"2022-07-29T07:26:23.492844Z","shell.execute_reply":"2022-07-29T07:26:29.997007Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_input_ids = []\ntest_attention_masks = []\n\n# 1文づつ処理\nfor sent in test_sent:\n    encoded_dict = tokenizer.encode_plus(\n                        sent,                      \n                        add_special_tokens = True, # Special Tokenの追加\n                        max_length = 55,           # 文章の長さを固定（Padding/Trancatinating）\n                        pad_to_max_length = True,# PADDINGで埋める\n                        return_attention_mask = True,   # Attention maksの作成\n                        return_tensors = 'pt',     #  Pytorch tensorsで返す\n                   )\n\n    # 単語IDを取得    \n    test_input_ids.append(encoded_dict['input_ids'])\n\n    # Attention　maskの取得\n    test_attention_masks.append(encoded_dict['attention_mask'])\n\n# リストに入ったtensorを縦方向（dim=0）へ結合\ntest_input_ids = torch.cat(test_input_ids, dim=0)\ntest_attention_masks = torch.cat(test_attention_masks, dim=0)\n\n# tenosor型に変換\ntest_lab = torch.tensor(test_lab)\n\n# 確認\nprint('Original: ', test_sent[0])\nprint('Token IDs:', test_input_ids[0])","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:26:30.000546Z","iopub.execute_input":"2022-07-29T07:26:30.001031Z","iopub.status.idle":"2022-07-29T07:26:32.656221Z","shell.execute_reply.started":"2022-07-29T07:26:30.000984Z","shell.execute_reply":"2022-07-29T07:26:32.654813Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import TensorDataset\nfrom torch.utils.data import DataLoader, RandomSampler, SequentialSampler\n\ntrain_dataset = TensorDataset(train_input_ids, train_attention_masks, train_lab)\ntest_dataset = TensorDataset(test_input_ids, test_attention_masks, test_lab)\n\nprint('訓練データ数：{}'.format(len(train_dataset)))\nprint('テストデータ数:　{} '.format(len(test_dataset)))\n\n# データローダーの作成\nbatch_size = 32\n\n# 訓練データローダー\ntrain_dataloader = DataLoader(\n            train_dataset,  \n            sampler = RandomSampler(train_dataset), # ランダムにデータを取得してバッチ化\n            batch_size = batch_size\n        )\n\n# 検証データローダー\ntest_dataloader = DataLoader(\n            test_dataset, \n            sampler = SequentialSampler(test_dataset), # 順番にデータを取得してバッチ化\n            batch_size = batch_size\n        )","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:26:32.658335Z","iopub.execute_input":"2022-07-29T07:26:32.658838Z","iopub.status.idle":"2022-07-29T07:26:32.671147Z","shell.execute_reply.started":"2022-07-29T07:26:32.658790Z","shell.execute_reply":"2022-07-29T07:26:32.669714Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(train_dataloader)","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:26:32.672981Z","iopub.execute_input":"2022-07-29T07:26:32.673717Z","iopub.status.idle":"2022-07-29T07:26:32.691162Z","shell.execute_reply.started":"2022-07-29T07:26:32.673670Z","shell.execute_reply":"2022-07-29T07:26:32.689887Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from transformers import BertForSequenceClassification, AdamW, BertConfig\n\n# BertForSequenceClassification 学習済みモデルのロード\nmodel = BertForSequenceClassification.from_pretrained(\n    \"bert-base-uncased\", # Pre trainedモデルの指定\n    num_labels = 2, # ラベル数（今回はBinayなので2、数値を増やせばマルチラベルも対応可）\n    output_attentions = False, # アテンションベクトルを出力するか\n    output_hidden_states = False, # 隠れ層を出力するか\n)\n\n# モデルをGPUへ転送\nmodel.cuda()","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:26:32.694219Z","iopub.execute_input":"2022-07-29T07:26:32.695615Z","iopub.status.idle":"2022-07-29T07:26:51.991180Z","shell.execute_reply.started":"2022-07-29T07:26:32.695539Z","shell.execute_reply":"2022-07-29T07:26:51.989010Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.nn import functional as F\n\n# 最適化手法の設定\noptimizer = AdamW(model.parameters(), lr=2e-5)\n\n# 訓練パートの定義\ndef training(model):\n    model.train() # 訓練モードで実行\n    train_loss = 0\n    for batch in train_dataloader:# train_dataloaderはword_id, mask, labelを出力する点に注意\n        b_input_ids = batch[0].to(device)\n        b_input_mask = batch[1].to(device)\n        b_labels = batch[2].to(device)\n        optimizer.zero_grad()\n        outputs = model(b_input_ids, \n                             token_type_ids=None, \n                             attention_mask=b_input_mask, \n                             labels=b_labels)\n        loss = F.cross_entropy(outputs.logits, b_labels)\n        loss.backward()\n        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        optimizer.step()\n        train_loss += loss.item()\n    return train_loss\n\n# # テストパートの定義\n# def validation(model):\n#     model.eval()# 訓練モードをオフ\n#     val_loss = 0\n#     with torch.no_grad(): # 勾配を計算しない\n#         for batch in test_dataloader:\n#             b_input_ids = batch[0]#.to(device)\n#             b_input_mask = batch[1]#.to(device)\n#             b_labels = batch[2]#.to(device)\n#             with torch.no_grad():        \n#                 (loss, logits) = model(b_input_ids, \n#                                     token_type_ids=None, \n#                                     attention_mask=b_input_mask,\n#                                     labels=b_labels)\n#             val_loss += loss.item()\n#     return val_loss","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:26:51.992658Z","iopub.execute_input":"2022-07-29T07:26:51.993350Z","iopub.status.idle":"2022-07-29T07:26:52.014428Z","shell.execute_reply.started":"2022-07-29T07:26:51.993290Z","shell.execute_reply":"2022-07-29T07:26:52.013348Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 学習の実行\nmax_epoch = 4\ntrain_loss_ = []\ntest_loss_ = []\n\nfor epoch in range(max_epoch):\n    train_ = training(model)\n    #test_ = validation(model)\n    train_loss_.append(train_)\n    #test_loss_.append(test_)","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:26:52.016372Z","iopub.execute_input":"2022-07-29T07:26:52.017066Z","iopub.status.idle":"2022-07-29T07:29:42.278028Z","shell.execute_reply.started":"2022-07-29T07:26:52.017022Z","shell.execute_reply":"2022-07-29T07:29:42.276688Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds_list = []\n\nmodel.eval()# 訓練モードをオフ\nfor batch in test_dataloader:\n    b_input_ids = batch[0].to(device)\n    b_input_mask = batch[1].to(device)\n    b_labels = batch[2].to(device)\n    with torch.no_grad():   \n        # 学習済みモデルによる予測結果をpredsで取得     \n        preds = model(b_input_ids, \n                            token_type_ids=None, \n                            attention_mask=b_input_mask)\n    preds_list.append(preds)","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:29:42.280943Z","iopub.execute_input":"2022-07-29T07:29:42.281851Z","iopub.status.idle":"2022-07-29T07:29:47.378530Z","shell.execute_reply.started":"2022-07-29T07:29:42.281797Z","shell.execute_reply":"2022-07-29T07:29:47.377508Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predict_label = []\n\nfor pred in preds_list:\n    for lo in pred.logits.cpu().numpy():\n        predict_label.append(np.argmax(lo))","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:29:47.380035Z","iopub.execute_input":"2022-07-29T07:29:47.380763Z","iopub.status.idle":"2022-07-29T07:29:47.428805Z","shell.execute_reply.started":"2022-07-29T07:29:47.380721Z","shell.execute_reply":"2022-07-29T07:29:47.427785Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_sub = pd.read_csv('../input/nlp-getting-started/sample_submission.csv')","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:29:47.430551Z","iopub.execute_input":"2022-07-29T07:29:47.430988Z","iopub.status.idle":"2022-07-29T07:29:47.444477Z","shell.execute_reply.started":"2022-07-29T07:29:47.430948Z","shell.execute_reply":"2022-07-29T07:29:47.443361Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 比較しやすい様にpd.dataframeへ整形\nimport pandas as pd\n\nsub=pd.DataFrame({'id': sample_sub['id'].values.tolist(), 'target': predict_label})\nsub.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:29:47.446152Z","iopub.execute_input":"2022-07-29T07:29:47.446562Z","iopub.status.idle":"2022-07-29T07:29:47.465565Z","shell.execute_reply.started":"2022-07-29T07:29:47.446522Z","shell.execute_reply":"2022-07-29T07:29:47.464638Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub.head()","metadata":{"execution":{"iopub.status.busy":"2022-07-29T07:29:47.468162Z","iopub.execute_input":"2022-07-29T07:29:47.468794Z","iopub.status.idle":"2022-07-29T07:29:47.480418Z","shell.execute_reply.started":"2022-07-29T07:29:47.468755Z","shell.execute_reply":"2022-07-29T07:29:47.479364Z"},"trusted":true},"execution_count":null,"outputs":[]}]}