{"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":"# Investment-wise Trainsformer\n\n**由于Kaggle Mean Pearson 得分计算存在bug** ，因此准确的test集得分计算请在本地进行，LB榜单不可信。计算方法可参考本文。","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader, Dataset\n\nimport matplotlib.pyplot as plt\n%matplotlib inline\n\nimport os\n\nimport gc\nfrom pathlib import Path\nfrom tqdm import tqdm","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:37:30.715333Z","iopub.execute_input":"2022-08-05T11:37:30.716611Z","iopub.status.idle":"2022-08-05T11:37:32.087642Z","shell.execute_reply.started":"2022-08-05T11:37:30.716474Z","shell.execute_reply":"2022-08-05T11:37:32.086571Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# EDA","metadata":{}},{"cell_type":"code","source":"train = pd.read_csv(\"../input/boolart-market-prediction/train.csv\")","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:37:32.089694Z","iopub.execute_input":"2022-08-05T11:37:32.090516Z","iopub.status.idle":"2022-08-05T11:38:06.034344Z","shell.execute_reply.started":"2022-08-05T11:37:32.090476Z","shell.execute_reply":"2022-08-05T11:38:06.033308Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train.head()","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:38:06.035724Z","iopub.execute_input":"2022-08-05T11:38:06.036385Z","iopub.status.idle":"2022-08-05T11:38:06.073701Z","shell.execute_reply.started":"2022-08-05T11:38:06.036348Z","shell.execute_reply":"2022-08-05T11:38:06.072643Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"特征为每个时刻投资标的的300个特征，需要预测target为投资标的收益率排序（经过数据放缩后处理），target越大，说明该投资标的在当前时刻收益率越靠前。\n\n我们有4种建模范式：\n\n1. 考虑单个投资标的当前时刻特征，输入信息为300个特征; 这种建模方式简单，但是由于我们预测的目标本质上是排序，如果只考虑单个投资标的，其实是很难预测其收益率排序的。\n\n2. 模型考虑当前时刻所有投资标的，同时预测target，Transformer在这里很适用，self-attention会计算投资间的两两关系，由于不存在时序问题，不需要做position embedding.\n\n3. 模型考虑投资标的历史特征序列，这种建模方法对该投资表的的趋势-波动建模，可以采用RNN/Seq2Seq/Transformer。\n\n4. 同时考虑时序和投资标的，可行的建模方式是，每个投资标的输入用一个RNN进行编码，编码后采用方案2. 进行建模。\n\n\n我们这里采用方案2进行建模，首先我们简单分析一下数据。","metadata":{}},{"cell_type":"code","source":"train.shape","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:38:06.076483Z","iopub.execute_input":"2022-08-05T11:38:06.076861Z","iopub.status.idle":"2022-08-05T11:38:06.084036Z","shell.execute_reply.started":"2022-08-05T11:38:06.076823Z","shell.execute_reply":"2022-08-05T11:38:06.082367Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 投资标的数量\ntrain.investment_id.nunique()","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:38:06.086098Z","iopub.execute_input":"2022-08-05T11:38:06.086554Z","iopub.status.idle":"2022-08-05T11:38:06.102443Z","shell.execute_reply.started":"2022-08-05T11:38:06.086517Z","shell.execute_reply":"2022-08-05T11:38:06.101215Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 投时间跨度\n\ntrain.groupby(\"investment_id\")['time_id'].agg(['min', 'max', 'count']).describe()","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:38:06.104697Z","iopub.execute_input":"2022-08-05T11:38:06.105291Z","iopub.status.idle":"2022-08-05T11:38:06.150062Z","shell.execute_reply.started":"2022-08-05T11:38:06.105263Z","shell.execute_reply":"2022-08-05T11:38:06.148877Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 时间戳从0-1000, 但是均存在一些缺失，可能是非交易日。","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:38:06.151808Z","iopub.execute_input":"2022-08-05T11:38:06.152300Z","iopub.status.idle":"2022-08-05T11:38:06.157322Z","shell.execute_reply.started":"2022-08-05T11:38:06.152229Z","shell.execute_reply":"2022-08-05T11:38:06.156019Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"_ = plt.hist(train.target, 100)","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:38:06.158890Z","iopub.execute_input":"2022-08-05T11:38:06.160050Z","iopub.status.idle":"2022-08-05T11:38:06.515152Z","shell.execute_reply.started":"2022-08-05T11:38:06.160014Z","shell.execute_reply":"2022-08-05T11:38:06.514082Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# target 分布很正态","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:38:06.516739Z","iopub.execute_input":"2022-08-05T11:38:06.517050Z","iopub.status.idle":"2022-08-05T11:38:06.522053Z","shell.execute_reply.started":"2022-08-05T11:38:06.517021Z","shell.execute_reply":"2022-08-05T11:38:06.520181Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"feat_cols = [f'f_{i}' for i in range(300)]\n\ntrain[feat_cols].mean().describe()","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:38:06.527528Z","iopub.execute_input":"2022-08-05T11:38:06.528243Z","iopub.status.idle":"2022-08-05T11:38:07.184117Z","shell.execute_reply.started":"2022-08-05T11:38:06.528203Z","shell.execute_reply":"2022-08-05T11:38:07.183122Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 特征的均值分布在-0.5 ～ 0.5之间，应该已经做过一些处理，在来查看一下min/max\n\n_ = plt.hist(train[feat_cols].min().describe(), 20)","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:38:07.185598Z","iopub.execute_input":"2022-08-05T11:38:07.186758Z","iopub.status.idle":"2022-08-05T11:38:08.050241Z","shell.execute_reply.started":"2022-08-05T11:38:07.186718Z","shell.execute_reply":"2022-08-05T11:38:08.049108Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"_ = plt.hist(train[feat_cols].max(), 20)","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:38:08.052046Z","iopub.execute_input":"2022-08-05T11:38:08.052454Z","iopub.status.idle":"2022-08-05T11:38:08.982993Z","shell.execute_reply.started":"2022-08-05T11:38:08.052414Z","shell.execute_reply":"2022-08-05T11:38:08.981610Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# min/max 的分布也比较集中","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:38:08.984711Z","iopub.execute_input":"2022-08-05T11:38:08.987305Z","iopub.status.idle":"2022-08-05T11:38:08.996067Z","shell.execute_reply.started":"2022-08-05T11:38:08.987256Z","shell.execute_reply":"2022-08-05T11:38:08.994208Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 随便查看16个特征的分布\n\n\nf, ax = plt.subplots(4,4,figsize=(16,16))\n\nfor i in range(4):\n    for j in range(4):\n        col = f'f_{np.random.randint(0, 300)}'\n        ax[i][j].hist(train[col], 100)\n        ax[i][j].set_title(col)","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:38:08.997685Z","iopub.execute_input":"2022-08-05T11:38:09.001219Z","iopub.status.idle":"2022-08-05T11:38:13.313330Z","shell.execute_reply.started":"2022-08-05T11:38:09.001150Z","shell.execute_reply":"2022-08-05T11:38:13.312031Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 特征分布均在0附近\n# 加载test看看\n\ntest = pd.read_csv(\"../input/boolart-market-prediction/test.csv\")","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:38:13.314469Z","iopub.execute_input":"2022-08-05T11:38:13.314848Z","iopub.status.idle":"2022-08-05T11:38:20.520373Z","shell.execute_reply.started":"2022-08-05T11:38:13.314803Z","shell.execute_reply":"2022-08-05T11:38:20.519138Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test.head()","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:38:20.522152Z","iopub.execute_input":"2022-08-05T11:38:20.522547Z","iopub.status.idle":"2022-08-05T11:38:20.570381Z","shell.execute_reply.started":"2022-08-05T11:38:20.522509Z","shell.execute_reply":"2022-08-05T11:38:20.564705Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 除了没有target, 其他与train完全一直\n# 查看一下 test 的时间跨度\n\ntest.groupby('investment_id')['time_id'].agg(['min', 'max', 'count'])","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:38:20.572953Z","iopub.execute_input":"2022-08-05T11:38:20.573296Z","iopub.status.idle":"2022-08-05T11:38:20.592073Z","shell.execute_reply.started":"2022-08-05T11:38:20.573269Z","shell.execute_reply":"2022-08-05T11:38:20.591191Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 时间戳为1001 -> 1219\n# 部分investment_id 在某些时间数据缺失, 例如`3722`只有205条数据","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:38:20.593628Z","iopub.execute_input":"2022-08-05T11:38:20.594004Z","iopub.status.idle":"2022-08-05T11:38:20.598414Z","shell.execute_reply.started":"2022-08-05T11:38:20.593967Z","shell.execute_reply":"2022-08-05T11:38:20.597226Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 数据准备\n\n首先我们填充investment_id 缺失的时间。","metadata":{}},{"cell_type":"code","source":"import itertools","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:38:20.600467Z","iopub.execute_input":"2022-08-05T11:38:20.600878Z","iopub.status.idle":"2022-08-05T11:38:20.607325Z","shell.execute_reply.started":"2022-08-05T11:38:20.600841Z","shell.execute_reply":"2022-08-05T11:38:20.606057Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_missing = set(itertools.product(train.investment_id.unique(), train.time_id.unique())) - \\\nset([tuple(l) for l in train[['investment_id', 'time_id']].values.tolist()])","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:38:20.608999Z","iopub.execute_input":"2022-08-05T11:38:20.610134Z","iopub.status.idle":"2022-08-05T11:38:22.349587Z","shell.execute_reply.started":"2022-08-05T11:38:20.610098Z","shell.execute_reply":"2022-08-05T11:38:22.348486Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(train_missing)","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:38:22.351070Z","iopub.execute_input":"2022-08-05T11:38:22.351759Z","iopub.status.idle":"2022-08-05T11:38:22.359473Z","shell.execute_reply.started":"2022-08-05T11:38:22.351718Z","shell.execute_reply":"2022-08-05T11:38:22.358319Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"miss_train = pd.DataFrame([list(l) for l in train_missing], columns=['investment_id', 'time_id'])\nmiss_train['row_id'] = miss_train['time_id'].astype(str) + \"_\" + miss_train['investment_id'].astype(str)","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:38:22.360972Z","iopub.execute_input":"2022-08-05T11:38:22.361981Z","iopub.status.idle":"2022-08-05T11:38:22.419743Z","shell.execute_reply.started":"2022-08-05T11:38:22.361928Z","shell.execute_reply":"2022-08-05T11:38:22.418853Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = pd.concat([train, miss_train], axis=0, ignore_index=True).sort_values(['time_id', 'investment_id'])","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:38:22.421051Z","iopub.execute_input":"2022-08-05T11:38:22.421945Z","iopub.status.idle":"2022-08-05T11:38:23.502578Z","shell.execute_reply.started":"2022-08-05T11:38:22.421907Z","shell.execute_reply":"2022-08-05T11:38:23.501436Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train.head()","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:38:23.504301Z","iopub.execute_input":"2022-08-05T11:38:23.504723Z","iopub.status.idle":"2022-08-05T11:38:23.532270Z","shell.execute_reply.started":"2022-08-05T11:38:23.504680Z","shell.execute_reply":"2022-08-05T11:38:23.531171Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_row_id = test.row_id.values\n\ntest_missing = set(itertools.product(test.investment_id.unique(), test.time_id.unique())) - \\\nset([tuple(l) for l in test[['investment_id', 'time_id']].values.tolist()])\n\nmiss_test = pd.DataFrame([list(l) for l in test_missing], columns=['investment_id', 'time_id'])\nmiss_test['row_id'] = miss_test['time_id'].astype(str) + \"_\" + miss_test['investment_id'].astype(str)","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:38:23.534518Z","iopub.execute_input":"2022-08-05T11:38:23.534985Z","iopub.status.idle":"2022-08-05T11:38:23.865218Z","shell.execute_reply.started":"2022-08-05T11:38:23.534944Z","shell.execute_reply":"2022-08-05T11:38:23.864215Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test = pd.concat([test, miss_test], axis=0, ignore_index=True).sort_values(['time_id', 'investment_id'])","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:38:23.866575Z","iopub.execute_input":"2022-08-05T11:38:23.867187Z","iopub.status.idle":"2022-08-05T11:38:24.112619Z","shell.execute_reply.started":"2022-08-05T11:38:23.867115Z","shell.execute_reply":"2022-08-05T11:38:24.111479Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test.head()","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:38:24.114340Z","iopub.execute_input":"2022-08-05T11:38:24.114749Z","iopub.status.idle":"2022-08-05T11:38:24.141824Z","shell.execute_reply.started":"2022-08-05T11:38:24.114707Z","shell.execute_reply":"2022-08-05T11:38:24.140660Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 构造mask列，0标识填充样本，1标识真实样本\n\ntrain['mask'] = 1-train['f_0'].isnull().astype(int)\ntest['mask'] = 1-test['f_0'].isnull().astype(int)","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:38:24.147835Z","iopub.execute_input":"2022-08-05T11:38:24.148125Z","iopub.status.idle":"2022-08-05T11:38:24.159234Z","shell.execute_reply.started":"2022-08-05T11:38:24.148098Z","shell.execute_reply":"2022-08-05T11:38:24.157962Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train['target'] = train['target'].fillna(-1)\ntrain.fillna(0., inplace=True)\ntest.fillna(0., inplace=True)","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:38:24.161349Z","iopub.execute_input":"2022-08-05T11:38:24.162238Z","iopub.status.idle":"2022-08-05T11:38:24.492661Z","shell.execute_reply.started":"2022-08-05T11:38:24.162178Z","shell.execute_reply":"2022-08-05T11:38:24.491576Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train.head()","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:38:24.494456Z","iopub.execute_input":"2022-08-05T11:38:24.494885Z","iopub.status.idle":"2022-08-05T11:38:24.521972Z","shell.execute_reply.started":"2022-08-05T11:38:24.494845Z","shell.execute_reply":"2022-08-05T11:38:24.520697Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 检查每个时间段investment_id数量一致\n\ntrain.groupby('time_id')['investment_id'].count().agg(['min', 'max'])","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:38:24.523665Z","iopub.execute_input":"2022-08-05T11:38:24.524057Z","iopub.status.idle":"2022-08-05T11:38:24.543113Z","shell.execute_reply.started":"2022-08-05T11:38:24.524017Z","shell.execute_reply":"2022-08-05T11:38:24.542198Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test.groupby('time_id')['investment_id'].count().agg(['min', 'max'])","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:38:24.544501Z","iopub.execute_input":"2022-08-05T11:38:24.545085Z","iopub.status.idle":"2022-08-05T11:38:24.558725Z","shell.execute_reply.started":"2022-08-05T11:38:24.545046Z","shell.execute_reply":"2022-08-05T11:38:24.557799Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 构建Dataset","metadata":{}},{"cell_type":"code","source":"# 建模前最关键的是定义input 数据结构\nFEATURE = [f'f_{i}' for i in range(300)]\n\nclass CustomDataSet(Dataset):\n    \n    def __init__(self, df):\n        self.gp = df.groupby('time_id')\n        self.time_id = df.time_id.unique()\n    \n    def __len__(self):\n        return len(self.time_id)\n    \n    def __getitem__(self, item):\n        df = self.gp.get_group(self.time_id[item])\n        feat = df[FEATURE].values\n        mask = df['mask'].values\n        if 'target' in df.columns:\n            target = df['target'].values\n            return {\"feat\": feat, \"mask\": mask}, target\n        return {\"feat\": feat, \"mask\": mask}","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:38:26.994813Z","iopub.execute_input":"2022-08-05T11:38:26.995431Z","iopub.status.idle":"2022-08-05T11:38:27.003859Z","shell.execute_reply.started":"2022-08-05T11:38:26.995383Z","shell.execute_reply":"2022-08-05T11:38:27.002567Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 检查一下\ntrain_ds = CustomDataSet(train)\nprint(train_ds[0])\n\ndel train_ds\nimport gc\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:38:40.292823Z","iopub.execute_input":"2022-08-05T11:38:40.293951Z","iopub.status.idle":"2022-08-05T11:38:40.549371Z","shell.execute_reply.started":"2022-08-05T11:38:40.293897Z","shell.execute_reply":"2022-08-05T11:38:40.548380Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 模型\n\n我们这里同一个时刻下的Investment_id是没有顺序含义的，因此不需要构造Position Embedding，但是需要注意，我们填充了缺失数据，需要在计算Self-Attention的时候，使用Mask信息，避免填充数据对我们的数据绕动，同时在计算Loss时候，不处理填充数据中的Label。","metadata":{}},{"cell_type":"code","source":"class CustomModel(nn.Module):\n    \n    def __init__(self, input_size, num_layers, d_model, ffn_size, n_head):\n        super().__init__()\n        self.input_fc = nn.Linear(input_size, d_model)\n        self.transformers = nn.Sequential(*[\n            nn.TransformerEncoderLayer(d_model, n_head, ffn_size, batch_first=True) for i in range(num_layers)])\n        self.head = nn.Sequential(\n            nn.Linear(d_model, d_model//2),\n            nn.GELU(),\n            nn.Linear(d_model//2, 1)\n        )\n        self.loss_fn = nn.MSELoss(reduce=False)\n        \n    def forward(self, feat, mask, label=None):\n        feat = self.input_fc(feat.float())\n        for layer in self.transformers:\n            feat = layer(feat, src_key_padding_mask=mask.long()) * 0.9 + feat * 0.1  # residual connect\n        pred = self.head(feat).squeeze()\n        \n        if label is not None:\n            loss = self.loss_fn(label.float(), pred)\n            loss = (loss * mask) / torch.sum(mask)\n            return pred, loss\n        return pred","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:38:43.492957Z","iopub.execute_input":"2022-08-05T11:38:43.493424Z","iopub.status.idle":"2022-08-05T11:38:43.503748Z","shell.execute_reply.started":"2022-08-05T11:38:43.493377Z","shell.execute_reply":"2022-08-05T11:38:43.502516Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 训练","metadata":{}},{"cell_type":"code","source":"import gc\nfrom pathlib import Path\n\nimport numpy as np\nimport torch\nfrom torch.utils.data import DataLoader\nfrom tqdm import tqdm\nimport random\n\ndef seed_everything(seed):\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n\n\nclass AverageMeter(object):\n    val = 0\n    avg = 0\n    sum = 0\n    count = 0\n    best = None\n\n    \"\"\"Computes and stores the average and current value\"\"\"\n    def __init__(self, larger_is_better=False):\n        self.larger_is_better = larger_is_better\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n        if self.larger_is_better:\n            self.best = -np.inf\n        else:\n            self.best = np.inf\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count\n        if self.larger_is_better and (val > self.best):\n            self.best = val\n            return True\n        elif (not self.larger_is_better) and (val < self.best):\n            self.best = val\n            return True\n        else:\n            return False\n\n\ndef to_device(d, device):\n    if isinstance(d, dict):\n        return {k: to_device(v, device) for k, v in d.items()}\n    elif isinstance(d, torch.Tensor):\n        return d.to(device)\n    elif isinstance(d, list):\n        return [to_device(i, device) for i in d]\n    else:\n        raise ValueError\n\n\nclass TrainerConfig:\n    max_epochs = 8\n    batch_size = 32\n    grad_norm_clip = None\n    # Data Loader\n    num_workers = 8\n    collate_fn = None\n    # Device\n    device = \"cuda:0\"\n    gpu_ids = None\n    save_last = False\n    # random\n    seed = 42\n\n    def __init__(self, save_path, **kwargs):\n        self.save_path = save_path\n        for k, v in kwargs.items():\n            setattr(self, k, v)\n\n\nclass Trainer:\n\n    def __init__(self, train_config, model, optimizer, train_dataset,\n                 test_dataset=None, metric_fn=None, metric_larger_better=True, lr_scheduler=None, collate_fn=None):\n        seed_everything(train_config.seed)\n        self.model = model\n        self.optimizer = optimizer\n        self.lr_scheduler = lr_scheduler\n        self.train_dataset = train_dataset\n        self.test_dataset = test_dataset\n        self.config = train_config\n        self.metric_fn = metric_fn\n        self.metric_larger_better = metric_larger_better\n        self.collate_fn = collate_fn\n        # take over whatever gpus are on the system\n        self.device = self.config.device\n        if self.config.gpu_ids is not None:\n            self.model = torch.nn.DataParallel(self.model, device_ids=self.config.gpu_ids).to(self.device)\n        else:\n            self.model.to(self.device)\n    def save_checkpoint(self, oof, last=False):\n        # DataParallel wrappers keep raw model object in .module attribute\n        Path(self.config.save_path).parent.mkdir(parents=True, exist_ok=True)\n        raw_model = self.model.module if hasattr(self.model, \"module\") else self.model\n        save_path = self.config.save_path\n        if last:\n            save_path = self.config.save_path + \".last\"\n        print(f\"saved in {save_path}\")\n        torch.save({'checkpoint': raw_model.state_dict(), 'oof': oof}, save_path)\n\n    @torch.no_grad()\n    def valid_epoch(self, loader):\n        model, config = self.model, self.config\n        model.eval()\n        losses = AverageMeter()\n        pbar = enumerate(loader)\n        score = None\n        preds, labels = [], []\n        for it, (x, y) in pbar:\n            x = to_device(x, self.device)\n            y = to_device(y, self.device)\n            # forward the model\n            logit, loss = model(**x, label=y)\n            loss = loss.mean()  # collapse all losses if they are scattered on multiple gpus\n            losses.update(loss.item())\n            preds.append(logit.cpu().numpy())\n            labels.append(y.cpu().numpy())\n            del x; del y; gc.collect()\n        preds = np.concatenate(preds, axis=0)\n        labels = np.concatenate(labels, axis=0)\n        if self.metric_fn is not None:\n            score = self.metric_fn(labels, preds)\n        return losses.avg, preds, score\n\n    def train_epoch(self, loader, num_epoch):\n        model, optimizer, lr_scheduler, config = self.model, self.optimizer, self.lr_scheduler, self.config\n        model.train()\n        score = None\n        losses = AverageMeter()\n        pbar = tqdm(enumerate(loader), total=len(loader))\n        preds = []\n        labels = []\n\n        for it, (x, y) in pbar:\n            # place data on the correct device\n            x = to_device(x, self.device)\n            y = to_device(y, self.device)\n\n            # forward the model\n            logit, loss = model(**x, label=y)\n            loss = loss.mean()  # collapse all losses if they are scattered on multiple gpus\n\n            # backprop and update the parameters\n            model.zero_grad()\n            loss.backward()\n            if self.config.grad_norm_clip is not None:\n                torch.nn.utils.clip_grad_norm_(model.parameters(), config.grad_norm_clip)\n            optimizer.step()\n\n            # record\n            loss = loss.item()\n            losses.update(loss)\n\n            preds.append(logit.detach().cpu().numpy())\n            labels.append(y.cpu().numpy())\n\n            # decay the learning rate based on our progress\n            if lr_scheduler is not None:\n                lr_scheduler.step(num_epoch+it/len(loader))\n            lr = []\n            for param_group in optimizer.param_groups:\n                lr.append(param_group['lr'])\n            lr = np.mean(lr)\n\n            # report progress\n            pbar.set_description(f\"train loss {loss:.4f} lr {lr:.6f}\")\n\n        labels = np.concatenate(labels, axis=0)\n        preds = np.concatenate(preds, axis=0)\n        if self.metric_fn is not None:\n            score = self.metric_fn(labels, preds)\n        return losses.avg, score\n\n    def fit(self):\n        config = self.config\n        scores = AverageMeter(self.metric_larger_better)\n\n        train_loader = DataLoader(\n            self.train_dataset,\n            shuffle=True,\n            pin_memory=True,\n            batch_size=config.batch_size,\n            num_workers=config.num_workers,\n            collate_fn=self.collate_fn\n        )\n\n        if self.test_dataset is not None:\n            test_loader = DataLoader(\n                self.test_dataset,\n                shuffle=True,\n                pin_memory=True,\n                batch_size=config.batch_size,\n                num_workers=config.num_workers,\n                collate_fn=self.collate_fn\n            )\n\n        for epoch in range(config.max_epochs):\n            trn_loss, trn_score = self.train_epoch(train_loader, epoch)\n            info = f\"Epoch {epoch + 1} train loss {trn_loss:.4f}\"\n            if self.metric_fn is not None:\n                info += f\" score {trn_score:.4f}\"\n            if self.test_dataset is not None:\n                val_loss, val_preds, val_score = self.valid_epoch(test_loader)\n                info += f\" val loss {val_loss:.4f}\"\n                if self.metric_fn is not None:\n                    info += f\" score {val_score: .4f}\"\n                good_model = scores.update(val_score)\n                if good_model:\n                    self.save_checkpoint(val_preds, False)\n            print(info)\n            # supports early stopping based on the test loss, or just save always if no test set is provided\n        if self.config.save_last:\n            self.save_checkpoint(None if self.test_dataset is None else val_preds, True)\n\n    @torch.no_grad()\n    def predict(self, dataset):\n        self.model.eval()\n        loader = DataLoader(\n            dataset,\n            shuffle=True,\n            pin_memory=True,\n            batch_size=self.config.batch_size,\n            num_workers=self.config.num_workers,\n            collate_fn=self.collate_fn\n        )\n        preds = []\n        for x in loader:\n            x = to_device(x, self.device)\n            logit = self.model(**x)\n            preds.append(logit.cpu().numpy())\n        preds = np.concatenate(preds, axis=0)\n        return preds\n","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:38:45.594080Z","iopub.execute_input":"2022-08-05T11:38:45.594659Z","iopub.status.idle":"2022-08-05T11:38:45.632367Z","shell.execute_reply.started":"2022-08-05T11:38:45.594620Z","shell.execute_reply":"2022-08-05T11:38:45.631274Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 数据划分","metadata":{}},{"cell_type":"code","source":"# 选取最后100个time_id作为验证集\ntime_ids = train.time_id.unique()\ntrain_time_ids = time_ids[:-100]\nval_time_ids = time_ids[-100:]\n\ntrain_ds = CustomDataSet(train[train.time_id.isin(train_time_ids)])\nval_ds = CustomDataSet(train[train.time_id.isin(val_time_ids)])\ntrain_ds_all = CustomDataSet(train)\ntest_ds = CustomDataSet(test)","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:39:01.806137Z","iopub.execute_input":"2022-08-05T11:39:01.806825Z","iopub.status.idle":"2022-08-05T11:39:02.462263Z","shell.execute_reply.started":"2022-08-05T11:39:01.806785Z","shell.execute_reply":"2022-08-05T11:39:02.461216Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 评价函数","metadata":{}},{"cell_type":"code","source":"from scipy.stats import pearsonr\n\ndef pearson_score(labels, preds):\n    # 遍历每个sample（某个time_id下的所有投资品收益），计算当前 time_id pearson score, 再计算平均值。\n    return np.mean([pearsonr(l, p)[0] for l, p in zip(labels, preds)])","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:39:21.751784Z","iopub.execute_input":"2022-08-05T11:39:21.753049Z","iopub.status.idle":"2022-08-05T11:39:22.328615Z","shell.execute_reply.started":"2022-08-05T11:39:21.752997Z","shell.execute_reply":"2022-08-05T11:39:22.327567Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 模型验证","metadata":{}},{"cell_type":"code","source":"model = CustomModel(input_size=300, num_layers=3, d_model=512, ffn_size=1024, n_head=8)\noptimizer = torch.optim.Adam(model.parameters(), lr=0.0001)\n\nconfig = TrainerConfig(save_path=\"./test_model.bin\", batch_size=32, max_epochs=2)\ntrainer = Trainer(config, model, optimizer, train_dataset=train_ds, test_dataset=val_ds, \n                  metric_fn=pearson_score, metric_larger_better=True)\ntrainer.fit()\n\ndel model; del trainer","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:39:37.660002Z","iopub.execute_input":"2022-08-05T11:39:37.660412Z","iopub.status.idle":"2022-08-05T11:39:59.056757Z","shell.execute_reply.started":"2022-08-05T11:39:37.660376Z","shell.execute_reply":"2022-08-05T11:39:59.055286Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 全量数据训练\n\n","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:22:41.948243Z","iopub.execute_input":"2022-08-05T11:22:41.948833Z","iopub.status.idle":"2022-08-05T11:22:47.487249Z","shell.execute_reply.started":"2022-08-05T11:22:41.948797Z","shell.execute_reply":"2022-08-05T11:22:47.485251Z"}}},{"cell_type":"code","source":"model = CustomModel(input_size=300, num_layers=3, d_model=512, ffn_size=1024, n_head=8)\noptimizer = torch.optim.Adam(model.parameters(), lr=0.0001)\n\nconfig = TrainerConfig(save_path=\"./test_model.bin\", batch_size=32, max_epochs=2)\ntrainer = Trainer(config, model, optimizer, train_dataset=train_ds_all, \n                  metric_fn=pearson_score, metric_larger_better=True)\ntrainer.fit()","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:40:03.814894Z","iopub.execute_input":"2022-08-05T11:40:03.815337Z","iopub.status.idle":"2022-08-05T11:40:18.560711Z","shell.execute_reply.started":"2022-08-05T11:40:03.815300Z","shell.execute_reply":"2022-08-05T11:40:18.559532Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 预测","metadata":{}},{"cell_type":"code","source":"# model = CustomModel(input_size=300, num_layers=3, d_model=512, ffn_size=1024, n_head=8)\n# model.load_state_dict(torch.load(\"./test_model.bin\")['checkpoint'])\n# trainer.model = model.to(config.device)","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:41:16.780402Z","iopub.execute_input":"2022-08-05T11:41:16.781103Z","iopub.status.idle":"2022-08-05T11:41:16.785928Z","shell.execute_reply.started":"2022-08-05T11:41:16.781064Z","shell.execute_reply":"2022-08-05T11:41:16.784903Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_pred = trainer.predict(test_ds)","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:41:17.978273Z","iopub.execute_input":"2022-08-05T11:41:17.979393Z","iopub.status.idle":"2022-08-05T11:41:19.638513Z","shell.execute_reply.started":"2022-08-05T11:41:17.979342Z","shell.execute_reply":"2022-08-05T11:41:19.636900Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test.shape","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:41:21.986387Z","iopub.execute_input":"2022-08-05T11:41:21.987250Z","iopub.status.idle":"2022-08-05T11:41:21.996246Z","shell.execute_reply.started":"2022-08-05T11:41:21.987205Z","shell.execute_reply":"2022-08-05T11:41:21.994934Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test['prediction'] = test_pred.reshape(-1)","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:41:23.286064Z","iopub.execute_input":"2022-08-05T11:41:23.286813Z","iopub.status.idle":"2022-08-05T11:41:23.293626Z","shell.execute_reply.started":"2022-08-05T11:41:23.286774Z","shell.execute_reply":"2022-08-05T11:41:23.292222Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test = test.loc[test['mask']==1, ['row_id', 'prediction']]\ntest.to_csv(\"submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:41:26.800733Z","iopub.execute_input":"2022-08-05T11:41:26.801140Z","iopub.status.idle":"2022-08-05T11:41:27.319574Z","shell.execute_reply.started":"2022-08-05T11:41:26.801106Z","shell.execute_reply":"2022-08-05T11:41:27.317226Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 确保提交与要求完全一致\nsample_submission = pd.read_csv(\"../input/boolart-market-prediction/sample_submission.csv\")\nassert np.all(test['row_id'] == sample_submission['row_id'])","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:41:32.241933Z","iopub.execute_input":"2022-08-05T11:41:32.242360Z","iopub.status.idle":"2022-08-05T11:41:32.324116Z","shell.execute_reply.started":"2022-08-05T11:41:32.242324Z","shell.execute_reply":"2022-08-05T11:41:32.323058Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_label = pd.read_csv(\"../input/boolart-market-prediction/test_label.csv\")","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:41:33.645941Z","iopub.execute_input":"2022-08-05T11:41:33.646358Z","iopub.status.idle":"2022-08-05T11:41:33.738254Z","shell.execute_reply.started":"2022-08-05T11:41:33.646321Z","shell.execute_reply":"2022-08-05T11:41:33.737223Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 计算 Test Mean Pearson Score","metadata":{}},{"cell_type":"code","source":"def LB_score(submission_file):\n    sub = pd.read_csv(submission_file)\n    test_label = pd.read_csv(\"../input/boolart-market-prediction/test_label.csv\")\n    test_label = test_label.merge(sub, on='row_id', how='left')\n    score= test_label.groupby('time_id').apply(lambda group: pearsonr(group['target'].values, group['prediction'].values)[0]).mean()\n    return score","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:41:42.049271Z","iopub.execute_input":"2022-08-05T11:41:42.049902Z","iopub.status.idle":"2022-08-05T11:41:42.056776Z","shell.execute_reply.started":"2022-08-05T11:41:42.049856Z","shell.execute_reply":"2022-08-05T11:41:42.055757Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"LB_score(\"submission.csv\")","metadata":{"execution":{"iopub.status.busy":"2022-08-05T11:41:42.592753Z","iopub.execute_input":"2022-08-05T11:41:42.593531Z","iopub.status.idle":"2022-08-05T11:41:42.828838Z","shell.execute_reply.started":"2022-08-05T11:41:42.593491Z","shell.execute_reply":"2022-08-05T11:41:42.827862Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 优化思路\n\n1. 尝试不同的模型结构或训练参数\n2. 尝试采用RNN或其他模型对每个investment_id 时序编码，再输入到Investment-Wise-Transformer\n3. 采用其他的机器学习或者深度学习模型建模，进行模型融合","metadata":{}}]}