{"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":"pip install pandas pyarrow\nimport pandas as pd","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn.functional as F","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-08-14T08:46:55.839144Z","iopub.execute_input":"2023-08-14T08:46:55.83957Z","iopub.status.idle":"2023-08-14T08:46:55.846612Z","shell.execute_reply.started":"2023-08-14T08:46:55.839537Z","shell.execute_reply":"2023-08-14T08:46:55.845157Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 学習用データ読み込み\ntrain = pd.read_csv('/kaggle/input/asl-fingerspelling/train.csv')\n# 動画データへのパスを重複無しで配列化する\nall_paths = train['path'].unique()\nall_paths[0]","metadata":{"execution":{"iopub.status.busy":"2023-08-14T08:45:09.722608Z","iopub.execute_input":"2023-08-14T08:45:09.723165Z","iopub.status.idle":"2023-08-14T08:45:09.940368Z","shell.execute_reply.started":"2023-08-14T08:45:09.723134Z","shell.execute_reply":"2023-08-14T08:45:09.939217Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 学習に使用する列リストを作成\nall_cols = []\ncol_names = ['x_left_hand_', 'x_right_hand_', 'y_left_hand_', 'y_right_hand_', 'z_left_hand_', 'z_right_hand_']\nfor col_name in col_names:\n    for i in range(21):\n        all_cols.append(col_name + str(i))\nall_cols","metadata":{"execution":{"iopub.status.busy":"2023-08-14T08:45:09.941747Z","iopub.execute_input":"2023-08-14T08:45:09.942393Z","iopub.status.idle":"2023-08-14T08:45:09.954409Z","shell.execute_reply.started":"2023-08-14T08:45:09.942351Z","shell.execute_reply":"2023-08-14T08:45:09.953089Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(all_cols)","metadata":{"execution":{"iopub.status.busy":"2023-08-14T08:50:08.208321Z","iopub.execute_input":"2023-08-14T08:50:08.208747Z","iopub.status.idle":"2023-08-14T08:50:08.216568Z","shell.execute_reply.started":"2023-08-14T08:50:08.208716Z","shell.execute_reply":"2023-08-14T08:50:08.215251Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 動画データフォルダのパス\nfolder_path = '/kaggle/input/asl-fingerspelling/'\n# 全ての動画データを処理する\nfor path in all_paths:\n    # フォルダパスとファイルパスを繋げる\n    parquet_path = folder_path + path\n    # 動画データをDataFrameとして読み込む\n    df = pd.read_parquet(parquet_path)\n    # 動画ごとの切れ目を配列として保持する\n    indexes = df[df['frame']==0].index\n    # デフォルトでIDがインデックスとして当てられているので、インデックスを振り直す\n    df = df.reset_index()\n    # 必要な列のみに絞る\n    df = df[all_cols]\n    # 0-1正規化\n    df = (df - df.min()) / (df.max() - df.min())\n    # NaNを0で置換する\n    df[df.isnull()] = 0\n    # ループの処理内容\n    for i in range(len(indexes)):\n        # スライスで切れ目ごとにデータを抽出してテンソルに変換する\n        tensor = df.iloc[indexes[i]:indexes[i+1]]\n        # テンソルのリサイズ(全てのテンソルを同じサイズにしてニューラルネットワークで処理する)\n            # まずは3次元テンソルを作る\n        tensor = tensor.reshape([3, indexes[i+1]-indexes[i], len(all_cols) / 3])  # 3次元の[フレーム数, 全列数/3]シェイプ\n            \n    \n    \n    \n    \n    ","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_parquet('/kaggle/input/asl-fingerspelling/train_landmarks/1019715464.parquet')\ndf[all_cols]","metadata":{"execution":{"iopub.status.busy":"2023-08-14T08:45:17.370557Z","iopub.execute_input":"2023-08-14T08:45:17.370963Z","iopub.status.idle":"2023-08-14T08:45:32.353793Z","shell.execute_reply.started":"2023-08-14T08:45:17.370931Z","shell.execute_reply":"2023-08-14T08:45:32.352636Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tensor = torch.tensor(df[all_cols].values)\nn_tensor = tensor.reshape([3, len(df), 42])\nn_tensor.shape","metadata":{"execution":{"iopub.status.busy":"2023-08-14T07:27:31.230629Z","iopub.execute_input":"2023-08-14T07:27:31.231061Z","iopub.status.idle":"2023-08-14T07:27:31.41973Z","shell.execute_reply.started":"2023-08-14T07:27:31.231026Z","shell.execute_reply":"2023-08-14T07:27:31.418535Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.DataFrame([[1, 2, 3, 4, 5, 6], [7, 8, 9, 10, 11, 12], [13 ,14, 15, 16, 17, 18]])\ndf","metadata":{"execution":{"iopub.status.busy":"2023-08-14T08:47:02.488146Z","iopub.execute_input":"2023-08-14T08:47:02.488679Z","iopub.status.idle":"2023-08-14T08:47:02.505045Z","shell.execute_reply.started":"2023-08-14T08:47:02.488644Z","shell.execute_reply":"2023-08-14T08:47:02.503668Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tensor = df","metadata":{"execution":{"iopub.status.busy":"2023-08-14T08:47:10.446255Z","iopub.execute_input":"2023-08-14T08:47:10.447488Z","iopub.status.idle":"2023-08-14T08:47:10.458837Z","shell.execute_reply.started":"2023-08-14T08:47:10.447445Z","shell.execute_reply":"2023-08-14T08:47:10.457737Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tensor = torch.tensor(df.values)\ntensor.reshape(3, 3, 2)","metadata":{"execution":{"iopub.status.busy":"2023-08-14T07:43:58.12808Z","iopub.execute_input":"2023-08-14T07:43:58.128554Z","iopub.status.idle":"2023-08-14T07:43:58.136722Z","shell.execute_reply.started":"2023-08-14T07:43:58.128517Z","shell.execute_reply":"2023-08-14T07:43:58.135811Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tensor = torch.tensor([[1, 2, 3, 4, 5, 6], [7, 8, 9, 10, 11, 12], [13 ,14, 15, 16, 17, 18]])\ntensor","metadata":{"execution":{"iopub.status.busy":"2023-08-14T07:41:09.490036Z","iopub.execute_input":"2023-08-14T07:41:09.490889Z","iopub.status.idle":"2023-08-14T07:41:09.498252Z","shell.execute_reply.started":"2023-08-14T07:41:09.490852Z","shell.execute_reply":"2023-08-14T07:41:09.497477Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tensor = torch.tensor(df.iloc[0:3].values)\ntensor.shape\ntensors = []\ntensors.append(tensor)\ntensors","metadata":{"execution":{"iopub.status.busy":"2023-08-14T06:26:50.129492Z","iopub.execute_input":"2023-08-14T06:26:50.129891Z","iopub.status.idle":"2023-08-14T06:26:50.140392Z","shell.execute_reply.started":"2023-08-14T06:26:50.12986Z","shell.execute_reply":"2023-08-14T06:26:50.139202Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tensors.append(tensor)\ntensors","metadata":{"execution":{"iopub.status.busy":"2023-08-14T06:27:24.837959Z","iopub.execute_input":"2023-08-14T06:27:24.839254Z","iopub.status.idle":"2023-08-14T06:27:24.850749Z","shell.execute_reply.started":"2023-08-14T06:27:24.839202Z","shell.execute_reply":"2023-08-14T06:27:24.849652Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = pd.read_csv('/kaggle/input/asl-fingerspelling/train.csv')\ntrain","metadata":{"execution":{"iopub.status.busy":"2023-08-14T06:25:14.05037Z","iopub.execute_input":"2023-08-14T06:25:14.051352Z","iopub.status.idle":"2023-08-14T06:25:14.165725Z","shell.execute_reply.started":"2023-08-14T06:25:14.051316Z","shell.execute_reply":"2023-08-14T06:25:14.16478Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for path in all_paths:\n    ","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"col_names = [\"x_left_hand_\", \"x_right_hand_\", \"y_left_hand_\", \"y_right_hand_\", \"y_left_hand_\", \"y_right_hand_\"]\nall_cols = []\nfor col_name in col_names:\n    for i in range(21):\n        all_cols.append(col_name + str(i))\ndf_2 = df[all_cols]\ndf_2.head()","metadata":{"execution":{"iopub.status.busy":"2023-08-11T06:31:13.813158Z","iopub.execute_input":"2023-08-11T06:31:13.813513Z","iopub.status.idle":"2023-08-11T06:31:13.853918Z","shell.execute_reply.started":"2023-08-11T06:31:13.813485Z","shell.execute_reply":"2023-08-11T06:31:13.852911Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_concat = pd.concat([df['frame'], df_2], axis=1)\ndf_concat.head()","metadata":{"execution":{"iopub.status.busy":"2023-08-11T06:31:15.985361Z","iopub.execute_input":"2023-08-11T06:31:15.986457Z","iopub.status.idle":"2023-08-11T06:31:16.021507Z","shell.execute_reply.started":"2023-08-11T06:31:15.98641Z","shell.execute_reply":"2023-08-11T06:31:16.020261Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_concat[df_concat.isnull()] = 0\ndf_concat","metadata":{"execution":{"iopub.status.busy":"2023-08-11T06:31:17.786228Z","iopub.execute_input":"2023-08-11T06:31:17.786588Z","iopub.status.idle":"2023-08-11T06:31:17.947685Z","shell.execute_reply.started":"2023-08-11T06:31:17.786563Z","shell.execute_reply":"2023-08-11T06:31:17.947037Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"execution":{"iopub.status.busy":"2023-08-11T07:11:32.940205Z","iopub.execute_input":"2023-08-11T07:11:32.940559Z","iopub.status.idle":"2023-08-11T07:11:33.08614Z","shell.execute_reply.started":"2023-08-11T07:11:32.940531Z","shell.execute_reply":"2023-08-11T07:11:33.085182Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_path = train['path'].unique()","metadata":{"execution":{"iopub.status.busy":"2023-08-11T07:14:20.889444Z","iopub.execute_input":"2023-08-11T07:14:20.889814Z","iopub.status.idle":"2023-08-11T07:14:20.900879Z","shell.execute_reply.started":"2023-08-11T07:14:20.889785Z","shell.execute_reply":"2023-08-11T07:14:20.899073Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#ライブラリのインポート\nimport math\nimport numpy as np\nimport pandas as pd\nimport seaborn as sns\nimport matplotlib.pyplot as plt\n\n\nfrom sklearn.preprocessing import StandardScaler\n\nimport torch\nimport torch.nn as nn\nfrom torch.nn import LayerNorm\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.nn import TransformerEncoder, TransformerDecoder, TransformerEncoderLayer, TransformerDecoderLayer\n\n#ランダムシードの設定\nfix_seed = 2023\nnp.random.seed(fix_seed)\ntorch.manual_seed(fix_seed)\n\n#デバイスの設定\ndevice = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu')","metadata":{"execution":{"iopub.status.busy":"2023-08-08T13:25:08.910612Z","iopub.execute_input":"2023-08-08T13:25:08.914781Z","iopub.status.idle":"2023-08-08T13:25:13.244283Z","shell.execute_reply.started":"2023-08-08T13:25:08.914688Z","shell.execute_reply":"2023-08-08T13:25:13.243077Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class AirPassengersDataset(Dataset):\n    def __init__(self, data, seq_len, pred_len):\n        #学習期間と予測期間の設定\n        self.seq_len = seq_len\n        self.pred_len = pred_len\n\n    def __getitem__(self, index):\n        #学習用の系列と予測用の系列を出力\n        s_begin = index\n        s_end = s_begin + self.seq_len\n        r_begin = s_end\n        r_end = r_begin + self.pred_len\n\n        src = self.data[s_begin:s_end]\n        tgt = self.data[r_begin:r_end]\n\n        return src, tgt\n    \n    def __len__(self):\n        return len(self.data) - self.seq_len - self.pred_len + 1","metadata":{"execution":{"iopub.status.busy":"2023-08-08T13:25:54.613142Z","iopub.execute_input":"2023-08-08T13:25:54.613998Z","iopub.status.idle":"2023-08-08T13:25:54.628976Z","shell.execute_reply.started":"2023-08-08T13:25:54.613955Z","shell.execute_reply":"2023-08-08T13:25:54.626057Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def data_provider(flag, seq_len, pred_len, batch_size):\n    #flagに合ったデータを出力\n    data_set = AirPassengersDataset(flag=flag, \n                                    seq_len=seq_len, \n                                    pred_len=pred_len\n                                   )\n    #データをバッチごとに分けて出力できるDataLoaderを使用\n    data_loader = DataLoader(data_set,\n                             batch_size=batch_size, \n                             shuffle=True\n                            )\n    \n    return data_loader","metadata":{"execution":{"iopub.status.busy":"2023-08-08T13:26:09.549852Z","iopub.execute_input":"2023-08-08T13:26:09.550287Z","iopub.status.idle":"2023-08-08T13:26:09.556784Z","shell.execute_reply.started":"2023-08-08T13:26:09.550254Z","shell.execute_reply":"2023-08-08T13:26:09.555608Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#位置エンコーディングの定義\nclass PositionalEncoding(nn.Module):\n\n    def __init__(self, d_model, dropout: float = 0.1, max_len: int = 5000) -> None:\n        super(PositionalEncoding, self).__init__()\n        self.dropout = nn.Dropout(p=dropout)\n        self.d_model = d_model\n\n        pe = torch.zeros(max_len, d_model)\n        position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)\n        div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model))\n        pe[:, 0::2] = torch.sin(position * div_term)\n        pe[:, 1::2] = torch.cos(position * div_term)\n        pe = pe.unsqueeze(0)\n        self.register_buffer('pe', pe)\n\n    def forward(self, x):\n        x = x + self.pe[:, :x.size(1)]\n        return self.dropout(x)\n\n#モデルに入力するために次元を拡張する\nclass TokenEmbedding(nn.Module):\n    def __init__(self, c_in, d_model):\n        super(TokenEmbedding, self).__init__()\n        self.tokenConv = nn.Linear(c_in, d_model) \n\n    def forward(self, x):\n        x = self.tokenConv(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2023-08-08T13:26:22.381909Z","iopub.execute_input":"2023-08-08T13:26:22.382303Z","iopub.status.idle":"2023-08-08T13:26:22.393923Z","shell.execute_reply.started":"2023-08-08T13:26:22.382274Z","shell.execute_reply":"2023-08-08T13:26:22.392353Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Transformer(nn.Module):\n    def __init__(self, num_encoder_layers, num_decoder_layers,\n        d_model, d_input, d_output,\n        dim_feedforward = 512, dropout = 0.1, nhead = 8):\n        \n        super(Transformer, self).__init__()\n        \n\n        #エンべディングの定義\n        self.token_embedding_src = TokenEmbedding(d_input, d_model)\n        self.token_embedding_tgt = TokenEmbedding(d_output, d_model)\n        self.positional_encoding = PositionalEncoding(d_model, dropout=dropout)\n        \n        #エンコーダの定義\n        encoder_layer = TransformerEncoderLayer(d_model=d_model, \n                                                nhead=nhead, \n                                                dim_feedforward=dim_feedforward,\n                                                dropout=dropout,\n                                                batch_first=True,\n                                                activation='gelu'\n                                               )\n        encoder_norm = LayerNorm(d_model)\n        self.transformer_encoder = TransformerEncoder(encoder_layer, \n                                                      num_layers=num_encoder_layers,\n                                                      norm=encoder_norm\n                                                     )\n        \n        #デコーダの定義\n        decoder_layer = TransformerDecoderLayer(d_model=d_model, \n                                                nhead=nhead, \n                                                dim_feedforward=dim_feedforward,\n                                                dropout=dropout,\n                                                batch_first=True,\n                                                activation='gelu'\n                                               )\n        decoder_norm = LayerNorm(d_model)\n        self.transformer_decoder = TransformerDecoder(decoder_layer,  num_layers=num_decoder_layers, \n                                                      norm=decoder_norm)\n        \n        #出力層の定義\n        self.output = nn.Linear(d_model, d_output)\n        \n\n    def forward(self, src, tgt, mask_src, mask_tgt):\n        #mask_src, mask_tgtはセルフアテンションの際に未来のデータにアテンションを向けないためのマスク\n        \n        embedding_src = self.positional_encoding(self.token_embedding_src(src))\n        memory = self.transformer_encoder(embedding_src, mask_src)\n        \n        embedding_tgt = self.positional_encoding(self.token_embedding_tgt(tgt))\n        outs = self.transformer_decoder(embedding_tgt, memory, mask_tgt)\n        \n        output = self.output(outs)\n        return output\n\n    def encode(self, src, mask_src):\n        return self.transformer_encoder(self.positional_encoding(self.token_embedding_src(src)), mask_src)\n\n    def decode(self, tgt, memory, mask_tgt):\n        return self.transformer_decoder(self.positional_encoding(self.token_embedding_tgt(tgt)), memory, mask_tgt)","metadata":{"execution":{"iopub.status.busy":"2023-08-08T13:27:49.418668Z","iopub.execute_input":"2023-08-08T13:27:49.419116Z","iopub.status.idle":"2023-08-08T13:27:49.431553Z","shell.execute_reply.started":"2023-08-08T13:27:49.419082Z","shell.execute_reply":"2023-08-08T13:27:49.430527Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_mask(src, tgt):\n    \n    seq_len_src = src.shape[1]\n    seq_len_tgt = tgt.shape[1]\n\n    mask_tgt = generate_square_subsequent_mask(seq_len_tgt).to(device)\n    mask_src = generate_square_subsequent_mask(seq_len_src).to(device)\n\n    return mask_src, mask_tgt\n\n\ndef generate_square_subsequent_mask(seq_len):\n    mask = torch.triu(torch.full((seq_len, seq_len), float('-inf')), diagonal=1)\n    return mask","metadata":{"execution":{"iopub.status.busy":"2023-08-08T13:28:42.728459Z","iopub.execute_input":"2023-08-08T13:28:42.72893Z","iopub.status.idle":"2023-08-08T13:28:42.736795Z","shell.execute_reply.started":"2023-08-08T13:28:42.728894Z","shell.execute_reply":"2023-08-08T13:28:42.73546Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train(model, data_provider, optimizer, criterion):\n    model.train()\n    total_loss = []\n    for src, tgt in data_provider:\n        \n        src = src.float().to(device)\n        tgt = tgt.float().to(device)\n\n        input_tgt = torch.cat((src[:,-1:,:],tgt[:,:-1,:]), dim=1)\n\n        mask_src, mask_tgt = create_mask(src, input_tgt)\n\n        output = model(\n            src=src, tgt=input_tgt, \n            mask_src=mask_src, mask_tgt=mask_tgt\n        )\n\n        optimizer.zero_grad()\n\n        loss = criterion(output, tgt)\n        loss.backward()\n        total_loss.append(loss.cpu().detach())\n        optimizer.step()\n        \n    return np.average(total_loss)\n\ndef evaluate(flag, model, data_provider, criterion):\n    model.eval()\n    total_loss = []\n    for src, tgt in data_provider:\n        \n        src = src.float().to(device)\n        tgt = tgt.float().to(device)\n\n        seq_len_src = src.shape[1]\n        mask_src = (torch.zeros(seq_len_src, seq_len_src)).type(torch.bool)\n        mask_src = mask_src.float().to(device)\n    \n        memory = model.encode(src, mask_src)\n        outputs = src[:, -1:, :]\n        seq_len_tgt = tgt.shape[1]\n    \n        for i in range(seq_len_tgt - 1):\n        \n            mask_tgt = (generate_square_subsequent_mask(outputs.size(1))).to(device)\n        \n            output = model.decode(outputs, memory, mask_tgt)\n            output = model.output(output)\n\n            outputs = torch.cat([outputs, output[:, -1:, :]], dim=1)\n        \n        loss = criterion(outputs, tgt)\n        total_loss.append(loss.cpu().detach())\n        \n    if flag=='test':\n        true = torch.cat((src, tgt), dim=1)\n        pred = torch.cat((src, output), dim=1)\n        plt.plot(true.squeeze().cpu().detach().numpy(), label='true')\n        plt.plot(pred.squeeze().cpu().detach().numpy(), label='pred')\n        plt.legend()\n        plt.savefig('test.pdf')\n        \n    return np.average(total_loss)","metadata":{"execution":{"iopub.status.busy":"2023-08-08T13:29:07.31275Z","iopub.execute_input":"2023-08-08T13:29:07.313197Z","iopub.status.idle":"2023-08-08T13:29:07.330492Z","shell.execute_reply.started":"2023-08-08T13:29:07.31316Z","shell.execute_reply":"2023-08-08T13:29:07.32933Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"d_input = 3\nd_output = 1\nd_model = 512\nnhead = 8\ndim_feedforward = 2048\nnum_encoder_layers = 1\nnum_decoder_layers = 1\ndropout = 0.01\nsrc_len = 36\ntgt_len = 12\nbatch_size = 1\nepochs = 30\nbest_loss = float('Inf')\nbest_model = None\n\nmodel = Transformer(num_encoder_layers=num_encoder_layers,\n                    num_decoder_layers=num_decoder_layers,\n                    d_model=d_model,\n                    d_input=d_input, \n                    d_output=d_output,\n                    dim_feedforward=dim_feedforward,\n                    dropout=dropout, nhead=nhead\n                   )\n\nfor p in model.parameters():\n    if p.dim() > 1:\n        nn.init.xavier_uniform_(p)\n\nmodel = model.to(device)\n\ncriterion = torch.nn.MSELoss()\n\noptimizer = torch.optim.RAdam(model.parameters(), lr=0.0001)","metadata":{"execution":{"iopub.status.busy":"2023-08-08T13:29:40.849012Z","iopub.execute_input":"2023-08-08T13:29:40.849409Z","iopub.status.idle":"2023-08-08T13:29:41.037044Z","shell.execute_reply.started":"2023-08-08T13:29:40.849378Z","shell.execute_reply":"2023-08-08T13:29:41.03594Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"valid_losses = []\nfor epoch in range(1, epochs + 1):\n    \n    loss_train = train(\n        model=model, data_provider=data_provider('train', src_len, tgt_len, batch_size), optimizer=optimizer,\n        criterion=criterion\n    )\n        \n    loss_valid = evaluate(\n        flag='val', model=model, data_provider=data_provider('val', src_len, tgt_len, batch_size), criterion=criterion\n    )\n    \n    if epoch%10==0:\n        print('[{}/{}] train loss: {:.2f}, valid loss: {:.2f}'.format(\n            epoch, epochs,\n            loss_train, loss_valid,\n        ))\n        \n    valid_losses.append(loss_valid)\n    \n    if best_loss > loss_valid:\n        best_loss = loss_valid\n        best_model = model","metadata":{"execution":{"iopub.status.busy":"2023-08-08T13:29:43.415181Z","iopub.execute_input":"2023-08-08T13:29:43.415572Z","iopub.status.idle":"2023-08-08T13:32:08.266221Z","shell.execute_reply.started":"2023-08-08T13:29:43.415534Z","shell.execute_reply":"2023-08-08T13:32:08.264399Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"evaluate(flag='test', model=best_model, data_provider=data_provider('test', src_len, tgt_len, batch_size), criterion=criterion)","metadata":{"execution":{"iopub.status.busy":"2023-08-08T13:32:08.268523Z","iopub.execute_input":"2023-08-08T13:32:08.269183Z","iopub.status.idle":"2023-08-08T13:32:09.109366Z","shell.execute_reply.started":"2023-08-08T13:32:08.269139Z","shell.execute_reply":"2023-08-08T13:32:09.108168Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}