{"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":"code","source":"!pip install mojimoji\n!pip install neologdn","metadata":{"execution":{"iopub.status.busy":"2022-07-13T09:51:45.158433Z","iopub.execute_input":"2022-07-13T09:51:45.159110Z","iopub.status.idle":"2022-07-13T09:52:19.197808Z","shell.execute_reply.started":"2022-07-13T09:51:45.159018Z","shell.execute_reply":"2022-07-13T09:52:19.196663Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport re\nimport string\nfrom typing import Tuple\nfrom collections import Counter\nimport pandas as pd\nimport numpy as np\nimport nltk\nfrom time import sleep\nimport matplotlib.pyplot as plt\n\nimport random\nimport glob\nfrom tqdm import tqdm\nimport json\nimport emoji\nimport mojimoji\nimport neologdn\nimport seaborn as sns\n\nimport torch\nfrom torch.utils.data import Dataset, DataLoader,random_split,RandomSampler,SequentialSampler,TensorDataset\nfrom transformers import BertJapaneseTokenizer, BertForSequenceClassification\nfrom pytorch_lightning.callbacks import ModelCheckpoint\nimport sys\nimport pandas as pd\nimport pytorch_lightning as pl\nfrom sklearn.model_selection import train_test_split","metadata":{"execution":{"iopub.status.busy":"2022-07-13T09:52:19.200221Z","iopub.execute_input":"2022-07-13T09:52:19.200932Z","iopub.status.idle":"2022-07-13T09:52:27.369826Z","shell.execute_reply.started":"2022-07-13T09:52:19.200888Z","shell.execute_reply":"2022-07-13T09:52:27.368792Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"execution":{"iopub.status.busy":"2022-07-13T09:52:27.371318Z","iopub.execute_input":"2022-07-13T09:52:27.372100Z","iopub.status.idle":"2022-07-13T09:52:27.381601Z","shell.execute_reply.started":"2022-07-13T09:52:27.372063Z","shell.execute_reply":"2022-07-13T09:52:27.380406Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.listdir(\"/kaggle/input/\")","metadata":{"execution":{"iopub.status.busy":"2022-07-13T09:52:27.384763Z","iopub.execute_input":"2022-07-13T09:52:27.385317Z","iopub.status.idle":"2022-07-13T09:52:27.395780Z","shell.execute_reply.started":"2022-07-13T09:52:27.385291Z","shell.execute_reply":"2022-07-13T09:52:27.394683Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#データの取得\ntweet = pd.read_csv('/kaggle/input/nlp-getting-started/train.csv')\ntest = pd.read_csv('/kaggle/input/nlp-getting-started/test.csv')\ntweet.head(10)\nprint(len(test))","metadata":{"execution":{"iopub.status.busy":"2022-07-13T09:52:27.398748Z","iopub.execute_input":"2022-07-13T09:52:27.399268Z","iopub.status.idle":"2022-07-13T09:52:27.455990Z","shell.execute_reply.started":"2022-07-13T09:52:27.399240Z","shell.execute_reply":"2022-07-13T09:52:27.454972Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#stopwords of Englishの取得\nnltk.download('stopwords')\nnltk.download('tagsets')\n\nstop_words = nltk.corpus.stopwords.words('english')\ntag_dict = nltk.data.load('help/tagsets/upenn_tagset.pickle')","metadata":{"execution":{"iopub.status.busy":"2022-07-13T09:52:27.457563Z","iopub.execute_input":"2022-07-13T09:52:27.457934Z","iopub.status.idle":"2022-07-13T09:52:27.620639Z","shell.execute_reply.started":"2022-07-13T09:52:27.457895Z","shell.execute_reply":"2022-07-13T09:52:27.619593Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def afk_lock(sec: float = 1.0):\n    ''' ユーザーの動作無し時にシャットダウンすることを防ぐ'''\n    c = 0\n    while True:\n        print(c, end='\\r')\n        c += sec\n        sleep(sec)","metadata":{"execution":{"iopub.status.busy":"2022-07-13T09:52:27.623721Z","iopub.execute_input":"2022-07-13T09:52:27.624023Z","iopub.status.idle":"2022-07-13T09:52:27.630780Z","shell.execute_reply.started":"2022-07-13T09:52:27.623997Z","shell.execute_reply":"2022-07-13T09:52:27.629789Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"データの内容を確認したりする作業たち","metadata":{}},{"cell_type":"code","source":"x=tweet.target.value_counts()\nsns.barplot(x.index,x)\nplt.gca().set_ylabel('samples')","metadata":{"execution":{"iopub.status.busy":"2022-07-13T09:52:27.632338Z","iopub.execute_input":"2022-07-13T09:52:27.632779Z","iopub.status.idle":"2022-07-13T09:52:27.833075Z","shell.execute_reply.started":"2022-07-13T09:52:27.632728Z","shell.execute_reply":"2022-07-13T09:52:27.832023Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#df[\"text\"]ノカタチ\n#ストップワードの個数のカウント\ndef get_stopword_count(ds):\n    pattern_stop_words = r'\\b({})\\b'.format('|'.join(stop_words))\n    return ds.str.count(pattern_stop_words) \n\n#ハイパーリンクのカウント\ndef get_hyperlink_count(ds):\n    pattern_hyperlink = r'(http)(s)?(://)'\n    return ds.str.count(pattern_hyperlink)\n\n#ハッシュタグのカウント\ndef get_hashtags_count(ds):\n    pattern_hashtags = '([#])'   \n    return ds.str.count(pattern_hashtags)\n\n#メンションのカウント\ndef get_mentions_count(ds):\n    pattern_mentions = '([@])'    \n    return ds.str.count(pattern_mentions)\n\n#文字数の長さの取得\ndef get_length(ds):\n    return ds.str.len()\n\n#単語数のカウント\ndef get_unique_count(ds):\n     return ds.apply(lambda x: len(set(x.split())))\n    \n#エンタイトルの数のカウント\ndef get_named_entities_count(ds):\n    def get_named_entities(quote: str):\n        ''' Extracts named entities from quote'''\n        words = nltk.word_tokenize(quote)\n        tags = nltk.pos_tag(words)\n        tree = nltk.ne_chunk(tags, binary=True)\n        return len(set(\n            ' '.join(i[0] for i in t)\n            for t in tree\n            if hasattr(t, 'label') and t.label() == 'NE'\n        ))\n    \n    return ds.apply(lambda x: get_named_entities(x))\n\n#補題の数のカウント\ndef get_lemma_count(ds):\n    lemmatizer = nltk.stem.WordNetLemmatizer()\n    series = ds.apply(nltk.word_tokenize)\n    series = ds.apply(lambda x: len([lemmatizer.lemmatize(w) for w in x]))\n\n    return series\n\n#dfノカタチ\ndef get_series_part_of_speech(df):\n    ''' POS(Part of speech)のタイプごとにカウントする'''\n    series = df.apply(lambda x: nltk.word_tokenize(x))\n    series = series.apply(lambda x: nltk.pos_tag(x))\n    pos_df = pd.json_normalize(series.apply(lambda x: Counter(elem[1] for elem in x)))\n\n    return pos_df\n\ndef fix_unseen_pos(df):\n    ''' 一部のPOSタグが送信セットに含まれていない可能性がある'''\n    for col in tag_dict.keys():\n        if col not in df.columns:\n            df[col] = 0\n\n    return df","metadata":{"execution":{"iopub.status.busy":"2022-07-13T09:52:27.834914Z","iopub.execute_input":"2022-07-13T09:52:27.835541Z","iopub.status.idle":"2022-07-13T09:52:27.851421Z","shell.execute_reply.started":"2022-07-13T09:52:27.835503Z","shell.execute_reply":"2022-07-13T09:52:27.850233Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def preprocess(df,target_column):\n    raw_series = df[target_column]\n    \n    df['stopword_count'] = get_stopword_count(raw_series)\n    df['hyperlink_count'] = get_hyperlink_count(raw_series)\n    df['hashtag_count'] = get_hashtags_count(raw_series)\n    df['mention_count'] = get_mentions_count(raw_series)\n    df['length'] = get_length(raw_series)\n    df['unique_count'] = get_unique_count(raw_series)\n    #df['named_entities_count'] = get_named_entities_count(raw_series)\n    #df['lemma_count'] = get_lemma_count(raw_series)\n    \n    df_notdisaster = df[df[\"target\"] == 0]\n    df_disaster = df[df[\"target\"] == 1]\n    \n    return df, df_disaster, df_notdisaster\n\n#devide not disaster or disaster tweets\ndf = tweet\ndf, df_disaster, df_notdisaster = preprocess(df,\"text\")","metadata":{"execution":{"iopub.status.busy":"2022-07-13T09:52:27.857501Z","iopub.execute_input":"2022-07-13T09:52:27.857885Z","iopub.status.idle":"2022-07-13T09:52:28.102099Z","shell.execute_reply.started":"2022-07-13T09:52:27.857842Z","shell.execute_reply":"2022-07-13T09:52:28.101129Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#show data stopwords\nfig,(ax1,ax2)=plt.subplots(1,2,figsize=(10,5))\nx = df_notdisaster.stopword_count.value_counts()\nsns.barplot(x.index,x,ax=ax1)\nax1.set_title('not disaster stopwords')\n\ny = df_disaster.stopword_count.value_counts()\nsns.barplot(y.index,y,ax=ax2)\nax2.set_title('disaster stopwords')","metadata":{"execution":{"iopub.status.busy":"2022-07-13T09:52:28.103500Z","iopub.execute_input":"2022-07-13T09:52:28.103866Z","iopub.status.idle":"2022-07-13T09:52:28.565145Z","shell.execute_reply.started":"2022-07-13T09:52:28.103829Z","shell.execute_reply":"2022-07-13T09:52:28.564159Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#show data hyperlinks\nfig,(ax1,ax2)=plt.subplots(1,2,figsize=(10,5))\nx = df_notdisaster.hyperlink_count.value_counts()\nsns.barplot(x.index,x,ax=ax1)\nax1.set_title('not disaster hyperlink')\n\ny = df_disaster.hyperlink_count.value_counts()\nsns.barplot(y.index,y,ax=ax2)\nax2.set_title('disaster hyperlink')","metadata":{"execution":{"iopub.status.busy":"2022-07-13T09:52:28.566468Z","iopub.execute_input":"2022-07-13T09:52:28.567617Z","iopub.status.idle":"2022-07-13T09:52:28.866107Z","shell.execute_reply.started":"2022-07-13T09:52:28.567576Z","shell.execute_reply":"2022-07-13T09:52:28.865150Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#show data hashtags\nfig,(ax1,ax2)=plt.subplots(1,2,figsize=(10,5))\nx = df_notdisaster.hashtag_count.value_counts()\nsns.barplot(x.index,x,ax=ax1)\nax1.set_title('not disaster hashtag')\n\ny = df_disaster.hashtag_count.value_counts()\nsns.barplot(y.index,y,ax=ax2)\nax2.set_title('disaster hashtag')","metadata":{"execution":{"iopub.status.busy":"2022-07-13T09:52:28.867675Z","iopub.execute_input":"2022-07-13T09:52:28.868369Z","iopub.status.idle":"2022-07-13T09:52:29.226638Z","shell.execute_reply.started":"2022-07-13T09:52:28.868322Z","shell.execute_reply":"2022-07-13T09:52:29.225715Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#show data mentions\nfig,(ax1,ax2)=plt.subplots(1,2,figsize=(10,5))\nx = df_notdisaster.mention_count.value_counts()\nsns.barplot(x.index,x,ax=ax1)\nax1.set_title(\"not disaster mentions\")\n\ny = df_disaster.mention_count.value_counts()\nsns.barplot(y.index,y,ax=ax2)\nax2.set_title(\"disaster mentions\")","metadata":{"execution":{"iopub.status.busy":"2022-07-13T09:52:29.228177Z","iopub.execute_input":"2022-07-13T09:52:29.228524Z","iopub.status.idle":"2022-07-13T09:52:29.530919Z","shell.execute_reply.started":"2022-07-13T09:52:29.228488Z","shell.execute_reply":"2022-07-13T09:52:29.529988Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"データをきれいにする処理たち","metadata":{}},{"cell_type":"code","source":"#TweetLが元になっているものを用いた正規化を行う\nclass CleansingTweets:\n\n    # replace and\\s to space\n    def cleansing_space(self, text):\n        return re.sub(\"\\u3000|\\s\", \" \", text)\n\n    #repeat string abbrebiation\n    def cleansing_repeat(self,text):\n        text = re.sub(\"!{2,}\",\"!\",text)\n        text = re.sub(\"\\?{2,}\",\"?\",text)\n        text = re.sub(\"(!\\?){2,}\",\"!?\",text)\n        text = re.sub(\"w{2,}\",\"w\",text)\n        text = re.sub(\"…{2,}\",\"…\",text)\n        text = text.replace(\"〝\",\"\\\"\")\n        \n        return text\n    \n    # remove hashtags\n    def cleansing_hash(self, text):\n        return re.sub(\"#[^\\s]+\", \"\", text)\n\n    # remove URLs\n    def cleansing_url(self, text):\n        return re.sub(r\"(https?|ftp)(:\\/\\/[-_\\.!~*\\'()a-zA-Z0-9;\\/?:\\@&=\\+$,%#]+)\", \"\" , text)\n\n    # remove pictographs\n    def cleansing_emoji(self, text):\n        return ''.join(c for c in text if c not in emoji.UNICODE_EMOJI)\n\n    # remove mentions\n    def cleansing_username(self, text):\n        return re.sub(r\"@([A-Za-z0-9_]+) \", \"\", text)\n\n    #  remove image strings\n    def cleansing_picture(self, text):\n        return re.sub(r\"pic.twitter.com/[-_\\.!~*\\'()a-zA-Z0-9;\\/?:\\@&=\\+$,%#]*\", \"\" , text)\n\n    # unify characters\n    def cleansing_unity(self, text):\n        text = text.lower()\n        text = mojimoji.zen_to_han(text, kana=True)\n        text = mojimoji.han_to_zen(text, digit=False, ascii=False)\n        return text\n\n    # replace number to zero\n    def cleansing_num(self, text):\n        text = re.sub(r'\\d+', \"0\", text)\n        return text\n\n    # remove rt\n    def cleansing_rt(self, text):\n        return re.sub(r\"RT @[-_\\.!~*\\'()a-zA-Z0-9;\\/?:\\@&=\\+$,%#]*?: \", \"\" , text)\n\n    def cleansing_text(self, text):\n        text = self.cleansing_rt(text)\n        text = self.cleansing_hash(text)\n        text = self.cleansing_space(text)\n        text = self.cleansing_url(text)\n        text = self.cleansing_emoji(text)\n        text = self.cleansing_username(text)\n        text = self.cleansing_picture(text)\n        text = self.cleansing_unity(text)\n        #text = self.cleansing_num(text)\n        text = neologdn.normalize(text)\n        text = self.cleansing_repeat(text)\n        return text\n\n    def cleansing_df(self, df, subset_cols=[\"text\"]):\n        if \"text\" in subset_cols:\n            # remove duplicates (because they might be RT.)\n            df = df.drop_duplicates(subset=\"text\", keep=False)\n\n        df_copy = df.copy()\n\n        for col in subset_cols:\n            # cleansing\n            df_copy[col] = df[col].apply(lambda x: self.cleansing_text(x))\n\n        if \"text\" in subset_cols:\n            # remove duplicates\n            df_copy = df_copy.drop_duplicates(subset=\"text\", keep=False).reset_index(drop=True)\n\n        return df_copy\n    \n    def cleansing_df_test(self, df, subset_cols=[\"text\"]):\n        df_copy = df.copy()\n        \n        for col in subset_cols:\n            # cleansing\n            df_copy[col] = df[col].apply(lambda x: self.cleansing_text(x))\n            \n        return df_copy","metadata":{"execution":{"iopub.status.busy":"2022-07-13T09:52:29.532707Z","iopub.execute_input":"2022-07-13T09:52:29.533059Z","iopub.status.idle":"2022-07-13T09:52:29.551830Z","shell.execute_reply.started":"2022-07-13T09:52:29.533023Z","shell.execute_reply":"2022-07-13T09:52:29.550515Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def makeclean_df(df, cols_name = \"text\"):   \n    #クラスのインスタンスの作成\n    tweet_cleaner = CleansingTweets()\n    #処理の対象となるコラムの名前\n    cols = [cols_name]\n    df_clean = tweet_cleaner.cleansing_df(df,subset_cols=cols)\n    return df_clean\n\ndef makeclean_df_test(df, cols_name = \"text\"):   \n    #クラスのインスタンスの作成\n    tweet_cleaner = CleansingTweets()\n    #処理の対象となるコラムの名前\n    cols = [cols_name]\n    df_clean = tweet_cleaner.cleansing_df_test(df,subset_cols=cols)\n    return df_clean","metadata":{"execution":{"iopub.status.busy":"2022-07-13T09:52:29.553417Z","iopub.execute_input":"2022-07-13T09:52:29.553795Z","iopub.status.idle":"2022-07-13T09:52:29.566281Z","shell.execute_reply.started":"2022-07-13T09:52:29.553758Z","shell.execute_reply":"2022-07-13T09:52:29.565027Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#クレイニングの実行\ndf = makeclean_df(tweet)\n#頭から何個か確認してみた\n#print(df.head(20))\n\ntest = makeclean_df_test(test)","metadata":{"execution":{"iopub.status.busy":"2022-07-13T09:52:29.569368Z","iopub.execute_input":"2022-07-13T09:52:29.570535Z","iopub.status.idle":"2022-07-13T09:52:31.261004Z","shell.execute_reply.started":"2022-07-13T09:52:29.570493Z","shell.execute_reply":"2022-07-13T09:52:31.259952Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sentences = df.text.values #文章の抽出\nlabels = df.target.values #ラベルの抽出\ntest_sentence = test.text.values #extract test sentence","metadata":{"execution":{"iopub.status.busy":"2022-07-13T09:52:31.262362Z","iopub.execute_input":"2022-07-13T09:52:31.262985Z","iopub.status.idle":"2022-07-13T09:52:31.268473Z","shell.execute_reply.started":"2022-07-13T09:52:31.262947Z","shell.execute_reply":"2022-07-13T09:52:31.267564Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(len(test_sentence))\nprint(len(test))","metadata":{"execution":{"iopub.status.busy":"2022-07-13T09:52:31.269789Z","iopub.execute_input":"2022-07-13T09:52:31.270393Z","iopub.status.idle":"2022-07-13T09:52:31.289209Z","shell.execute_reply.started":"2022-07-13T09:52:31.270339Z","shell.execute_reply":"2022-07-13T09:52:31.287170Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MODEL_NAME = \"bert-large-uncased\" #コーパス\n# https://huggingface.co/transformers/model_doc/bert.html#<class名>\ntokenizer = BertJapaneseTokenizer.from_pretrained(MODEL_NAME) #bertの学習済みモデルをtokenizerとする\nbert_sc = BertForSequenceClassification.from_pretrained(MODEL_NAME, num_labels=2) #学習済みモデルを用いた分類、num_labelで分類する種類の数\nbert_sc = bert_sc.cuda() #gpuに乗せて高速化させるための関数","metadata":{"execution":{"iopub.status.busy":"2022-07-13T09:52:31.290830Z","iopub.execute_input":"2022-07-13T09:52:31.291941Z","iopub.status.idle":"2022-07-13T09:53:15.902235Z","shell.execute_reply.started":"2022-07-13T09:52:31.291881Z","shell.execute_reply":"2022-07-13T09:53:15.901222Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#データの符号化\ndef Encoding(text_list, labels, max_length=128):\n    \"\"\"\n    text_list : text\n    labels : label\n    max_length : max length of text (set indivisual,64,128,256)\n    \"\"\"\n    dataset_for_loader = []\n    input_ids = []\n    attention_mask = []\n    for idx,text in enumerate(text_list):\n        encoding = tokenizer(text, max_length=max_length,padding=\"max_length\",truncation=True) #textを形態素解析、\"pt\"でtensor出力,辞書型でreturn\n        encoding[\"labels\"] = labels[idx] #add label\n        encoding = {k: torch.tensor(v) for k,v in encoding.items()}  \n        dataset_for_loader.append(encoding)\n        input_ids.append(encoding[\"input_ids\"])\n        attention_mask.append(encoding[\"attention_mask\"])\n    return dataset_for_loader, input_ids, attention_mask\n\nmax_length = 64\ndataset, input_ids, attention_mask = Encoding(sentences,labels,max_length)\n","metadata":{"execution":{"iopub.status.busy":"2022-07-13T09:53:15.903603Z","iopub.execute_input":"2022-07-13T09:53:15.904386Z","iopub.status.idle":"2022-07-13T09:53:20.270228Z","shell.execute_reply.started":"2022-07-13T09:53:15.904345Z","shell.execute_reply":"2022-07-13T09:53:20.269218Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 確認\nprint('Original: ', sentences[0])\nprint('Token IDs:', input_ids[0])","metadata":{"execution":{"iopub.status.busy":"2022-07-13T09:53:20.272094Z","iopub.execute_input":"2022-07-13T09:53:20.272503Z","iopub.status.idle":"2022-07-13T09:53:20.280331Z","shell.execute_reply.started":"2022-07-13T09:53:20.272461Z","shell.execute_reply":"2022-07-13T09:53:20.278993Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# IDを取得\ntrain_size = int(0.8 * len(dataset))\nval_size = len(dataset) - train_size\n# データセットを分割\ntrain_dataset, val_dataset = random_split(dataset, [train_size, val_size])\n\n# データローダーの作成\nbatch_size = 16\n\n# 訓練データローダー\ntrain_dataloader = DataLoader(\n            train_dataset,  \n            sampler = RandomSampler(train_dataset), # ランダムにデータを取得してバッチ化\n            batch_size = batch_size\n        )\n\n# 検証データローダー\nvalidation_dataloader = DataLoader(\n            val_dataset, \n            sampler = SequentialSampler(val_dataset), # 順番にデータを取得してバッチ化\n            batch_size = batch_size\n        )","metadata":{"execution":{"iopub.status.busy":"2022-07-13T09:53:20.282286Z","iopub.execute_input":"2022-07-13T09:53:20.282760Z","iopub.status.idle":"2022-07-13T09:53:20.293845Z","shell.execute_reply.started":"2022-07-13T09:53:20.282703Z","shell.execute_reply":"2022-07-13T09:53:20.292717Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class BertForSequenceClassification_pl(pl.LightningModule):\n    \n    def __init__(self, model_name, num_labels, lr):\n        #model_name: name of transformers model\n        #num_labels: num of labels\n        #lr: learning rate\n        \n        super().__init__()\n        #num_labels,lrを保存\n        self.save_hyperparameters()\n        \n        #BERTのロード\n        self.bert_sc = BertForSequenceClassification.from_pretrained(model_name, num_labels=num_labels)\n    \n    #テストデータのミニバッチが与えられたときテストデータを評価する指標を計算する関数を書く\n    def training_step(self, batch, batch_idx):\n        output = self.bert_sc(**batch)\n        loss = output.loss\n        self.log(\"train_loss\", loss) #損失をtrain_lossの名前でログを取る\n        return loss\n    \n    #検証データ版の評価関数\n    def validation_step(self, batch, batch_idx):\n        output = self.bert_sc(**batch)\n        val_loss = output.loss\n        self.log(\"val_loss\", val_loss) #ログの保存\n        \n    def test_step(self, batch, batch_idx):\n        labels = batch.pop(\"labels\") #バッチからラベルの取得\n        output = self.bert_sc(**batch)\n        labels_predicted = output.logits.argmax(-1)\n        return labels_predicted\n\n    def configure_optimizers(self):\n        return torch.optim.AdamW(self.parameters(), lr = self.hparams.lr)","metadata":{"execution":{"iopub.status.busy":"2022-07-13T09:53:20.295632Z","iopub.execute_input":"2022-07-13T09:53:20.296622Z","iopub.status.idle":"2022-07-13T09:53:20.308883Z","shell.execute_reply.started":"2022-07-13T09:53:20.296578Z","shell.execute_reply":"2022-07-13T09:53:20.307876Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"checkpoint = ModelCheckpoint(monitor=\"val_loss\", mode=\"min\", save_top_k=1,save_weights_only=True, dirpath=\"model/\")\n\n#学習方法の指定\ntrainer = pl.Trainer(gpus=1, max_epochs=10 ,callbacks = [checkpoint])","metadata":{"execution":{"iopub.status.busy":"2022-07-13T09:53:20.310769Z","iopub.execute_input":"2022-07-13T09:53:20.311239Z","iopub.status.idle":"2022-07-13T09:53:20.327554Z","shell.execute_reply.started":"2022-07-13T09:53:20.311199Z","shell.execute_reply":"2022-07-13T09:53:20.326322Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = BertForSequenceClassification_pl(MODEL_NAME, num_labels=2, lr=1e-5)\n\n#fine-Tuning\ntrainer.fit(model, train_dataloader, validation_dataloader)","metadata":{"execution":{"iopub.status.busy":"2022-07-13T09:53:20.329873Z","iopub.execute_input":"2022-07-13T09:53:20.330983Z","iopub.status.idle":"2022-07-13T09:53:50.707590Z","shell.execute_reply.started":"2022-07-13T09:53:20.330895Z","shell.execute_reply":"2022-07-13T09:53:50.703092Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%load_ext tensorboard\n%tensorboard --logdir ./","metadata":{"execution":{"iopub.status.busy":"2022-07-13T09:53:50.710258Z","iopub.execute_input":"2022-07-13T09:53:50.710669Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.bert_sc.save_pretrained(\"./model_ex\")\ntokenizer = BertJapaneseTokenizer.from_pretrained(MODEL_NAME) #bertの学習済みモデルをtokenizerとする\nbert_sc = BertForSequenceClassification.from_pretrained(\"./model_ex\",num_labels=2)\nbert_sc = bert_sc.cuda()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#process test data ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#ert_sc.predict(test_dataset)\ntest_sentence = test_sentence.tolist()\nprint(len(test_sentence))\n#推論,ラベルの取得\nencoding = tokenizer(test_sentence,max_length=max_length, padding=\"max_length\",return_tensors=\"pt\")#形態素解析、tensor出力\nencoding = {k : v.cuda() for k,v in encoding.items()}    \nwith torch.no_grad():\n    output = bert_sc.forward(**encoding)\nscores = output.logits\nlabels_predicted = scores.argmax(-1)\np_labels = labels_predicted.tolist()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test[\"target\"] = p_labels","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"output = test[['id','target']]\noutput.to_csv(\"./submission.csv\", index = False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(len(output))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}