{"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"},{"sourceId":207293230,"sourceType":"kernelVersion"}],"dockerImageVersionId":30787,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install -q --requirement /kaggle/input/yunbase/Yunbase/requirements.txt  \\\n--no-index --find-links file:/kaggle/input/yunbase/","metadata":{"execution":{"iopub.status.busy":"2024-11-14T07:31:34.760201Z","iopub.execute_input":"2024-11-14T07:31:34.760517Z","iopub.status.idle":"2024-11-14T07:31:47.594642Z","shell.execute_reply.started":"2024-11-14T07:31:34.760482Z","shell.execute_reply":"2024-11-14T07:31:47.593408Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"source_file_path = '/kaggle/input/yunbase/Yunbase/baseline.py'\ntarget_file_path = '/kaggle/working/baseline.py'\nwith open(source_file_path, 'r', encoding='utf-8') as file:\n    content = file.read()\nwith open(target_file_path, 'w', encoding='utf-8') as file:\n    file.write(content)","metadata":{"execution":{"iopub.status.busy":"2024-11-14T07:31:47.59711Z","iopub.execute_input":"2024-11-14T07:31:47.597983Z","iopub.status.idle":"2024-11-14T07:31:47.610824Z","shell.execute_reply.started":"2024-11-14T07:31:47.597914Z","shell.execute_reply":"2024-11-14T07:31:47.610072Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from baseline import Yunbase\nimport polars as pl#similar to pandas, but with better performance when dealing with large datasets.\nimport pandas as pd#read csv,parquet\nimport numpy as np#for scientific computation of matrices\n#model\nfrom  lightgbm import LGBMRegressor\nfrom catboost import CatBoostRegressor\nfrom xgboost import XGBRegressor\nimport os#Libraries that interact with the operating system\nimport gc#rubbish collection\n#environment provided by competition hoster\nimport kaggle_evaluation.jane_street_inference_server\nimport time\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm","metadata":{"execution":{"iopub.status.busy":"2024-11-14T07:31:47.611865Z","iopub.execute_input":"2024-11-14T07:31:47.612217Z","iopub.status.idle":"2024-11-14T07:31:52.898141Z","shell.execute_reply.started":"2024-11-14T07:31:47.612184Z","shell.execute_reply":"2024-11-14T07:31:52.897155Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# load data\n# t1 = time.time()\n# df_0 = pl.read_parquet('../input/jane-street-real-time-market-data-forecasting/train.parquet/partition_id=0/part-0.parquet')\n# df_1 = pl.read_parquet('../input/jane-street-real-time-market-data-forecasting/train.parquet/partition_id=1/part-0.parquet')\n# df_2 = pl.read_parquet('../input/jane-street-real-time-market-data-forecasting/train.parquet/partition_id=2/part-0.parquet')\n# df_3 = pl.read_parquet('../input/jane-street-real-time-market-data-forecasting/train.parquet/partition_id=3/part-0.parquet')\n# df_4 = pl.read_parquet('../input/jane-street-real-time-market-data-forecasting/train.parquet/partition_id=4/part-0.parquet')\n# df_5 = pl.read_parquet('../input/jane-street-real-time-market-data-forecasting/train.parquet/partition_id=5/part-0.parquet')\n# df_6 = pl.read_parquet('../input/jane-street-real-time-market-data-forecasting/train.parquet/partition_id=6/part-0.parquet')\n# df_7 = pl.read_parquet('../input/jane-street-real-time-market-data-forecasting/train.parquet/partition_id=7/part-0.parquet')\n# df_8 = pl.read_parquet('../input/jane-street-real-time-market-data-forecasting/train.parquet/partition_id=8/part-0.parquet')\n# df_9 = pl.read_parquet('../input/jane-street-real-time-market-data-forecasting/train.parquet/partition_id=9/part-0.parquet')\n# t2 = time.time()\n# print('Elapsed time [s]:', np.round(t2-t1,4))","metadata":{"execution":{"iopub.status.busy":"2024-11-14T07:31:52.900329Z","iopub.execute_input":"2024-11-14T07:31:52.900828Z","iopub.status.idle":"2024-11-14T07:31:52.905775Z","shell.execute_reply.started":"2024-11-14T07:31:52.900793Z","shell.execute_reply":"2024-11-14T07:31:52.904892Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# combine in one data frame\n# df = pl.concat([df_0,df_1,df_2,df_3,df_4,df_5,df_6,df_7,df_8,df_9])\n# df = pl.concat([df_5,df_6,df_7,df_8,df_9])","metadata":{"execution":{"iopub.status.busy":"2024-11-14T07:31:52.907061Z","iopub.execute_input":"2024-11-14T07:31:52.90741Z","iopub.status.idle":"2024-11-14T07:31:52.922896Z","shell.execute_reply.started":"2024-11-14T07:31:52.907366Z","shell.execute_reply":"2024-11-14T07:31:52.922008Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# clean up\n# del df_5,df_6,df_7,df_8,df_9\n# gc.collect();","metadata":{"execution":{"iopub.status.busy":"2024-11-14T07:31:52.923978Z","iopub.execute_input":"2024-11-14T07:31:52.924275Z","iopub.status.idle":"2024-11-14T07:31:52.933582Z","shell.execute_reply.started":"2024-11-14T07:31:52.924241Z","shell.execute_reply":"2024-11-14T07:31:52.932673Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# def plot_correlation_heatmap(df, figsize=(10, 8), cmap='coolwarm'):\n#     \"\"\"\n#     Plots a correlation heatmap for the given DataFrame.\n    \n#     Parameters:\n#     - df (pd.DataFrame): The input DataFrame with numerical features.\n#     - figsize (tuple): Figure size for the heatmap plot.\n#     - cmap (str): Colormap for the heatmap.\n#     \"\"\"\n#     # Calculate the correlation matrix\n#     corr_matrix = df.corr()\n    \n#     # Set up the matplotlib figure\n#     plt.figure(figsize=figsize)\n    \n#     # Draw the heatmap with Seaborn\n#     sns.heatmap(corr_matrix, cmap=cmap, vmin=-1, vmax=1, square=True, linewidths=0.5)\n    \n#     # Display the plot\n#     plt.title(\"Feature Correlation Heatmap\")\n#     plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-11-14T07:31:52.934647Z","iopub.execute_input":"2024-11-14T07:31:52.93496Z","iopub.status.idle":"2024-11-14T07:31:52.946303Z","shell.execute_reply.started":"2024-11-14T07:31:52.9349Z","shell.execute_reply":"2024-11-14T07:31:52.945498Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# df_0 = pl.read_parquet('../input/jane-street-real-time-market-data-forecasting/train.parquet/partition_id=7/part-0.parquet')\n# #df_0=df_0.to_pandas()\n# df_0.head()","metadata":{"execution":{"iopub.status.busy":"2024-11-14T07:31:52.947278Z","iopub.execute_input":"2024-11-14T07:31:52.947549Z","iopub.status.idle":"2024-11-14T07:31:52.960708Z","shell.execute_reply.started":"2024-11-14T07:31:52.947513Z","shell.execute_reply":"2024-11-14T07:31:52.960006Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# #fill missing values\n# df_0=df_0.fillna(method='ffill')\n# df_0=df_0.fillna(0)\n# df_0.head()","metadata":{"execution":{"iopub.status.busy":"2024-11-14T07:31:52.961777Z","iopub.execute_input":"2024-11-14T07:31:52.962073Z","iopub.status.idle":"2024-11-14T07:31:52.971722Z","shell.execute_reply.started":"2024-11-14T07:31:52.962042Z","shell.execute_reply":"2024-11-14T07:31:52.970892Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# plot_correlation_heatmap(df_0)","metadata":{"execution":{"iopub.status.busy":"2024-11-14T07:31:52.975125Z","iopub.execute_input":"2024-11-14T07:31:52.97538Z","iopub.status.idle":"2024-11-14T07:31:52.983076Z","shell.execute_reply.started":"2024-11-14T07:31:52.975352Z","shell.execute_reply":"2024-11-14T07:31:52.982221Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# df_1 = pl.read_parquet('../input/jane-street-real-time-market-data-forecasting/train.parquet/partition_id=8/part-0.parquet')\n# df_1=df_1.to_pandas()\n# df_1.head()","metadata":{"execution":{"iopub.status.busy":"2024-11-14T07:31:52.984088Z","iopub.execute_input":"2024-11-14T07:31:52.984431Z","iopub.status.idle":"2024-11-14T07:31:52.993127Z","shell.execute_reply.started":"2024-11-14T07:31:52.98439Z","shell.execute_reply":"2024-11-14T07:31:52.992327Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# #fill missing values\n# df_1=df_1.fillna(method='ffill')\n# df_1=df_1.fillna(0)\n# df_1.head()","metadata":{"execution":{"iopub.status.busy":"2024-11-14T07:31:52.994147Z","iopub.execute_input":"2024-11-14T07:31:52.994413Z","iopub.status.idle":"2024-11-14T07:31:53.00445Z","shell.execute_reply.started":"2024-11-14T07:31:52.994384Z","shell.execute_reply":"2024-11-14T07:31:53.00366Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# plot_correlation_heatmap(df_1)","metadata":{"execution":{"iopub.status.busy":"2024-11-14T07:31:53.005389Z","iopub.execute_input":"2024-11-14T07:31:53.00564Z","iopub.status.idle":"2024-11-14T07:31:53.015267Z","shell.execute_reply.started":"2024-11-14T07:31:53.005611Z","shell.execute_reply":"2024-11-14T07:31:53.014563Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# def find_high_correlation_columns(df: pl.DataFrame, threshold: float):\n#     \"\"\"\n#     Finds pairs of columns in the DataFrame with correlation higher than the given threshold.\n\n#     Parameters:\n#     - df (pl.DataFrame): The input Polars DataFrame with numerical features.\n#     - threshold (float): The correlation threshold.\n\n#     Returns:\n#     - List of tuples with column pairs and their correlation values.\n#     \"\"\"\n#     # Calculate the correlation matrix\n#     corr_df = df.corr()\n\n#     # Convert the correlation matrix to a Pandas DataFrame for easy manipulation\n#     #corr_df = corr_matrix.to_pandas()\n\n#     # Find column pairs with correlation above the threshold\n#     high_corr_pairs = []\n#     for i in range(len(corr_df.columns)):\n#         for j in range(i + 1, len(corr_df.columns)):\n#             corr_value = corr_df.iloc[i, j]\n#             if abs(corr_value) > threshold:\n#                 col_pair = (corr_df.columns[i], corr_df.columns[j], corr_value)\n#                 high_corr_pairs.append(col_pair)\n\n#     return high_corr_pairs","metadata":{"execution":{"iopub.status.busy":"2024-11-14T07:31:53.016316Z","iopub.execute_input":"2024-11-14T07:31:53.016599Z","iopub.status.idle":"2024-11-14T07:31:53.027235Z","shell.execute_reply.started":"2024-11-14T07:31:53.016568Z","shell.execute_reply":"2024-11-14T07:31:53.026375Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# def find_high_correlation_columns_polars(df: pl.DataFrame, threshold: float):\n#     \"\"\"\n#     Finds pairs of columns in the Polars DataFrame with correlation higher than the given threshold.\n\n#     Parameters:\n#     - df (pl.DataFrame): The input Polars DataFrame with numerical features.\n#     - threshold (float): The correlation threshold.\n\n#     Returns:\n#     - List of tuples with column pairs and their correlation values.\n#     \"\"\"\n#     # Calculate the correlation matrix in Polars\n#     corr_matrix = df.corr()\n\n#     # Initialize a list to store column pairs with high correlation\n#     high_corr_pairs = []\n\n#     # Get the column names of the correlation matrix\n#     columns = corr_matrix.columns\n\n#     # Loop through each pair of columns in the upper triangle of the correlation matrix\n#     for i in range(len(columns)):\n#         for j in range(i + 1, len(columns)):\n#             # Access correlation value without converting to another format\n#             corr_value = corr_matrix[columns[i]][j].item()\n#             if abs(corr_value) > threshold:\n#                 col_pair = (columns[i], columns[j], corr_value)\n#                 high_corr_pairs.append(col_pair)\n\n#     return high_corr_pairs\n\ndef find_high_pairwise_correlation(df: pl.DataFrame, threshold: float):\n    \"\"\"\n    Finds pairs of columns in the Polars DataFrame with pairwise correlation higher than the given threshold.\n\n    Parameters:\n    - df (pl.DataFrame): The input Polars DataFrame with numerical features.\n    - threshold (float): The correlation threshold.\n\n    Returns:\n    - List of tuples with column pairs and their correlation values.\n    \"\"\"\n    # Initialize a list to store column pairs with high correlation\n    high_corr_pairs = []\n    columns_to_drop=[]\n\n    # Get the column names\n    columns = df.columns\n\n    # Loop through each pair of columns\n    for i in tqdm(range(len(columns)), desc=\"Calculating correlations\"):\n        for j in range(i + 1, len(columns)):\n            # Compute the correlation between the two columns\n            corr_value = df.select(pl.corr(columns[i], columns[j])).item()\n            \n            # Check if the correlation is above the threshold\n            if abs(corr_value) > threshold:\n                col_pair = (columns[i], columns[j], corr_value)\n                high_corr_pairs.append(col_pair)\n                columns_to_drop.append(columns[j])\n            \n            #del corr_value\n            #gc.collect()\n\n    return high_corr_pairs,columns_to_drop","metadata":{"execution":{"iopub.status.busy":"2024-11-14T07:31:53.028353Z","iopub.execute_input":"2024-11-14T07:31:53.028995Z","iopub.status.idle":"2024-11-14T07:31:53.042644Z","shell.execute_reply.started":"2024-11-14T07:31:53.028937Z","shell.execute_reply":"2024-11-14T07:31:53.041901Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# high_corr_columns = find_high_correlation_columns(df_0, threshold=0.8)\n# print(\"Highly correlated columns:\", high_corr_columns)","metadata":{"execution":{"iopub.status.busy":"2024-11-14T07:31:53.043791Z","iopub.execute_input":"2024-11-14T07:31:53.044119Z","iopub.status.idle":"2024-11-14T07:31:53.056128Z","shell.execute_reply.started":"2024-11-14T07:31:53.044088Z","shell.execute_reply":"2024-11-14T07:31:53.055243Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# columns=df_0.columns\n# print(columns[10])\n# print(df_0.select(pl.corr(columns[10],columns[11])).item())","metadata":{"execution":{"iopub.status.busy":"2024-11-14T07:31:53.057138Z","iopub.execute_input":"2024-11-14T07:31:53.057406Z","iopub.status.idle":"2024-11-14T07:31:53.075257Z","shell.execute_reply.started":"2024-11-14T07:31:53.057376Z","shell.execute_reply":"2024-11-14T07:31:53.074485Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# high_corr_columns,columns_to_drop = find_high_pairwise_correlation(df, threshold=0.9)\n# print(\"Highly correlated columns:\", high_corr_columns)\n# print(\"Columns that needs to be dropped:\",columns_to_drop)","metadata":{"execution":{"iopub.status.busy":"2024-11-14T07:31:53.076296Z","iopub.execute_input":"2024-11-14T07:31:53.077064Z","iopub.status.idle":"2024-11-14T07:31:53.085465Z","shell.execute_reply.started":"2024-11-14T07:31:53.077021Z","shell.execute_reply":"2024-11-14T07:31:53.084694Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"columns_to_drop=['feature_02', 'feature_03', 'feature_03', 'feature_67', 'feature_70', 'feature_69', 'feature_72', 'feature_17', 'feature_31', 'feature_34', 'feature_35', 'feature_35', 'feature_60', 'feature_74', 'feature_77', 'feature_78', 'feature_78', 'feature_76', 'feature_78']","metadata":{"execution":{"iopub.status.busy":"2024-11-14T07:31:53.086413Z","iopub.execute_input":"2024-11-14T07:31:53.086694Z","iopub.status.idle":"2024-11-14T07:31:53.096765Z","shell.execute_reply.started":"2024-11-14T07:31:53.086664Z","shell.execute_reply":"2024-11-14T07:31:53.096005Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data=[]\n# partition_ids=[x for x in range(10)]\n# responders=['responder_'+str(x) for x in range(9)]\npartition_ids=[7,8,9]\nresponders=['responder_6']\nyunbase=Yunbase()\nfor i in partition_ids:\n    train=pl.read_parquet(f\"/kaggle/input/jane-street-real-time-market-data-forecasting/train.parquet/partition_id={i}/part-0.parquet\")\n    train=train.to_pandas()\n    train['sin_time_id']=np.sin(2*np.pi*train['time_id']/967)\n    train['cos_time_id']=np.cos(2*np.pi*train['time_id']/967)\n    train['sin_time_id_halfday']=np.sin(2*np.pi*train['time_id']/483)\n    train['cos_time_id_halfday']=np.cos(2*np.pi*train['time_id']/483)\n    #train=train.fillna(method='ffill')\n    #train=train.fillna(0)\n    train=yunbase.reduce_mem_usage(train,float16_as32=False)\n    data.append(train)\ntrain=pd.concat(data)\nprint(f\"train.shape:{train.shape}\")\ndel data\ngc.collect()\nall_feature=['symbol_id','sin_time_id','cos_time_id','weight','sin_time_id_halfday','cos_time_id_halfday']+[f'feature_0{i}' if i<10 else f'feature_{i}' for i in range(79)]\nfinal_feature=[col for col in all_feature if col not in columns_to_drop]\n\ntrain=train[responders+final_feature]\ntrain.head()","metadata":{"execution":{"iopub.status.busy":"2024-11-14T07:32:03.363286Z","iopub.execute_input":"2024-11-14T07:32:03.36368Z","iopub.status.idle":"2024-11-14T07:33:07.381218Z","shell.execute_reply.started":"2024-11-14T07:32:03.36364Z","shell.execute_reply":"2024-11-14T07:33:07.380203Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#fill the initial null values with zero\n# train.fillna(0)","metadata":{"execution":{"iopub.status.busy":"2024-11-14T07:33:07.382684Z","iopub.execute_input":"2024-11-14T07:33:07.383025Z","iopub.status.idle":"2024-11-14T07:33:07.386704Z","shell.execute_reply.started":"2024-11-14T07:33:07.382988Z","shell.execute_reply":"2024-11-14T07:33:07.385823Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"lgb_params={\"boosting_type\": \"gbdt\",\"metric\": 'rmse',\n            'random_state': 2025,  \"max_depth\": 10,\"learning_rate\": 0.1,\n            \"n_estimators\": 120,\"colsample_bytree\": 0.6,\"colsample_bynode\": 0.6,\"verbose\": -1,\"reg_alpha\": 0.2,\n            \"reg_lambda\": 5,\"extra_trees\":True,'num_leaves':64,\"max_bin\":255,\n            'device':'gpu','gpu_use_dp':True,\n            }\n\ncat_params={'task_type':'GPU',\n           'random_state':2025,\n           'eval_metric'         : 'RMSE',\n           'bagging_temperature' : 0.50,\n           'iterations'          : 200,\n           'learning_rate'       : 0.1,\n           'max_depth'           : 12,\n           'l2_leaf_reg'         : 1.25,\n           'min_data_in_leaf'    : 24,\n           'random_strength'     : 0.25, \n           'verbose'             : 0,\n          }\nxgb_params={'random_state': 2025, 'n_estimators': 125, \n            'learning_rate': 0.1, 'max_depth': 10,\n            'reg_alpha': 0.08, 'reg_lambda': 0.8, \n            'subsample': 0.95, 'colsample_bytree': 0.6, \n            'min_child_weight': 3,\n            'tree_method':'gpu_hist',\n           }\nprint(\"lgb\")\nlgb=LGBMRegressor(**lgb_params)\nlgb.fit(train[final_feature].values,train['responder_6'].values)\nprint(\"cat\")\ncat=CatBoostRegressor(**cat_params)\ncat.fit(train[final_feature].values,train['responder_6'].values)\nprint(\"xgb\")\nxgb=XGBRegressor(**xgb_params)\nxgb.fit(train[final_feature].values,train['responder_6'].values)","metadata":{"execution":{"iopub.status.busy":"2024-11-14T07:33:07.387736Z","iopub.execute_input":"2024-11-14T07:33:07.388029Z","iopub.status.idle":"2024-11-14T07:43:02.744751Z","shell.execute_reply.started":"2024-11-14T07:33:07.387997Z","shell.execute_reply":"2024-11-14T07:43:02.74375Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# test=yunbase.reduce_mem_usage(test,float16_as32=False)","metadata":{"execution":{"iopub.status.busy":"2024-11-12T07:05:47.168783Z","iopub.execute_input":"2024-11-12T07:05:47.169492Z","iopub.status.idle":"2024-11-12T07:05:47.173705Z","shell.execute_reply.started":"2024-11-12T07:05:47.169444Z","shell.execute_reply":"2024-11-12T07:05:47.172702Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def predict(test,lags):\n    global lgb,cat,xgb\n    \n    predictions = test.select(\n        'row_id',\n        pl.lit(0.0).alias('responder_6'),\n    )\n    test=test.to_pandas()\n    test['sin_time_id']=np.sin(2*np.pi*test['time_id']/967)\n    test['cos_time_id']=np.cos(2*np.pi*test['time_id']/967)\n    test['sin_time_id_halfday']=np.sin(2*np.pi*test['time_id']/483)\n    test['cos_time_id_halfday']=np.cos(2*np.pi*test['time_id']/483)\n    test=test.fillna(-1)\n    test=test[final_feature]\n    eps=1e-10\n    test_preds=0.55*lgb.predict(test)+0.2*cat.predict(test)+0.25*xgb.predict(test)\n    test_preds=np.clip(test_preds,-5+eps,5-eps)\n    predictions = predictions.with_columns(pl.Series('responder_6', test_preds.ravel()))\n    return predictions\n\ninference_server = kaggle_evaluation.jane_street_inference_server.JSInferenceServer(predict)\n\nif os.getenv('KAGGLE_IS_COMPETITION_RERUN'):\n    inference_server.serve()\nelse:\n    inference_server.run_local_gateway(\n        (\n            '/kaggle/input/jane-street-real-time-market-data-forecasting/test.parquet',\n            '/kaggle/input/jane-street-real-time-market-data-forecasting/lags.parquet',\n        )\n    )","metadata":{"execution":{"iopub.status.busy":"2024-11-14T07:49:42.386451Z","iopub.execute_input":"2024-11-14T07:49:42.387534Z","iopub.status.idle":"2024-11-14T07:49:42.71216Z","shell.execute_reply.started":"2024-11-14T07:49:42.387483Z","shell.execute_reply":"2024-11-14T07:49:42.71054Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#submission = pl.read_parquet('/kaggle/working/submission.parquet')","metadata":{"execution":{"iopub.status.busy":"2024-11-14T07:49:47.997679Z","iopub.execute_input":"2024-11-14T07:49:47.998182Z","iopub.status.idle":"2024-11-14T07:49:48.002999Z","shell.execute_reply.started":"2024-11-14T07:49:47.998142Z","shell.execute_reply":"2024-11-14T07:49:48.002008Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#submission=submission.to_pandas()","metadata":{"execution":{"iopub.status.busy":"2024-11-14T07:49:52.311318Z","iopub.execute_input":"2024-11-14T07:49:52.311709Z","iopub.status.idle":"2024-11-14T07:49:52.317073Z","shell.execute_reply.started":"2024-11-14T07:49:52.311671Z","shell.execute_reply":"2024-11-14T07:49:52.316273Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#submission.head()","metadata":{"execution":{"iopub.status.busy":"2024-11-14T07:49:56.419495Z","iopub.execute_input":"2024-11-14T07:49:56.4199Z","iopub.status.idle":"2024-11-14T07:49:56.43287Z","shell.execute_reply.started":"2024-11-14T07:49:56.419861Z","shell.execute_reply":"2024-11-14T07:49:56.431738Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#submission.to_csv('submission(safa)_ver1.csv', index=False)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install pytorch_forecasting","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-18T08:20:39.225499Z","iopub.execute_input":"2024-11-18T08:20:39.2259Z","iopub.status.idle":"2024-11-18T08:20:40.619719Z","shell.execute_reply.started":"2024-11-18T08:20:39.225861Z","shell.execute_reply":"2024-11-18T08:20:40.618736Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[{"name":"stdout","text":"ERROR: unknown command \"installpytorch_forecasting\"\n","output_type":"stream"}],"execution_count":7},{"cell_type":"code","source":"!pip install pytorch-forecasting pytorch-lightning torch torchvision\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-18T08:21:20.494507Z","iopub.execute_input":"2024-11-18T08:21:20.494935Z","iopub.status.idle":"2024-11-18T08:21:35.172537Z","shell.execute_reply.started":"2024-11-18T08:21:20.494893Z","shell.execute_reply":"2024-11-18T08:21:35.171541Z"}},"outputs":[{"name":"stdout","text":"Collecting pytorch-forecasting\n  Downloading pytorch_forecasting-1.1.1-py3-none-any.whl.metadata (13 kB)\nRequirement already satisfied: pytorch-lightning in /opt/conda/lib/python3.10/site-packages (2.4.0)\nRequirement already satisfied: torch in /opt/conda/lib/python3.10/site-packages (2.4.0)\nRequirement already satisfied: torchvision in /opt/conda/lib/python3.10/site-packages (0.19.0)\nRequirement already satisfied: numpy<2.0.0 in /opt/conda/lib/python3.10/site-packages (from pytorch-forecasting) (1.26.4)\nCollecting lightning<3.0.0,>=2.0.0 (from pytorch-forecasting)\n  Downloading lightning-2.4.0-py3-none-any.whl.metadata (38 kB)\nRequirement already satisfied: scipy<2.0,>=1.8 in /opt/conda/lib/python3.10/site-packages (from pytorch-forecasting) (1.14.1)\nRequirement already satisfied: pandas<3.0.0,>=1.3.0 in /opt/conda/lib/python3.10/site-packages (from pytorch-forecasting) (2.2.2)\nRequirement already satisfied: scikit-learn<2.0,>=1.2 in /opt/conda/lib/python3.10/site-packages (from pytorch-forecasting) (1.2.2)\nRequirement already satisfied: tqdm>=4.57.0 in /opt/conda/lib/python3.10/site-packages (from pytorch-lightning) (4.66.4)\nRequirement already satisfied: PyYAML>=5.4 in /opt/conda/lib/python3.10/site-packages (from pytorch-lightning) (6.0.2)\nRequirement already satisfied: fsspec>=2022.5.0 in /opt/conda/lib/python3.10/site-packages (from fsspec[http]>=2022.5.0->pytorch-lightning) (2024.6.1)\nRequirement already satisfied: torchmetrics>=0.7.0 in /opt/conda/lib/python3.10/site-packages (from pytorch-lightning) (1.4.2)\nRequirement already satisfied: packaging>=20.0 in /opt/conda/lib/python3.10/site-packages (from pytorch-lightning) (21.3)\nRequirement already satisfied: typing-extensions>=4.4.0 in /opt/conda/lib/python3.10/site-packages (from pytorch-lightning) (4.12.2)\nRequirement already satisfied: lightning-utilities>=0.10.0 in /opt/conda/lib/python3.10/site-packages (from pytorch-lightning) (0.11.7)\nRequirement already satisfied: filelock in /opt/conda/lib/python3.10/site-packages (from torch) (3.15.1)\nRequirement already satisfied: sympy in /opt/conda/lib/python3.10/site-packages (from torch) (1.13.3)\nRequirement already satisfied: networkx in /opt/conda/lib/python3.10/site-packages (from torch) (3.3)\nRequirement already satisfied: jinja2 in /opt/conda/lib/python3.10/site-packages (from torch) (3.1.4)\nRequirement already satisfied: pillow!=8.3.*,>=5.3.0 in /opt/conda/lib/python3.10/site-packages (from torchvision) (10.3.0)\nRequirement already satisfied: aiohttp!=4.0.0a0,!=4.0.0a1 in /opt/conda/lib/python3.10/site-packages (from fsspec[http]>=2022.5.0->pytorch-lightning) (3.9.5)\nRequirement already satisfied: setuptools in /opt/conda/lib/python3.10/site-packages (from lightning-utilities>=0.10.0->pytorch-lightning) (70.0.0)\nRequirement already satisfied: pyparsing!=3.0.5,>=2.0.2 in /opt/conda/lib/python3.10/site-packages (from packaging>=20.0->pytorch-lightning) (3.1.2)\nRequirement already satisfied: python-dateutil>=2.8.2 in /opt/conda/lib/python3.10/site-packages (from pandas<3.0.0,>=1.3.0->pytorch-forecasting) (2.9.0.post0)\nRequirement already satisfied: pytz>=2020.1 in /opt/conda/lib/python3.10/site-packages (from pandas<3.0.0,>=1.3.0->pytorch-forecasting) (2024.1)\nRequirement already satisfied: tzdata>=2022.7 in /opt/conda/lib/python3.10/site-packages (from pandas<3.0.0,>=1.3.0->pytorch-forecasting) (2024.1)\nRequirement already satisfied: joblib>=1.1.1 in /opt/conda/lib/python3.10/site-packages (from scikit-learn<2.0,>=1.2->pytorch-forecasting) (1.4.2)\nRequirement already satisfied: threadpoolctl>=2.0.0 in /opt/conda/lib/python3.10/site-packages (from scikit-learn<2.0,>=1.2->pytorch-forecasting) (3.5.0)\nRequirement already satisfied: MarkupSafe>=2.0 in /opt/conda/lib/python3.10/site-packages (from jinja2->torch) (2.1.5)\nRequirement already satisfied: mpmath<1.4,>=1.1.0 in /opt/conda/lib/python3.10/site-packages (from sympy->torch) (1.3.0)\nRequirement already satisfied: aiosignal>=1.1.2 in /opt/conda/lib/python3.10/site-packages (from aiohttp!=4.0.0a0,!=4.0.0a1->fsspec[http]>=2022.5.0->pytorch-lightning) (1.3.1)\nRequirement already satisfied: attrs>=17.3.0 in /opt/conda/lib/python3.10/site-packages (from aiohttp!=4.0.0a0,!=4.0.0a1->fsspec[http]>=2022.5.0->pytorch-lightning) (23.2.0)\nRequirement already satisfied: frozenlist>=1.1.1 in /opt/conda/lib/python3.10/site-packages (from aiohttp!=4.0.0a0,!=4.0.0a1->fsspec[http]>=2022.5.0->pytorch-lightning) (1.4.1)\nRequirement already satisfied: multidict<7.0,>=4.5 in /opt/conda/lib/python3.10/site-packages (from aiohttp!=4.0.0a0,!=4.0.0a1->fsspec[http]>=2022.5.0->pytorch-lightning) (6.0.5)\nRequirement already satisfied: yarl<2.0,>=1.0 in /opt/conda/lib/python3.10/site-packages (from aiohttp!=4.0.0a0,!=4.0.0a1->fsspec[http]>=2022.5.0->pytorch-lightning) (1.9.4)\nRequirement already satisfied: async-timeout<5.0,>=4.0 in /opt/conda/lib/python3.10/site-packages (from aiohttp!=4.0.0a0,!=4.0.0a1->fsspec[http]>=2022.5.0->pytorch-lightning) (4.0.3)\nRequirement already satisfied: six>=1.5 in /opt/conda/lib/python3.10/site-packages (from python-dateutil>=2.8.2->pandas<3.0.0,>=1.3.0->pytorch-forecasting) (1.16.0)\nRequirement already satisfied: idna>=2.0 in /opt/conda/lib/python3.10/site-packages (from yarl<2.0,>=1.0->aiohttp!=4.0.0a0,!=4.0.0a1->fsspec[http]>=2022.5.0->pytorch-lightning) (3.7)\nDownloading pytorch_forecasting-1.1.1-py3-none-any.whl (177 kB)\n\u001b[2K   \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m177.6/177.6 kB\u001b[0m \u001b[31m14.0 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n\u001b[?25hDownloading lightning-2.4.0-py3-none-any.whl (810 kB)\n\u001b[2K   \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m811.0/811.0 kB\u001b[0m \u001b[31m44.1 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n\u001b[?25hInstalling collected packages: lightning, pytorch-forecasting\nSuccessfully installed lightning-2.4.0 pytorch-forecasting-1.1.1\n","output_type":"stream"}],"execution_count":8},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport torch\nfrom pytorch_forecasting import TimeSeriesDataSet, TemporalFusionTransformer, Baseline\nfrom pytorch_forecasting.data import GroupNormalizer, MultiNormalizer\nfrom pytorch_forecasting.metrics import QuantileLoss\nfrom pytorch_lightning import Trainer\nfrom pytorch_lightning.callbacks import EarlyStopping\nimport gc","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-18T08:21:53.496331Z","iopub.execute_input":"2024-11-18T08:21:53.496709Z","iopub.status.idle":"2024-11-18T08:21:53.507784Z","shell.execute_reply.started":"2024-11-18T08:21:53.496675Z","shell.execute_reply":"2024-11-18T08:21:53.506937Z"}},"outputs":[],"execution_count":11},{"cell_type":"code","source":"pd.read_parquet(f\"/kaggle/input/jane-street-real-time-market-data-forecasting/train.parquet/partition_id=0/part-0.parquet\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-18T08:32:31.046029Z","iopub.execute_input":"2024-11-18T08:32:31.046904Z","iopub.status.idle":"2024-11-18T08:32:34.26869Z","shell.execute_reply.started":"2024-11-18T08:32:31.046858Z","shell.execute_reply":"2024-11-18T08:32:34.267588Z"}},"outputs":[{"execution_count":14,"output_type":"execute_result","data":{"text/plain":"         date_id  time_id  symbol_id    weight  feature_00  feature_01  \\\n0              0        0          1  3.889038         NaN         NaN   \n1              0        0          7  1.370613         NaN         NaN   \n2              0        0          9  2.285698         NaN         NaN   \n3              0        0         10  0.690606         NaN         NaN   \n4              0        0         14  0.440570         NaN         NaN   \n...          ...      ...        ...       ...         ...         ...   \n1944205      169      848         19  3.438631         NaN         NaN   \n1944206      169      848         30  0.768528         NaN         NaN   \n1944207      169      848         33  1.354696         NaN         NaN   \n1944208      169      848         34  1.021797         NaN         NaN   \n1944209      169      848         38  1.570022         NaN         NaN   \n\n         feature_02  feature_03  feature_04  feature_05  ...  feature_78  \\\n0               NaN         NaN         NaN    0.851033  ...   -0.281498   \n1               NaN         NaN         NaN    0.676961  ...   -0.302441   \n2               NaN         NaN         NaN    1.056285  ...   -0.096792   \n3               NaN         NaN         NaN    1.139366  ...   -0.296244   \n4               NaN         NaN         NaN    0.955200  ...    3.418133   \n...             ...         ...         ...         ...  ...         ...   \n1944205         NaN         NaN         NaN   -0.028087  ...   -0.166964   \n1944206         NaN         NaN         NaN   -0.022584  ...   -0.352810   \n1944207         NaN         NaN         NaN   -0.024804  ...   -0.239716   \n1944208         NaN         NaN         NaN   -0.016138  ...   -0.442859   \n1944209         NaN         NaN         NaN   -0.017634  ...   -0.174461   \n\n         responder_0  responder_1  responder_2  responder_3  responder_4  \\\n0           0.738489    -0.069556     1.380875     2.005353     0.186018   \n1           2.965889     1.190077    -0.523998     3.849921     2.626981   \n2          -0.864488    -0.280303    -0.326697     0.375781     1.271291   \n3           0.408499     0.223992     2.294888     1.097444     1.225872   \n4          -0.373387    -0.502764    -0.348021    -3.928148    -1.591366   \n...              ...          ...          ...          ...          ...   \n1944205     0.983339    -0.669860     0.272615    -3.676842    -1.221126   \n1944206     0.992615     0.961595     1.089402     0.796034     0.488380   \n1944207     1.701618     0.757672    -5.000000    -3.174266    -1.110790   \n1944208    -2.036891    -0.064228     1.919665     1.827681     0.872019   \n1944209     0.323230     0.018376    -3.457667    -0.305218    -0.181438   \n\n         responder_5  responder_6  responder_7  responder_8  \n0           1.218368     0.775981     0.346999     0.095504  \n1           5.000000     0.703665     0.216683     0.778639  \n2           0.099793     2.109352     0.670881     0.772828  \n3           1.225376     1.114137     0.775199    -1.379516  \n4          -5.000000    -3.572820    -1.089123    -5.000000  \n...              ...          ...          ...          ...  \n1944205     1.070584     0.465345     0.207483     0.874975  \n1944206     1.846634    -0.088542    -0.008324    -0.153451  \n1944207    -3.349107    -0.407801    -0.185842    -0.931004  \n1944208     3.248694     0.254584     0.090288     0.434726  \n1944209    -0.791345     0.347400     0.241875     0.987731  \n\n[1944210 rows x 92 columns]","text/html":"<div>\n<style scoped>\n    .dataframe tbody tr th:only-of-type {\n        vertical-align: middle;\n    }\n\n    .dataframe tbody tr th {\n        vertical-align: top;\n    }\n\n    .dataframe thead th {\n        text-align: right;\n    }\n</style>\n<table border=\"1\" class=\"dataframe\">\n  <thead>\n    <tr style=\"text-align: right;\">\n      <th></th>\n      <th>date_id</th>\n      <th>time_id</th>\n      <th>symbol_id</th>\n      <th>weight</th>\n      <th>feature_00</th>\n      <th>feature_01</th>\n      <th>feature_02</th>\n      <th>feature_03</th>\n      <th>feature_04</th>\n      <th>feature_05</th>\n      <th>...</th>\n      <th>feature_78</th>\n      <th>responder_0</th>\n      <th>responder_1</th>\n      <th>responder_2</th>\n      <th>responder_3</th>\n      <th>responder_4</th>\n      <th>responder_5</th>\n      <th>responder_6</th>\n      <th>responder_7</th>\n      <th>responder_8</th>\n    </tr>\n  </thead>\n  <tbody>\n    <tr>\n      <th>0</th>\n      <td>0</td>\n      <td>0</td>\n      <td>1</td>\n      <td>3.889038</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>0.851033</td>\n      <td>...</td>\n      <td>-0.281498</td>\n      <td>0.738489</td>\n      <td>-0.069556</td>\n      <td>1.380875</td>\n      <td>2.005353</td>\n      <td>0.186018</td>\n      <td>1.218368</td>\n      <td>0.775981</td>\n      <td>0.346999</td>\n      <td>0.095504</td>\n    </tr>\n    <tr>\n      <th>1</th>\n      <td>0</td>\n      <td>0</td>\n      <td>7</td>\n      <td>1.370613</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>0.676961</td>\n      <td>...</td>\n      <td>-0.302441</td>\n      <td>2.965889</td>\n      <td>1.190077</td>\n      <td>-0.523998</td>\n      <td>3.849921</td>\n      <td>2.626981</td>\n      <td>5.000000</td>\n      <td>0.703665</td>\n      <td>0.216683</td>\n      <td>0.778639</td>\n    </tr>\n    <tr>\n      <th>2</th>\n      <td>0</td>\n      <td>0</td>\n      <td>9</td>\n      <td>2.285698</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>1.056285</td>\n      <td>...</td>\n      <td>-0.096792</td>\n      <td>-0.864488</td>\n      <td>-0.280303</td>\n      <td>-0.326697</td>\n      <td>0.375781</td>\n      <td>1.271291</td>\n      <td>0.099793</td>\n      <td>2.109352</td>\n      <td>0.670881</td>\n      <td>0.772828</td>\n    </tr>\n    <tr>\n      <th>3</th>\n      <td>0</td>\n      <td>0</td>\n      <td>10</td>\n      <td>0.690606</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>1.139366</td>\n      <td>...</td>\n      <td>-0.296244</td>\n      <td>0.408499</td>\n      <td>0.223992</td>\n      <td>2.294888</td>\n      <td>1.097444</td>\n      <td>1.225872</td>\n      <td>1.225376</td>\n      <td>1.114137</td>\n      <td>0.775199</td>\n      <td>-1.379516</td>\n    </tr>\n    <tr>\n      <th>4</th>\n      <td>0</td>\n      <td>0</td>\n      <td>14</td>\n      <td>0.440570</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>0.955200</td>\n      <td>...</td>\n      <td>3.418133</td>\n      <td>-0.373387</td>\n      <td>-0.502764</td>\n      <td>-0.348021</td>\n      <td>-3.928148</td>\n      <td>-1.591366</td>\n      <td>-5.000000</td>\n      <td>-3.572820</td>\n      <td>-1.089123</td>\n      <td>-5.000000</td>\n    </tr>\n    <tr>\n      <th>...</th>\n      <td>...</td>\n      <td>...</td>\n      <td>...</td>\n      <td>...</td>\n      <td>...</td>\n      <td>...</td>\n      <td>...</td>\n      <td>...</td>\n      <td>...</td>\n      <td>...</td>\n      <td>...</td>\n      <td>...</td>\n      <td>...</td>\n      <td>...</td>\n      <td>...</td>\n      <td>...</td>\n      <td>...</td>\n      <td>...</td>\n      <td>...</td>\n      <td>...</td>\n      <td>...</td>\n    </tr>\n    <tr>\n      <th>1944205</th>\n      <td>169</td>\n      <td>848</td>\n      <td>19</td>\n      <td>3.438631</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>-0.028087</td>\n      <td>...</td>\n      <td>-0.166964</td>\n      <td>0.983339</td>\n      <td>-0.669860</td>\n      <td>0.272615</td>\n      <td>-3.676842</td>\n      <td>-1.221126</td>\n      <td>1.070584</td>\n      <td>0.465345</td>\n      <td>0.207483</td>\n      <td>0.874975</td>\n    </tr>\n    <tr>\n      <th>1944206</th>\n      <td>169</td>\n      <td>848</td>\n      <td>30</td>\n      <td>0.768528</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>-0.022584</td>\n      <td>...</td>\n      <td>-0.352810</td>\n      <td>0.992615</td>\n      <td>0.961595</td>\n      <td>1.089402</td>\n      <td>0.796034</td>\n      <td>0.488380</td>\n      <td>1.846634</td>\n      <td>-0.088542</td>\n      <td>-0.008324</td>\n      <td>-0.153451</td>\n    </tr>\n    <tr>\n      <th>1944207</th>\n      <td>169</td>\n      <td>848</td>\n      <td>33</td>\n      <td>1.354696</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>-0.024804</td>\n      <td>...</td>\n      <td>-0.239716</td>\n      <td>1.701618</td>\n      <td>0.757672</td>\n      <td>-5.000000</td>\n      <td>-3.174266</td>\n      <td>-1.110790</td>\n      <td>-3.349107</td>\n      <td>-0.407801</td>\n      <td>-0.185842</td>\n      <td>-0.931004</td>\n    </tr>\n    <tr>\n      <th>1944208</th>\n      <td>169</td>\n      <td>848</td>\n      <td>34</td>\n      <td>1.021797</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>-0.016138</td>\n      <td>...</td>\n      <td>-0.442859</td>\n      <td>-2.036891</td>\n      <td>-0.064228</td>\n      <td>1.919665</td>\n      <td>1.827681</td>\n      <td>0.872019</td>\n      <td>3.248694</td>\n      <td>0.254584</td>\n      <td>0.090288</td>\n      <td>0.434726</td>\n    </tr>\n    <tr>\n      <th>1944209</th>\n      <td>169</td>\n      <td>848</td>\n      <td>38</td>\n      <td>1.570022</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>NaN</td>\n      <td>-0.017634</td>\n      <td>...</td>\n      <td>-0.174461</td>\n      <td>0.323230</td>\n      <td>0.018376</td>\n      <td>-3.457667</td>\n      <td>-0.305218</td>\n      <td>-0.181438</td>\n      <td>-0.791345</td>\n      <td>0.347400</td>\n      <td>0.241875</td>\n      <td>0.987731</td>\n    </tr>\n  </tbody>\n</table>\n<p>1944210 rows × 92 columns</p>\n</div>"},"metadata":{}}],"execution_count":14},{"cell_type":"code","source":"# Load data\ndata = []\npartition_ids = [6, 7, 8, 9]  # Use desired partitions\nfor i in partition_ids:\n    train = pd.read_parquet(f\"/kaggle/input/jane-street-real-time-market-data-forecasting/train.parquet/partition_id={i}/part-0.parquet\")\n    train['sin_time_id'] = np.sin(2 * np.pi * train['time_id'] / 967)\n    train['cos_time_id'] = np.cos(2 * np.pi * train['time_id'] / 967)\n    data.append(train)\n\ntrain = pd.concat(data)\ndel data\ngc.collect()\n\n# Features\nall_feature = ['symbol_id', 'sin_time_id', 'cos_time_id'] + [f'feature_0{i}' if i<10 else f'feature_{i}' for i in range(79)]\nfinal_feature = all_feature\n\nprint(final_feature)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-18T08:35:32.096937Z","iopub.execute_input":"2024-11-18T08:35:32.097341Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train = train[final_feature + ['responder_6', 'date_id', 'time_id']]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-18T08:33:04.061772Z","iopub.execute_input":"2024-11-18T08:33:04.062715Z","iopub.status.idle":"2024-11-18T08:33:04.146662Z","shell.execute_reply.started":"2024-11-18T08:33:04.062671Z","shell.execute_reply":"2024-11-18T08:33:04.145339Z"}},"outputs":[{"traceback":["\u001b[0;31m---------------------------------------------------------------------------\u001b[0m","\u001b[0;31mKeyError\u001b[0m                                  Traceback (most recent call last)","Cell \u001b[0;32mIn[15], line 1\u001b[0m\n\u001b[0;32m----> 1\u001b[0m train \u001b[38;5;241m=\u001b[39m \u001b[43mtrain\u001b[49m\u001b[43m[\u001b[49m\u001b[43mfinal_feature\u001b[49m\u001b[43m \u001b[49m\u001b[38;5;241;43m+\u001b[39;49m\u001b[43m \u001b[49m\u001b[43m[\u001b[49m\u001b[38;5;124;43m'\u001b[39;49m\u001b[38;5;124;43mresponder_6\u001b[39;49m\u001b[38;5;124;43m'\u001b[39;49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[38;5;124;43m'\u001b[39;49m\u001b[38;5;124;43mdate_id\u001b[39;49m\u001b[38;5;124;43m'\u001b[39;49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[38;5;124;43m'\u001b[39;49m\u001b[38;5;124;43mtime_id\u001b[39;49m\u001b[38;5;124;43m'\u001b[39;49m\u001b[43m]\u001b[49m\u001b[43m]\u001b[49m\n","File \u001b[0;32m/opt/conda/lib/python3.10/site-packages/pandas/core/frame.py:4108\u001b[0m, in \u001b[0;36mDataFrame.__getitem__\u001b[0;34m(self, key)\u001b[0m\n\u001b[1;32m   4106\u001b[0m     \u001b[38;5;28;01mif\u001b[39;00m is_iterator(key):\n\u001b[1;32m   4107\u001b[0m         key \u001b[38;5;241m=\u001b[39m \u001b[38;5;28mlist\u001b[39m(key)\n\u001b[0;32m-> 4108\u001b[0m     indexer \u001b[38;5;241m=\u001b[39m \u001b[38;5;28;43mself\u001b[39;49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43mcolumns\u001b[49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43m_get_indexer_strict\u001b[49m\u001b[43m(\u001b[49m\u001b[43mkey\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[38;5;124;43m\"\u001b[39;49m\u001b[38;5;124;43mcolumns\u001b[39;49m\u001b[38;5;124;43m\"\u001b[39;49m\u001b[43m)\u001b[49m[\u001b[38;5;241m1\u001b[39m]\n\u001b[1;32m   4110\u001b[0m \u001b[38;5;66;03m# take() does not accept boolean indexers\u001b[39;00m\n\u001b[1;32m   4111\u001b[0m \u001b[38;5;28;01mif\u001b[39;00m \u001b[38;5;28mgetattr\u001b[39m(indexer, \u001b[38;5;124m\"\u001b[39m\u001b[38;5;124mdtype\u001b[39m\u001b[38;5;124m\"\u001b[39m, \u001b[38;5;28;01mNone\u001b[39;00m) \u001b[38;5;241m==\u001b[39m \u001b[38;5;28mbool\u001b[39m:\n","File \u001b[0;32m/opt/conda/lib/python3.10/site-packages/pandas/core/indexes/base.py:6200\u001b[0m, in \u001b[0;36mIndex._get_indexer_strict\u001b[0;34m(self, key, axis_name)\u001b[0m\n\u001b[1;32m   6197\u001b[0m \u001b[38;5;28;01melse\u001b[39;00m:\n\u001b[1;32m   6198\u001b[0m     keyarr, indexer, new_indexer \u001b[38;5;241m=\u001b[39m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39m_reindex_non_unique(keyarr)\n\u001b[0;32m-> 6200\u001b[0m \u001b[38;5;28;43mself\u001b[39;49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43m_raise_if_missing\u001b[49m\u001b[43m(\u001b[49m\u001b[43mkeyarr\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[43mindexer\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[43maxis_name\u001b[49m\u001b[43m)\u001b[49m\n\u001b[1;32m   6202\u001b[0m keyarr \u001b[38;5;241m=\u001b[39m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39mtake(indexer)\n\u001b[1;32m   6203\u001b[0m \u001b[38;5;28;01mif\u001b[39;00m \u001b[38;5;28misinstance\u001b[39m(key, Index):\n\u001b[1;32m   6204\u001b[0m     \u001b[38;5;66;03m# GH 42790 - Preserve name from an Index\u001b[39;00m\n","File \u001b[0;32m/opt/conda/lib/python3.10/site-packages/pandas/core/indexes/base.py:6252\u001b[0m, in \u001b[0;36mIndex._raise_if_missing\u001b[0;34m(self, key, indexer, axis_name)\u001b[0m\n\u001b[1;32m   6249\u001b[0m     \u001b[38;5;28;01mraise\u001b[39;00m \u001b[38;5;167;01mKeyError\u001b[39;00m(\u001b[38;5;124mf\u001b[39m\u001b[38;5;124m\"\u001b[39m\u001b[38;5;124mNone of [\u001b[39m\u001b[38;5;132;01m{\u001b[39;00mkey\u001b[38;5;132;01m}\u001b[39;00m\u001b[38;5;124m] are in the [\u001b[39m\u001b[38;5;132;01m{\u001b[39;00maxis_name\u001b[38;5;132;01m}\u001b[39;00m\u001b[38;5;124m]\u001b[39m\u001b[38;5;124m\"\u001b[39m)\n\u001b[1;32m   6251\u001b[0m not_found \u001b[38;5;241m=\u001b[39m \u001b[38;5;28mlist\u001b[39m(ensure_index(key)[missing_mask\u001b[38;5;241m.\u001b[39mnonzero()[\u001b[38;5;241m0\u001b[39m]]\u001b[38;5;241m.\u001b[39munique())\n\u001b[0;32m-> 6252\u001b[0m \u001b[38;5;28;01mraise\u001b[39;00m \u001b[38;5;167;01mKeyError\u001b[39;00m(\u001b[38;5;124mf\u001b[39m\u001b[38;5;124m\"\u001b[39m\u001b[38;5;132;01m{\u001b[39;00mnot_found\u001b[38;5;132;01m}\u001b[39;00m\u001b[38;5;124m not in index\u001b[39m\u001b[38;5;124m\"\u001b[39m)\n","\u001b[0;31mKeyError\u001b[0m: \"['feature_0', 'feature_1', 'feature_2', 'feature_3', 'feature_4', 'feature_5', 'feature_6', 'feature_7', 'feature_8', 'feature_9'] not in index\""],"ename":"KeyError","evalue":"\"['feature_0', 'feature_1', 'feature_2', 'feature_3', 'feature_4', 'feature_5', 'feature_6', 'feature_7', 'feature_8', 'feature_9'] not in index\"","output_type":"error"}],"execution_count":15},{"cell_type":"code","source":"\n\n# Prepare dataset for TFT\n\ntrain = train.rename(columns={\"date_id\": \"time_idx\", \"symbol_id\": \"group_id\"})  # Required column names for TFT\n\n# Define dataset\nmax_prediction_length = 1  # Predict next step\nmax_encoder_length = 30  # Use past 30 steps for prediction\n\n# Use MultiNormalizer to normalize both target and features\nmulti_normalizer = MultiNormalizer(\n    [GroupNormalizer(groups=[\"group_id\"], transformation=\"softplus\"),\n     GroupNormalizer(groups=[\"group_id\"], transformation=\"standard\", target=\"responder_6\")]\n)\n\ntraining = TimeSeriesDataSet(\n    train,\n    time_idx=\"time_idx\",\n    target=\"responder_6\",\n    group_ids=[\"group_id\"],\n    min_encoder_length=10,  # Allow partial history\n    max_encoder_length=max_encoder_length,\n    min_prediction_length=max_prediction_length,\n    max_prediction_length=max_prediction_length,\n    static_categoricals=[\"group_id\"],\n    time_varying_known_reals=[\"time_idx\", \"sin_time_id\", \"cos_time_id\"],\n    time_varying_unknown_reals=[col for col in final_feature if col not in [\"group_id\"]],\n    target_normalizer=multi_normalizer\n)\n\n# Create DataLoader\nbatch_size = 64  # Adjust based on your GPU memory\ntrain_dataloader = training.to_dataloader(train=True, batch_size=batch_size, num_workers=4)\n\n# Define TFT model\ntft = TemporalFusionTransformer.from_dataset(\n    training,\n    learning_rate=1e-3,\n    hidden_size=32,\n    attention_head_size=4,\n    dropout=0.1,\n    hidden_continuous_size=16,\n    output_size=7,  # Quantiles for QuantileLoss\n    loss=QuantileLoss(),\n    log_interval=10,\n    reduce_on_plateau_patience=4\n)\n\n# Training with early stopping\nearly_stop_callback = EarlyStopping(monitor=\"val_loss\", patience=5, mode=\"min\")\n\ntrainer = Trainer(\n    max_epochs=30,\n    gpus=1 if torch.cuda.is_available() else 0,\n    gradient_clip_val=0.1,\n    callbacks=[early_stop_callback],\n)\n\n# Train model\ntrainer.fit(\n    tft,\n    train_dataloaders=train_dataloader\n)\n\n# Model Interpretation\n# SHAP values for interpretability\ninterpret = tft.interpret_output(tft.predict(train_dataloader, mode=\"raw\"))\ntft.plot_interpretation(interpret)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-18T08:25:58.63348Z","iopub.execute_input":"2024-11-18T08:25:58.634421Z","iopub.status.idle":"2024-11-18T08:26:46.558174Z","shell.execute_reply.started":"2024-11-18T08:25:58.634359Z","shell.execute_reply":"2024-11-18T08:26:46.556423Z"}},"outputs":[{"traceback":["\u001b[0;31m---------------------------------------------------------------------------\u001b[0m","\u001b[0;31mKeyError\u001b[0m                                  Traceback (most recent call last)","Cell \u001b[0;32mIn[12], line 19\u001b[0m\n\u001b[1;32m     16\u001b[0m final_feature \u001b[38;5;241m=\u001b[39m all_feature\n\u001b[1;32m     18\u001b[0m \u001b[38;5;66;03m# Prepare dataset for TFT\u001b[39;00m\n\u001b[0;32m---> 19\u001b[0m train \u001b[38;5;241m=\u001b[39m \u001b[43mtrain\u001b[49m\u001b[43m[\u001b[49m\u001b[43mfinal_feature\u001b[49m\u001b[43m \u001b[49m\u001b[38;5;241;43m+\u001b[39;49m\u001b[43m \u001b[49m\u001b[43m[\u001b[49m\u001b[38;5;124;43m'\u001b[39;49m\u001b[38;5;124;43mresponder_6\u001b[39;49m\u001b[38;5;124;43m'\u001b[39;49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[38;5;124;43m'\u001b[39;49m\u001b[38;5;124;43mdate_id\u001b[39;49m\u001b[38;5;124;43m'\u001b[39;49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[38;5;124;43m'\u001b[39;49m\u001b[38;5;124;43mtime_id\u001b[39;49m\u001b[38;5;124;43m'\u001b[39;49m\u001b[43m]\u001b[49m\u001b[43m]\u001b[49m\n\u001b[1;32m     20\u001b[0m train \u001b[38;5;241m=\u001b[39m train\u001b[38;5;241m.\u001b[39mrename(columns\u001b[38;5;241m=\u001b[39m{\u001b[38;5;124m\"\u001b[39m\u001b[38;5;124mdate_id\u001b[39m\u001b[38;5;124m\"\u001b[39m: \u001b[38;5;124m\"\u001b[39m\u001b[38;5;124mtime_idx\u001b[39m\u001b[38;5;124m\"\u001b[39m, \u001b[38;5;124m\"\u001b[39m\u001b[38;5;124msymbol_id\u001b[39m\u001b[38;5;124m\"\u001b[39m: \u001b[38;5;124m\"\u001b[39m\u001b[38;5;124mgroup_id\u001b[39m\u001b[38;5;124m\"\u001b[39m})  \u001b[38;5;66;03m# Required column names for TFT\u001b[39;00m\n\u001b[1;32m     22\u001b[0m \u001b[38;5;66;03m# Define dataset\u001b[39;00m\n","File \u001b[0;32m/opt/conda/lib/python3.10/site-packages/pandas/core/frame.py:4108\u001b[0m, in \u001b[0;36mDataFrame.__getitem__\u001b[0;34m(self, key)\u001b[0m\n\u001b[1;32m   4106\u001b[0m     \u001b[38;5;28;01mif\u001b[39;00m is_iterator(key):\n\u001b[1;32m   4107\u001b[0m         key \u001b[38;5;241m=\u001b[39m \u001b[38;5;28mlist\u001b[39m(key)\n\u001b[0;32m-> 4108\u001b[0m     indexer \u001b[38;5;241m=\u001b[39m \u001b[38;5;28;43mself\u001b[39;49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43mcolumns\u001b[49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43m_get_indexer_strict\u001b[49m\u001b[43m(\u001b[49m\u001b[43mkey\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[38;5;124;43m\"\u001b[39;49m\u001b[38;5;124;43mcolumns\u001b[39;49m\u001b[38;5;124;43m\"\u001b[39;49m\u001b[43m)\u001b[49m[\u001b[38;5;241m1\u001b[39m]\n\u001b[1;32m   4110\u001b[0m \u001b[38;5;66;03m# take() does not accept boolean indexers\u001b[39;00m\n\u001b[1;32m   4111\u001b[0m \u001b[38;5;28;01mif\u001b[39;00m \u001b[38;5;28mgetattr\u001b[39m(indexer, \u001b[38;5;124m\"\u001b[39m\u001b[38;5;124mdtype\u001b[39m\u001b[38;5;124m\"\u001b[39m, \u001b[38;5;28;01mNone\u001b[39;00m) \u001b[38;5;241m==\u001b[39m \u001b[38;5;28mbool\u001b[39m:\n","File \u001b[0;32m/opt/conda/lib/python3.10/site-packages/pandas/core/indexes/base.py:6200\u001b[0m, in \u001b[0;36mIndex._get_indexer_strict\u001b[0;34m(self, key, axis_name)\u001b[0m\n\u001b[1;32m   6197\u001b[0m \u001b[38;5;28;01melse\u001b[39;00m:\n\u001b[1;32m   6198\u001b[0m     keyarr, indexer, new_indexer \u001b[38;5;241m=\u001b[39m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39m_reindex_non_unique(keyarr)\n\u001b[0;32m-> 6200\u001b[0m \u001b[38;5;28;43mself\u001b[39;49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43m_raise_if_missing\u001b[49m\u001b[43m(\u001b[49m\u001b[43mkeyarr\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[43mindexer\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[43maxis_name\u001b[49m\u001b[43m)\u001b[49m\n\u001b[1;32m   6202\u001b[0m keyarr \u001b[38;5;241m=\u001b[39m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39mtake(indexer)\n\u001b[1;32m   6203\u001b[0m \u001b[38;5;28;01mif\u001b[39;00m \u001b[38;5;28misinstance\u001b[39m(key, Index):\n\u001b[1;32m   6204\u001b[0m     \u001b[38;5;66;03m# GH 42790 - Preserve name from an Index\u001b[39;00m\n","File \u001b[0;32m/opt/conda/lib/python3.10/site-packages/pandas/core/indexes/base.py:6252\u001b[0m, in \u001b[0;36mIndex._raise_if_missing\u001b[0;34m(self, key, indexer, axis_name)\u001b[0m\n\u001b[1;32m   6249\u001b[0m     \u001b[38;5;28;01mraise\u001b[39;00m \u001b[38;5;167;01mKeyError\u001b[39;00m(\u001b[38;5;124mf\u001b[39m\u001b[38;5;124m\"\u001b[39m\u001b[38;5;124mNone of [\u001b[39m\u001b[38;5;132;01m{\u001b[39;00mkey\u001b[38;5;132;01m}\u001b[39;00m\u001b[38;5;124m] are in the [\u001b[39m\u001b[38;5;132;01m{\u001b[39;00maxis_name\u001b[38;5;132;01m}\u001b[39;00m\u001b[38;5;124m]\u001b[39m\u001b[38;5;124m\"\u001b[39m)\n\u001b[1;32m   6251\u001b[0m not_found \u001b[38;5;241m=\u001b[39m \u001b[38;5;28mlist\u001b[39m(ensure_index(key)[missing_mask\u001b[38;5;241m.\u001b[39mnonzero()[\u001b[38;5;241m0\u001b[39m]]\u001b[38;5;241m.\u001b[39munique())\n\u001b[0;32m-> 6252\u001b[0m \u001b[38;5;28;01mraise\u001b[39;00m \u001b[38;5;167;01mKeyError\u001b[39;00m(\u001b[38;5;124mf\u001b[39m\u001b[38;5;124m\"\u001b[39m\u001b[38;5;132;01m{\u001b[39;00mnot_found\u001b[38;5;132;01m}\u001b[39;00m\u001b[38;5;124m not in index\u001b[39m\u001b[38;5;124m\"\u001b[39m)\n","\u001b[0;31mKeyError\u001b[0m: \"['feature_0', 'feature_1', 'feature_2', 'feature_3', 'feature_4', 'feature_5', 'feature_6', 'feature_7', 'feature_8', 'feature_9'] not in index\""],"ename":"KeyError","evalue":"\"['feature_0', 'feature_1', 'feature_2', 'feature_3', 'feature_4', 'feature_5', 'feature_6', 'feature_7', 'feature_8', 'feature_9'] not in index\"","output_type":"error"}],"execution_count":12},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}