{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":84493,"databundleVersionId":9871156,"sourceType":"competition"}],"dockerImageVersionId":30786,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport polars as pl\nfrom sklearn.preprocessing import StandardScaler, RobustScaler\nimport os\nfrom tqdm import tqdm\nimport gc\n\nfrom torch.utils.data import Dataset,DataLoader\nimport torch\nimport torch.nn as nn\n\nimport pickle\n\nfrom accelerate import Accelerator","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2025-01-11T04:48:02.790905Z","iopub.execute_input":"2025-01-11T04:48:02.791594Z","iopub.status.idle":"2025-01-11T04:48:02.796140Z","shell.execute_reply.started":"2025-01-11T04:48:02.791560Z","shell.execute_reply":"2025-01-11T04:48:02.795174Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import time\n\ndef calculate_time(func):\n    def wrapper(*args, **kwargs):\n        start_time = time.time()\n        result = func(*args, **kwargs)\n        end_time = time.time()\n        print(\"Function {} took {} seconds to execute.\".format(func.__name__, end_time - start_time))\n        return result\n    return wrapper","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T04:48:02.800405Z","iopub.execute_input":"2025-01-11T04:48:02.801143Z","iopub.status.idle":"2025-01-11T04:48:02.810341Z","shell.execute_reply.started":"2025-01-11T04:48:02.801104Z","shell.execute_reply":"2025-01-11T04:48:02.809268Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"1. use one-hot code to handle symbol_id\n2. change the date_id and time_id\n3. how to deal with weight\n4. standard scale\n5. implement the model","metadata":{}},{"cell_type":"markdown","source":"# 1.prepare data","metadata":{}},{"cell_type":"code","source":"def prepare_data(path):\n    N =50\n    raw_data = pl.read_parquet(path)\n    raw_data = raw_data.fill_null(0)\n    raw_data = raw_data.with_columns( \n        (pl.col('time_id').cast(pl.Float32)/848).alias('time_ids'),\n        (pl.col('date_id').cast(pl.Float32)/252).alias('date_ids'),\n        pl.col('symbol_id').map_elements(lambda x : [ 0 if y!=(x) else 1 for y in range(N)],return_dtype = pl.List(pl.Int64)).cast(pl.Array(pl.Int8,N)).alias('one_hot')\n        )\n    \n    for i in range(N):\n        raw_data = raw_data.with_columns(\n        pl.col('one_hot').arr.get(i).alias(\"symbol\"+str(i))\n        )\n\n    raw_data = raw_data.select(\n            pl.col(\n         'date_id',\n         'time_id',\n         'symbol_id',\n         'weight',\n         'feature_00',\n         'feature_01',\n         'feature_02',\n         'feature_03',\n         'feature_04',\n         'feature_05',\n         'feature_06',\n         'feature_07',\n         'feature_08',\n         'feature_09',\n         'feature_10',\n         'feature_11',\n         'feature_12',\n         'feature_13',\n         'feature_14',\n         'feature_15',\n         'feature_16',\n         'feature_17',\n         'feature_18',\n         'feature_19',\n         'feature_20',\n         'feature_21',\n         'feature_22',\n         'feature_23',\n         'feature_24',\n         'feature_25',\n         'feature_26',\n         'feature_27',\n         'feature_28',\n         'feature_29',\n         'feature_30',\n         'feature_31',\n         'feature_32',\n         'feature_33',\n         'feature_34',\n         'feature_35',\n         'feature_36',\n         'feature_37',\n         'feature_38',\n         'feature_39',\n         'feature_40',\n         'feature_41',\n         'feature_42',\n         'feature_43',\n         'feature_44',\n         'feature_45',\n         'feature_46',\n         'feature_47',\n         'feature_48',\n         'feature_49',\n         'feature_50',\n         'feature_51',\n         'feature_52',\n         'feature_53',\n         'feature_54',\n         'feature_55',\n         'feature_56',\n         'feature_57',\n         'feature_58',\n         'feature_59',\n         'feature_60',\n         'feature_61',\n         'feature_62',\n         'feature_63',\n         'feature_64',\n         'feature_65',\n         'feature_66',\n         'feature_67',\n         'feature_68',\n         'feature_69',\n         'feature_70',\n         'feature_71',\n         'feature_72',\n         'feature_73',\n         'feature_74',\n         'feature_75',\n         'feature_76',\n         'feature_77',\n         'feature_78',\n         'time_ids',\n         'date_ids',\n         'symbol0',\n         'symbol1',\n         'symbol2',\n         'symbol3',\n         'symbol4',\n         'symbol5',\n         'symbol6',\n         'symbol7',\n         'symbol8',\n         'symbol9',\n         'symbol10',\n         'symbol11',\n         'symbol12',\n         'symbol13',\n         'symbol14',\n         'symbol15',\n         'symbol16',\n         'symbol17',\n         'symbol18',\n         'symbol19',\n         'symbol20',\n         'symbol21',\n         'symbol22',\n         'symbol23',\n         'symbol24',\n         'symbol25',\n         'symbol26',\n         'symbol27',\n         'symbol28',\n         'symbol29',\n         'symbol30',\n         'symbol31',\n         'symbol32',\n         'symbol33',\n         'symbol34',\n         'symbol35',\n         'symbol36',\n         'symbol37',\n         'symbol38',\n            'symbol39',\n            'symbol40',\n            'symbol41',\n            'symbol42',\n            'symbol43',\n            'symbol44',\n            'symbol45',\n            'symbol46',\n            'symbol47',\n            'symbol48',\n            'symbol49',\n            # 'symbol50',\n            # 'symbol51',\n            # 'symbol52',\n            # 'symbol53',\n            # 'symbol54',\n            # 'symbol55',\n            # 'symbol56',\n            # 'symbol57',\n            # 'symbol58',\n            # 'symbol59',\n            # 'symbol60',\n            # 'symbol61',\n            # 'symbol62',\n            # 'symbol63',\n            # 'symbol64',\n            # 'symbol65',\n            # 'symbol66',\n            # 'symbol67',\n            # 'symbol68',\n            # 'symbol69',\n            # 'symbol70',\n            # 'symbol71',\n            # 'symbol72',\n            # 'symbol73',\n            # 'symbol74',\n            # 'symbol75',\n            # 'symbol76',\n            # 'symbol77',\n            # 'symbol78',\n            # 'symbol79',\n            # 'symbol80',\n            # 'symbol81',\n            # 'symbol82',\n            # 'symbol83',\n            # 'symbol84',\n            # 'symbol85',\n            # 'symbol86',\n            # 'symbol87',\n            # 'symbol88',\n            # 'symbol89',\n            # 'symbol90',\n            # 'symbol91',\n            # 'symbol92',\n            # 'symbol93',\n            # 'symbol94',\n            # 'symbol95',\n            # 'symbol96',\n            # 'symbol97',\n            # 'symbol98',\n            # 'symbol99',\n            'responder_0',\n         'responder_1',\n         'responder_2',\n         'responder_3',\n         'responder_4',\n         'responder_5',\n         'responder_6',\n         'responder_7',\n         'responder_8'\n            )\n        )\n\n    return raw_data\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T04:48:02.811856Z","iopub.execute_input":"2025-01-11T04:48:02.812102Z","iopub.status.idle":"2025-01-11T04:48:02.826887Z","shell.execute_reply.started":"2025-01-11T04:48:02.812079Z","shell.execute_reply":"2025-01-11T04:48:02.826018Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class datapreprocessor:\n    def __init__(self,data,process_list):\n        self.raw_data = data\n        self.need_processing_list = process_list\n\n        self.scaler_dict = {}\n        self.scaled_df = pl.DataFrame()\n        \n    def training_scale(self):\n        _i = 0\n        for tag in self.raw_data.columns:\n            if tag in self.need_processing_list:\n                self.scaler_dict[tag] = StandardScaler()\n                _series = np.expand_dims(self.raw_data[tag],axis=1)\n                self.scaler_dict[tag].fit(_series)\n                self.scaled_df = self.scaled_df.insert_column(\n                    _i,pl.Series(tag, np.squeeze(self.scaler_dict[tag].transform(_series),axis = 1))\n                )\n            else:\n                self.scaled_df = self.scaled_df.insert_column(\n                    _i,pl.Series(tag, self.raw_data[tag].alias(tag) )\n                )\n            _i += 1\n    \n    def inferrence_scale(self,data):\n        _i = 0\n        self.scaled_df = pl.DataFrame()\n        for tag in data.columns:\n            if tag in self.need_processing_list:\n                # self.scaler_dict[tag] = StandardScaler()\n                _series = np.expand_dims(data[tag],axis=1)\n                # self.scaler_dict[tag].fit(_series)\n                self.scaled_df = self.scaled_df.insert_column(\n                    _i,pl.Series(tag, np.squeeze(self.scaler_dict[tag].transform(_series),axis = 1))\n                )\n            else:\n                self.scaled_df = self.scaled_df.insert_column(\n                    _i,pl.Series(tag, data[tag].alias(tag) )\n                )\n            _i += 1\n    \n    def reverse_scale(self):\n        pass","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T04:48:02.827990Z","iopub.execute_input":"2025-01-11T04:48:02.828373Z","iopub.status.idle":"2025-01-11T04:48:02.844395Z","shell.execute_reply.started":"2025-01-11T04:48:02.828337Z","shell.execute_reply":"2025-01-11T04:48:02.843595Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"tag = [ 'feature_00',\n 'feature_01',\n 'feature_02',\n 'feature_03',\n 'feature_04',\n 'feature_05',\n 'feature_06',\n 'feature_07',\n 'feature_08',\n 'feature_09',\n 'feature_10',\n 'feature_11',\n 'feature_12',\n 'feature_13',\n 'feature_14',\n 'feature_15',\n 'feature_16',\n 'feature_17',\n 'feature_18',\n 'feature_19',\n 'feature_20',\n 'feature_21',\n 'feature_22',\n 'feature_23',\n 'feature_24',\n 'feature_25',\n 'feature_26',\n 'feature_27',\n 'feature_28',\n 'feature_29',\n 'feature_30',\n 'feature_31',\n 'feature_32',\n 'feature_33',\n 'feature_34',\n 'feature_35',\n 'feature_36',\n 'feature_37',\n 'feature_38',\n 'feature_39',\n 'feature_40',\n 'feature_41',\n 'feature_42',\n 'feature_43',\n 'feature_44',\n 'feature_45',\n 'feature_46',\n 'feature_47',\n 'feature_48',\n 'feature_49',\n 'feature_50',\n 'feature_51',\n 'feature_52',\n 'feature_53',\n 'feature_54',\n 'feature_55',\n 'feature_56',\n 'feature_57',\n 'feature_58',\n 'feature_59',\n 'feature_60',\n 'feature_61',\n 'feature_62',\n 'feature_63',\n 'feature_64',\n 'feature_65',\n 'feature_66',\n 'feature_67',\n 'feature_68',\n 'feature_69',\n 'feature_70',\n 'feature_71',\n 'feature_72',\n 'feature_73',\n 'feature_74',\n 'feature_75',\n 'feature_76',\n 'feature_77',\n 'feature_78']","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T04:48:02.846010Z","iopub.execute_input":"2025-01-11T04:48:02.846261Z","iopub.status.idle":"2025-01-11T04:48:02.857592Z","shell.execute_reply.started":"2025-01-11T04:48:02.846236Z","shell.execute_reply":"2025-01-11T04:48:02.856747Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 2 dataset and dataloader","metadata":{}},{"cell_type":"code","source":"class fast_dataset(Dataset):\n    def __init__(self,data,T=64):\n        self.data = data.with_row_index('id_original')\n        self.T = T\n        self.symbols = set(self.data['symbol_id'].value_counts()['symbol_id'])\n        # self.lag = len(self.symbols)*self.T\n        self.data_dict = {}\n\n        for tag in self.symbols:\n            self.data_dict[tag] = self.data.lazy().filter(\n                pl.col('symbol_id')==tag\n            ).collect()\n\n        _idx_record = 0\n\n        self.hash_df = pl.DataFrame(schema={'idx': pl.Int64, 'tag':  pl.Int32})\n        for tag in self.symbols:\n            _series = np.arange(len(self.data_dict[tag]))- self.T + _idx_record\n            self.data_dict[tag].insert_column(0,pl.Series('idx',_series))\n            _idx_record = _series[-1]\n        \n            self.hash_df = self.hash_df.vstack(pl.DataFrame({'idx':_series[T+1:],'tag': tag}))\n        \n        \n    def __len__(self):\n        return len(self.hash_df) \n        \n    def __getitem__(self,idx):\n        idx = idx + 1\n\n        _item = torch.FloatTensor(np.array( (self.data_dict[self.hash_df.lazy().filter(\n                                                pl.col('idx') == idx\n                                            ).collect()['tag'][0]]).lazy().filter(\n                                                (pl.col('idx')>=(idx-self.T))&(pl.col('idx')<=(idx))\n                                            ).collect()))\n        return _item[:,6:-9], _item[-1,-9:],_item[-1,5]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T04:48:02.858615Z","iopub.execute_input":"2025-01-11T04:48:02.858868Z","iopub.status.idle":"2025-01-11T04:48:02.870521Z","shell.execute_reply.started":"2025-01-11T04:48:02.858845Z","shell.execute_reply":"2025-01-11T04:48:02.869708Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# batch_size = 128\n# seq_length = 64\n\n# training_set = fast_dataset(dfprocessor.scaled_df,seq_length)\n\n# t_loader = DataLoader(training_set, batch_size=batch_size,shuffle=True)\n\n# start = time.time()\n# test_x , test_y, test_weight = next(iter(t_loader))\n# print('run_time',time.time()-start)\n\n# print(test_x.shape , test_y.shape)\n\n# print(training_set.data_dict[0][0])\n\n# fenzi = (test_weight*((test_y[:,6] - test_y_pred[:,6])**2)).sum()\n# fenmu = (test_weight*((test_y[:,6])**2)).sum()\n# 1-(fenzi/fenmu)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T04:48:02.871466Z","iopub.execute_input":"2025-01-11T04:48:02.871726Z","iopub.status.idle":"2025-01-11T04:48:02.883107Z","shell.execute_reply.started":"2025-01-11T04:48:02.871695Z","shell.execute_reply":"2025-01-11T04:48:02.882113Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# del(training_set)\n# del(t_loader)\n# del(test_x)\n# del(test_y)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T04:48:02.884138Z","iopub.execute_input":"2025-01-11T04:48:02.884413Z","iopub.status.idle":"2025-01-11T04:48:02.893221Z","shell.execute_reply.started":"2025-01-11T04:48:02.884388Z","shell.execute_reply":"2025-01-11T04:48:02.892542Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 3. model","metadata":{}},{"cell_type":"code","source":"class PositionalEncoding(nn.Module):\n    def __init__(self, d_model, max_len=5000):\n        super().__init__()\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() * (-np.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        return x + self.pe[:, :x.size(1)]\n\n\nclass T_Transformer(nn.Module):\n    def __init__(self, input_dim, d_model, nhead, num_layers,dropout):\n        super().__init__()\n\n        self.onehot_embedding = nn.Linear(50, 4)\n        \n        self.embedding = nn.Linear(input_dim, d_model)\n        self.pos_encoder = PositionalEncoding(d_model)\n        \n        encoder_layer = nn.TransformerEncoderLayer(d_model=d_model, nhead=nhead,batch_first=True,dropout=dropout)\n        self.transformer_encoder = nn.TransformerEncoder(encoder_layer, num_layers=num_layers)\n        \n        self.output_layer1 = nn.Linear(d_model, d_model//2)\n        self.output_layer = nn.Linear(d_model//2, 9)\n        \n    def forward(self,src_whole):\n        vec1 = src_whole[:,:,:-50]\n        vec2 = src_whole[:,:,-50:]\n\n        vec2_2 = self.onehot_embedding(vec2)\n\n        src = torch.cat((vec1,vec2_2),2)\n        \n        src = self.embedding(src)\n        src = self.pos_encoder(src)\n        encoder_output = self.transformer_encoder(src)\n        \n        output1 = self.output_layer1(encoder_output[:,-1,:])\n        output = self.output_layer(output1)\n        \n        return output\n\n    # def forward(self,src):\n\n    #     src = self.embedding(src)\n    #     src = self.pos_encoder(src)\n    #     encoder_output = self.transformer_encoder(src)\n        \n    #     output1 = self.output_layer1(encoder_output[:,-1,:])\n    #     output = self.output_layer(output1)\n        \n    #     return output\n\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T04:48:02.972252Z","iopub.execute_input":"2025-01-11T04:48:02.972603Z","iopub.status.idle":"2025-01-11T04:48:02.982025Z","shell.execute_reply.started":"2025-01-11T04:48:02.972574Z","shell.execute_reply":"2025-01-11T04:48:02.980997Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# from accelerate import Accelerator\n# accelerator = Accelerator()\n# device = accelerator.device","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T04:48:02.983387Z","iopub.execute_input":"2025-01-11T04:48:02.983886Z","iopub.status.idle":"2025-01-11T04:48:02.996185Z","shell.execute_reply.started":"2025-01-11T04:48:02.983858Z","shell.execute_reply":"2025-01-11T04:48:02.995363Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def training(data_id,model,device,criterion,optimizer,num_epochs,t_loader,v_loader,accelerator):\n\n    train_losses = []\n    val_losses = []\n    best_val_loss = float('inf')\n\n    val_loss = 0\n\n    \n    # Validation\n    if data_id>=0:\n        model.eval()\n        val_loss = 0\n        with torch.no_grad():\n            for x_batch, y_batch ,weight_batch in v_loader:\n                x_batch, y_batch,weight_batch = x_batch.to(device), y_batch.to(device),weight_batch.to(device)\n                y_pred = model(x_batch)\n                val_loss += criterion(y_pred, y_batch).item()\n    \n        val_loss /= len(v_loader)\n    \n        \n        \n        sse_vali = ((weight_batch*((y_batch[:,6] - y_pred[:,6])**2)).sum()).item()\n        sst_vali = ((weight_batch*((y_batch[:,6])**2)).sum()).item()\n    \n        print('parquet_id: ', data_id,' validation_loss: ',val_loss,' R_square: ',1-(sse_vali/sst_vali))\n\n    val_losses.append(val_loss)\n    # Training\n    \n    model, optimizer, t_loader = accelerator.prepare(model, optimizer, t_loader)\n    \n    for epoch in range(num_epochs):\n        \n        model.train()\n        train_loss = 0\n\n        sst = 0\n        sse = 0\n        \n        loop = tqdm(enumerate(t_loader), total=len(t_loader))\n        \n        for step,(x_batch, y_batch,weight_batch) in loop:\n            # x_batch, y_batch, weight_batch = x_batch.to(device), y_batch.to(device),weight_batch.to(device)\n\n            optimizer.zero_grad()\n            y_pred = model(x_batch)\n            loss = criterion(y_pred, y_batch)\n\n            # loss.backward()\n            accelerator.backward(loss)\n            \n            optimizer.step()\n\n            train_loss += loss.item() \n\n            sse += ((weight_batch*((y_batch[:,6] - y_pred[:,6])**2)).sum()).item()\n            sst += ((weight_batch*((y_batch[:,6])**2)).sum()).item()\n            \n\n            loop.set_description(f'Epoch [{epoch+1}/{num_epochs}]')\n            loop.set_postfix(loss=train_loss/(step+1),r_square= 1-(sse/sst))\n            \n            \n        torch.save(model.state_dict(), 'model_'+'parquet_'+str(data_id)+'_epoch_'+str(epoch+1)+'.pth')\n        \n        train_loss /= len(t_loader)\n\n        train_losses.append(train_loss)\n        \n\n        # if val_loss < best_val_loss:\n        #     best_val_loss = val_loss\n        #     torch.save(model.state_dict(), 'best_model'+str(epoch)+'.pth')\n\n        # if epoch==9 :\n        #     best_val_loss = val_loss\n        \n\n        # if epoch % 10 == 0:\n        if True:\n            print(f'Epoch {epoch+1}: Train Loss = {train_loss:.4f}, Val Loss = {val_loss:.4f}')\n    \n    return train_losses,val_losses\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T07:42:20.032861Z","iopub.execute_input":"2025-01-11T07:42:20.033525Z","iopub.status.idle":"2025-01-11T07:42:20.044888Z","shell.execute_reply.started":"2025-01-11T07:42:20.033492Z","shell.execute_reply":"2025-01-11T07:42:20.043956Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# def training(data_id,model,device,criterion,optimizer,num_epochs,t_loader,v_loader):\n\n#     train_losses = []\n#     val_losses = []\n#     best_val_loss = float('inf')\n\n#     val_loss = 0\n#     # Validation\n#     if data_id>=0:\n#         model.eval()\n#         val_loss = 0\n#         with torch.no_grad():\n#             for x_batch, y_batch ,weight_batch in v_loader:\n#                 x_batch, y_batch,weight_batch = x_batch.to(device), y_batch.to(device),weight_batch.to(device)\n#                 y_pred = model(x_batch)\n#                 val_loss += criterion(y_pred, y_batch).item()\n    \n#         val_loss /= len(v_loader)\n    \n        \n        \n#         sse_vali = ((weight_batch*((y_batch[:,6] - y_pred[:,6])**2)).sum()).item()\n#         sst_vali = ((weight_batch*((y_batch[:,6])**2)).sum()).item()\n    \n#         print('parquet_id: ', data_id,' validation_loss: ',val_loss,' R_square: ',1-(sse_vali/sst_vali))\n\n#     val_losses.append(val_loss)\n#     # Training\n#     for epoch in range(num_epochs):\n        \n#         model.train()\n#         train_loss = 0\n\n#         sst = 0\n#         sse = 0\n        \n#         loop = tqdm(enumerate(t_loader), total=len(t_loader))\n        \n#         for step,(x_batch, y_batch,weight_batch) in loop:\n#             x_batch, y_batch, weight_batch = x_batch.to(device), y_batch.to(device),weight_batch.to(device)\n\n#             optimizer.zero_grad()\n#             y_pred = model(x_batch)\n#             loss = criterion(y_pred, y_batch)\n\n#             loss.backward()\n#             optimizer.step()\n\n#             train_loss += loss.item() \n\n#             sse += ((weight_batch*((y_batch[:,6] - y_pred[:,6])**2)).sum()).item()\n#             sst += ((weight_batch*((y_batch[:,6])**2)).sum()).item()\n            \n\n#             loop.set_description(f'Epoch [{epoch+1}/{num_epochs}]')\n#             loop.set_postfix(loss=train_loss/(step+1),r_square= 1-(sse/sst))\n            \n#             torch.save(model.state_dict(), 'model_'+'parquet_'+str(data_id)+'_epoch_'+str(epoch+1)+'.pth')\n\n#         train_loss /= len(t_loader)\n\n\n\n#         train_losses.append(train_loss)\n        \n\n#         # if val_loss < best_val_loss:\n#         #     best_val_loss = val_loss\n#         #     torch.save(model.state_dict(), 'best_model'+str(epoch)+'.pth')\n\n#         # if epoch==9 :\n#         #     best_val_loss = val_loss\n        \n\n#         # if epoch % 10 == 0:\n#         if True:\n#             print(f'Epoch {epoch+1}: Train Loss = {train_loss:.4f}, Val Loss = {val_loss:.4f}')\n    \n#     return train_losses,val_losses\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T04:48:03.012124Z","iopub.execute_input":"2025-01-11T04:48:03.012403Z","iopub.status.idle":"2025-01-11T04:48:03.025755Z","shell.execute_reply.started":"2025-01-11T04:48:03.012378Z","shell.execute_reply":"2025-01-11T04:48:03.025004Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def test_model(model,device,criterion,test_loader):\n    model.eval()\n    test_loss = 0\n    predictions = []\n    actuals = []\n    with torch.no_grad():\n        for x_batch, y_batch in test_loader:\n            x_batch, y_batch = x_batch.to(device), y_batch.to(device)\n            y_pred = model(x_batch)\n            test_loss += criterion(y_pred, y_batch).item()\n\n            # print(np.shape(y_pred.cpu().numpy()))\n\n            predictions.extend(y_pred.cpu().numpy())\n            actuals.extend(y_batch.cpu().numpy())\n\n    test_loss /= len(test_loader)\n    return test_loss,np.array(predictions),np.array(actuals)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T04:48:03.027931Z","iopub.execute_input":"2025-01-11T04:48:03.028674Z","iopub.status.idle":"2025-01-11T04:48:03.042631Z","shell.execute_reply.started":"2025-01-11T04:48:03.028635Z","shell.execute_reply":"2025-01-11T04:48:03.041854Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# dfprocessor = datapreprocessor(raw_data,tag)\n# dfprocessor.training_scale()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T04:48:03.043773Z","iopub.execute_input":"2025-01-11T04:48:03.044355Z","iopub.status.idle":"2025-01-11T04:48:03.055094Z","shell.execute_reply.started":"2025-01-11T04:48:03.044295Z","shell.execute_reply":"2025-01-11T04:48:03.054343Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sequence_length = 64\n\n# batch_size = 128\nbatch_size = 1024\n\nd_imput = 85\nd_model = 48\nnhead = 8\nnum_layers = 6\nlearn_rate = 1e-4\ndropout = 0.1","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T04:48:03.056062Z","iopub.execute_input":"2025-01-11T04:48:03.056362Z","iopub.status.idle":"2025-01-11T04:48:03.068085Z","shell.execute_reply.started":"2025-01-11T04:48:03.056300Z","shell.execute_reply":"2025-01-11T04:48:03.067288Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"accelerator = Accelerator()\ndevice = accelerator.device","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T04:48:03.069094Z","iopub.execute_input":"2025-01-11T04:48:03.069371Z","iopub.status.idle":"2025-01-11T04:48:03.110397Z","shell.execute_reply.started":"2025-01-11T04:48:03.069344Z","shell.execute_reply":"2025-01-11T04:48:03.109523Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## load model and scaler","metadata":{}},{"cell_type":"code","source":"t_model = T_Transformer(d_imput,d_model,nhead,num_layers,dropout)\nt_model.load_state_dict(torch.load('/kaggle/input/save/pytorch/default/1/model_save.pth'))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T04:48:03.111697Z","iopub.execute_input":"2025-01-11T04:48:03.112034Z","iopub.status.idle":"2025-01-11T04:48:03.464250Z","shell.execute_reply.started":"2025-01-11T04:48:03.111999Z","shell.execute_reply":"2025-01-11T04:48:03.463354Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# t_model = T_Transformer(d_imput,d_model,nhead,num_layers,dropout)\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nt_model = t_model.to(device)\ncriterion = nn.MSELoss()\noptimizer = torch.optim.Adam(t_model.parameters(), lr=learn_rate)\n\n# num_epochs = 10\nnum_epochs = 1","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T04:48:03.465346Z","iopub.execute_input":"2025-01-11T04:48:03.465615Z","iopub.status.idle":"2025-01-11T04:48:04.358808Z","shell.execute_reply.started":"2025-01-11T04:48:03.465590Z","shell.execute_reply":"2025-01-11T04:48:04.358118Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T04:48:04.361142Z","iopub.execute_input":"2025-01-11T04:48:04.361580Z","iopub.status.idle":"2025-01-11T04:48:04.366294Z","shell.execute_reply.started":"2025-01-11T04:48:04.361551Z","shell.execute_reply":"2025-01-11T04:48:04.365360Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# size = len(dfprocessor.scaled_df)\n\n# training_size = int(0.8*size)\n# vali_size = int(0.1*size)\n# test_size = int(0.1*size)\n\n# training_set = fast_dataset(dfprocessor.scaled_df[0:training_size],sequence_length)\n# vali_set = fast_dataset(dfprocessor.scaled_df[training_size-sequence_length:training_size+vali_size],sequence_length)\n# test_set = fast_dataset(dfprocessor.scaled_df[training_size+vali_size-sequence_length:],sequence_length)\n\n# len(training_set),len(vali_set),len(test_set)\n\n# t_loader = DataLoader(training_set, batch_size=batch_size,shuffle=True)\n# v_loader = DataLoader(vali_set, batch_size=batch_size)\n# test_loader = DataLoader(test_set, batch_size=batch_size)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T04:48:04.367398Z","iopub.execute_input":"2025-01-11T04:48:04.367699Z","iopub.status.idle":"2025-01-11T04:48:04.377736Z","shell.execute_reply.started":"2025-01-11T04:48:04.367664Z","shell.execute_reply":"2025-01-11T04:48:04.377047Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# train_loss,val_loss = training(\n#     model = t_model,\n#     device = device,\n#     criterion = criterion,\n#     optimizer = optimizer,\n#     num_epochs =num_epochs,\n#     t_loader = t_loader,\n#     v_loader = v_loader\n# )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T04:48:04.378912Z","iopub.execute_input":"2025-01-11T04:48:04.379726Z","iopub.status.idle":"2025-01-11T04:48:04.388541Z","shell.execute_reply.started":"2025-01-11T04:48:04.379646Z","shell.execute_reply":"2025-01-11T04:48:04.387769Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 4.training ","metadata":{}},{"cell_type":"markdown","source":"可以通过在新的数据集上选头几列先进行validation，看看验证效果，这样在训练时也能利用全部数据","metadata":{}},{"cell_type":"code","source":"# 从文件中加载对象\nwith open(\"dfprocessor.pickle\", \"rb\") as file:\n    dfprocessor = pickle.load(file)\n\nprint(dfprocessor.scaler_dict) # 输出：Alice\nprint(dfprocessor.scaled_df.head()) # 输出：25","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"raw_data = prepare_data('/kaggle/input/jane-street-real-time-market-data-forecasting/train.parquet/partition_id=0/part-0.parquet')\ndfprocessor = datapreprocessor(raw_data,tag)\ndfprocessor.training_scale()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# A = range(5)\n# A = [3,4]\nB = range(4,10)\nfor i in B:\n\n    print('parquet_data_id : ',i)\n\n    del(raw_data)\n    gc.collect()\n\n    # read data\n    \n    raw_data = prepare_data('/kaggle/input/jane-street-real-time-market-data-forecasting/train.parquet/partition_id='+str(i)+'/part-0.parquet')\n\n    print(i,'：data_loaded')\n    # scale \n    \n    if i==0:\n        dfprocessor = datapreprocessor(raw_data,tag)\n        dfprocessor.training_scale()\n    else:\n        dfprocessor.inferrence_scale(raw_data)\n\n    print(i,':data_scaled')\n    # dataset\n\n    size = len(dfprocessor.scaled_df)\n\n    # training_size = int(0.9*size)\n    vali_size = int(0.05*size)\n    # test_size = int(0.1*size)\n    \n    training_set = fast_dataset(dfprocessor.scaled_df,sequence_length)\n    vali_set = fast_dataset(dfprocessor.scaled_df[:vali_size],sequence_length)\n    # test_set = fast_dataset(dfprocessor.scaled_df[training_size+vali_size-sequence_length:],sequence_length)\n    \n    # len(training_set),len(vali_set)\n    # ,len(test_set)\n    \n    t_loader = DataLoader(training_set, batch_size=batch_size,shuffle=True,pin_memory=True)\n    v_loader = DataLoader(vali_set, batch_size=batch_size)\n    # test_loader = DataLoader(test_set, batch_size=batch_size)\n\n    # train\n    print(i,':training start!!')\n    \n    train_loss,val_loss = training(\n        data_id = i,\n        model = t_model,\n        device = device,\n        criterion = criterion,\n        optimizer = optimizer,\n        num_epochs =num_epochs,\n        t_loader = t_loader,\n        v_loader = v_loader,\n        accelerator = accelerator\n    )\n    \n    # train_loss,val_loss = training(\n    #     data_id = i,\n    #     model = t_model,\n    #     device = device,\n    #     criterion = criterion,\n    #     optimizer = optimizer,\n    #     num_epochs =num_epochs,\n    #     t_loader = t_loader,\n    #     v_loader = v_loader\n    # )\n    \n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T07:43:23.408071Z","iopub.execute_input":"2025-01-11T07:43:23.408411Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\ntorch.save(t_model.state_dict(), 'model_save'+str(saveid)+'.pth')\nsaveid += 1","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T07:41:52.598999Z","iopub.execute_input":"2025-01-11T07:41:52.599773Z","iopub.status.idle":"2025-01-11T07:41:52.620757Z","shell.execute_reply.started":"2025-01-11T07:41:52.599738Z","shell.execute_reply":"2025-01-11T07:41:52.620055Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"with open(\"dfprocessor.pickle\", \"wb\") as file:\n    pickle.dump(dfprocessor, file)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# t_model_load.eval()\n# val_loss = 0\n# with torch.no_grad():\n#     for x_batch, y_batch ,weight_batch in v_loader:\n#         x_batch, y_batch,weight_batch = x_batch.to(device), y_batch.to(device),weight_batch.to(device)\n#         y_pred = t_model_load(x_batch)\n#         val_loss += criterion(y_pred, y_batch).item()\n\n# val_loss /= len(v_loader)\n\n\n\n# sse_vali = ((weight_batch*((y_batch[:,6] - y_pred[:,6])**2)).sum()).item()\n# sst_vali = ((weight_batch*((y_batch[:,6])**2)).sum()).item()\n\n# print('parquet_id: ', 1,' validation_loss: ',val_loss,' R_square: ',1-(sse_vali/sst_vali))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-11T04:48:39.845470Z","iopub.status.idle":"2025-01-11T04:48:39.845784Z","shell.execute_reply.started":"2025-01-11T04:48:39.845638Z","shell.execute_reply":"2025-01-11T04:48:39.845653Z"}},"outputs":[],"execution_count":null}]}