{"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":"none","dataSources":[{"sourceId":84493,"databundleVersionId":9871156,"sourceType":"competition"},{"sourceId":10167324,"sourceType":"datasetVersion","datasetId":6278812},{"sourceId":208703,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":177931,"modelId":200235}],"dockerImageVersionId":30805,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import polars as pl\nimport numpy as np\n\nfeatures = [\"feature_06\", \"feature_36\", \"feature_04\", \"feature_56\", \"feature_19\", \"feature_59\", \n           \"feature_25\", \"feature_45\", \"feature_60\", \"feature_58\", \"feature_39\", \"feature_66\",\n           \"feature_08\", \"feature_68\", \"feature_52\", \"feature_70\", \"feature_48\", \"feature_24\", \n            \"feature_65\", \"feature_74\"]\n\nclass CONFIG:\n    target_col = \"responder_6\"\n    start_dt = 1000\n    selected_columns = [\"date_id\", 'time_id', 'symbol_id', 'responder_6', 'weight'] + features\n\n\ndef load_and_split_data(parquet_path, config):\n    df = pl.scan_parquet(parquet_path).select(\n        config.selected_columns\n    ).select(\n        pl.int_range(pl.len(), dtype=pl.UInt32).alias(\"id\"),\n        pl.all(),\n    ).filter(\n        pl.col(\"date_id\").gt(config.start_dt)\n    ).collect()\n    \n    # Create lag feature\n    df = df.sort(['symbol_id', 'date_id', 'time_id'])\n   \n    print(f\"日期範圍: {df['date_id'].min()} - {df['date_id'].max()}\")\n    return df\n\n# 使用示例\ndf = load_and_split_data(\n    \"/kaggle/input/jane-street-real-time-market-data-forecasting/train.parquet\",\n    CONFIG\n)\ndf = df.to_pandas()\ndf = df.drop('id', axis=1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-24T11:16:03.669244Z","iopub.execute_input":"2024-12-24T11:16:03.669671Z","iopub.status.idle":"2024-12-24T11:16:35.954646Z","shell.execute_reply.started":"2024-12-24T11:16:03.669622Z","shell.execute_reply":"2024-12-24T11:16:35.953525Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import multiprocessing\nmultiprocessing.set_start_method('spawn')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-24T11:16:35.957045Z","iopub.execute_input":"2024-12-24T11:16:35.957491Z","iopub.status.idle":"2024-12-24T11:16:35.965548Z","shell.execute_reply.started":"2024-12-24T11:16:35.957455Z","shell.execute_reply":"2024-12-24T11:16:35.964337Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install pytorch-forecasting\n!pip --quiet install pytorch_lightning","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-24T11:16:35.966938Z","iopub.execute_input":"2024-12-24T11:16:35.967385Z","iopub.status.idle":"2024-12-24T11:17:05.721844Z","shell.execute_reply.started":"2024-12-24T11:16:35.967336Z","shell.execute_reply":"2024-12-24T11:17:05.719439Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import copy\nfrom pathlib import Path\nimport warnings\n\nimport lightning.pytorch as pl\nfrom lightning.pytorch.callbacks import EarlyStopping, LearningRateMonitor\nfrom lightning.pytorch.loggers import TensorBoardLogger\nimport numpy as np\nimport pandas as pd\nimport torch\n\nfrom pytorch_forecasting import Baseline, TemporalFusionTransformer, TimeSeriesDataSet\nfrom pytorch_forecasting.data import GroupNormalizer\nfrom pytorch_forecasting.metrics import MAE, SMAPE, PoissonLoss, QuantileLoss\nfrom pytorch_forecasting.models.temporal_fusion_transformer.tuning import optimize_hyperparameters\n\ncolumns_with_na = df.columns[df.isna().any()].tolist()\nfor column in columns_with_na:\n    df[column] = df.groupby('symbol_id')[column].ffill()\n    df[column] = df.groupby('symbol_id')[column].bfill()\n    df[column] = df[column].fillna(0)\n\ndef create_sequential_id_with_duplicates(df):\n    \"\"\"\n    將 date_id 和 time_id 合併成連續的 ID，相同的組合會得到相同的 ID\n    \n    Parameters:\n    df (pandas.DataFrame): 包含 date_id、time_id 的 DataFrame\n    \n    Returns:\n    pandas.DataFrame: 添加了 sequential_id 的 DataFrame\n    \"\"\"\n    # 獲取唯一的 (date_id, time_id) 組合並排序\n    unique_combinations = (\n        df[['date_id', 'time_id']]\n        .drop_duplicates()\n        .sort_values(['date_id', 'time_id'])\n    )\n    \n    # 為唯一組合創建 ID 映射字典\n    id_mapping = {\n        (date_id, time_id): i \n        for i, (date_id, time_id) in enumerate(\n            zip(unique_combinations['date_id'], \n                unique_combinations['time_id'])\n        )\n    }\n    \n    # 為每一行添加 sequential_id\n    df['sequential_id'] = df.apply(\n        lambda row: id_mapping[(row['date_id'], row['time_id'])], \n        axis=1\n    )\n    \n    return df.sort_values(['symbol_id', 'sequential_id'])\ndf = create_sequential_id_with_duplicates(df)\n# 1. 獲取唯一的date_id並排序\nunique_dates = sorted(df['date_id'].unique())\n\n# 2. 計算10%分位點的索引\ncutoff_index = int(len(unique_dates) * 0.8)  # 取90%位置，即最後10%的起始點\n\n# 3. 獲取切分的date_id\ncutoff_date = unique_dates[cutoff_index]\n\n# 4. 找出對應的最小sequential_id\ncutoff_sequential_id = df[df['date_id'] >= cutoff_date]['sequential_id'].min()\n\nprint(f\"切分date_id: {cutoff_date}\")\nprint(f\"對應的sequential_id: {cutoff_sequential_id}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-24T11:17:05.726164Z","iopub.execute_input":"2024-12-24T11:17:05.726796Z","iopub.status.idle":"2024-12-24T11:24:29.476299Z","shell.execute_reply.started":"2024-12-24T11:17:05.726732Z","shell.execute_reply":"2024-12-24T11:24:29.474989Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#max_encoder_length = 500\n#encoder_data = df[lambda x: x.sequential_id > x.sequential_id.max() - max_encoder_length]\n#encoder_data.to_csv(f'encoder_data_{max_encoder_length}.csv', index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-24T11:24:29.477840Z","iopub.execute_input":"2024-12-24T11:24:29.478421Z","iopub.status.idle":"2024-12-24T11:24:29.483346Z","shell.execute_reply.started":"2024-12-24T11:24:29.478385Z","shell.execute_reply":"2024-12-24T11:24:29.482163Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"warnings.filterwarnings('ignore', message='X does not have valid feature names')\nwarnings.filterwarnings(\"ignore\", category=FutureWarning)\nwarnings.filterwarnings(\"ignore\", category=UserWarning)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-24T11:24:29.484686Z","iopub.execute_input":"2024-12-24T11:24:29.485310Z","iopub.status.idle":"2024-12-24T11:24:29.497741Z","shell.execute_reply.started":"2024-12-24T11:24:29.485275Z","shell.execute_reply":"2024-12-24T11:24:29.496617Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df['symbol_id'] = df['symbol_id'].astype(str)\nmax_prediction_length = 10\nmax_encoder_length = 500\n\n\n\ntraining = TimeSeriesDataSet(\n    df[lambda x: x.sequential_id <= cutoff_sequential_id],\n    time_idx=\"sequential_id\",\n    target=\"responder_6\",\n    group_ids=[\"symbol_id\"],\n    min_encoder_length=1, #max_encoder_length // 2,  \n    max_encoder_length=max_encoder_length,\n    min_prediction_length=1,\n    max_prediction_length=max_prediction_length,\n    static_categoricals=[\"symbol_id\"],\n    time_varying_known_reals=[\"date_id\", \"time_id\", 'sequential_id']+features,\n    time_varying_unknown_reals=['responder_6'],\n    weight = \"weight\",\n    #target_normalizer=GroupNormalizer(\n        #groups=[\"symbol_id\"], transformation=\"softplus\"\n    #),\n    add_encoder_length=True,\n    allow_missing_timesteps=True,\n)\n\n# create validation set (predict=True) which means to predict the last max_prediction_length points in time\n# for each series\nvalidation = TimeSeriesDataSet.from_dataset(training, df, predict=True, stop_randomization=True)\n\n# create dataloaders for model\nbatch_size = 128  # set this between 32 to 128\ntrain_dataloader = training.to_dataloader(train=True, batch_size=batch_size, num_workers=0)\nval_dataloader = validation.to_dataloader(train=False, batch_size=batch_size * 10, num_workers=0)\nearly_stop_callback = EarlyStopping(monitor=\"val_loss\", min_delta=1e-4, patience=10, verbose=False, mode=\"min\")\nlr_logger = LearningRateMonitor()  # log the learning rate\nlogger = TensorBoardLogger(\"lightning_logs\")  # logging results to a tensorboard\ntrainer = pl.Trainer(\n    max_epochs=300,\n    enable_model_summary=True,\n    gradient_clip_val=0.1,\n    limit_train_batches=50,  # coment in for training, running valiation every 30 batches\n    # fast_dev_run=True,  # comment in to check that networkor dataset has no serious bugs\n    callbacks=[lr_logger, early_stop_callback],\n    logger=logger,\n)\n\ntft = TemporalFusionTransformer.from_dataset(\n    training,\n    learning_rate=0.0011,\n    hidden_size=16,#94\n    attention_head_size=8,\n    dropout=0.25,\n    hidden_continuous_size=16,#8\n    loss=QuantileLoss(),\n    log_interval=10,  # uncomment for learning rate finder and otherwise, e.g. to 10 for logging every 10 batches\n    optimizer=\"AdamW\",\n    reduce_on_plateau_patience=4,\n)\n\nimport logging\nlogging.getLogger(\"lightning.pytorch\").setLevel(logging.ERROR)\n#trainer.fit(\n    #tft,\n    #train_dataloaders=train_dataloader,\n    #val_dataloaders=val_dataloader,\n#)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-24T11:24:29.501621Z","iopub.execute_input":"2024-12-24T11:24:29.502199Z","iopub.status.idle":"2024-12-24T11:27:30.974096Z","shell.execute_reply.started":"2024-12-24T11:24:29.502150Z","shell.execute_reply":"2024-12-24T11:27:30.973026Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"state_dict = torch.load('/kaggle/input/tft_10/pytorch/default/1/kaggle/working/tft_pre_10.pth', map_location='cpu')\ntft.load_state_dict(state_dict)\npredictions = tft.predict(val_dataloader, return_y=True, trainer_kwargs=dict(accelerator=\"cpu\"))\nMAE()(predictions.output, predictions.y)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-24T11:43:54.452808Z","iopub.execute_input":"2024-12-24T11:43:54.453236Z","iopub.status.idle":"2024-12-24T11:44:14.741461Z","shell.execute_reply.started":"2024-12-24T11:43:54.453201Z","shell.execute_reply":"2024-12-24T11:44:14.740274Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"raw_predictions = tft.predict(val_dataloader, mode=\"raw\", return_x=True)\nfor idx in range(4):  \n    tft.plot_prediction(raw_predictions.x, raw_predictions.output, idx=idx, add_loss_to_title=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-24T11:46:12.876652Z","iopub.execute_input":"2024-12-24T11:46:12.877461Z","iopub.status.idle":"2024-12-24T11:46:16.796241Z","shell.execute_reply.started":"2024-12-24T11:46:12.877419Z","shell.execute_reply":"2024-12-24T11:46:16.795081Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"predictions = tft.predict(val_dataloader, return_x=True)\npredictions_vs_actuals = tft.calculate_prediction_actual_by_variable(predictions.x, predictions.output)\ntft.plot_prediction_actual_by_variable(predictions_vs_actuals)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-24T11:48:33.758563Z","iopub.execute_input":"2024-12-24T11:48:33.759025Z","iopub.status.idle":"2024-12-24T11:48:51.216682Z","shell.execute_reply.started":"2024-12-24T11:48:33.758989Z","shell.execute_reply":"2024-12-24T11:48:51.215304Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"interpretation = tft.interpret_output(raw_predictions.output, reduction=\"sum\")\ntft.plot_interpretation(interpretation)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-24T11:51:07.937371Z","iopub.execute_input":"2024-12-24T11:51:07.938261Z","iopub.status.idle":"2024-12-24T11:51:09.324992Z","shell.execute_reply.started":"2024-12-24T11:51:07.938218Z","shell.execute_reply":"2024-12-24T11:51:09.323801Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#切分date_id: 1559\n#對應的sequential_id: 540144","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dependency = tft.predict_dependency(\n    val_dataloader.dataset, \"date_id\", np.linspace(1559, df['date_id'].max(),  100), show_progress_bar=True, mode=\"dataframe\"\n)\n# plotting median and 25% and 75% percentile\nagg_dependency = dependency.groupby(\"date_id\").normalized_prediction.agg(\n    median=\"median\", q25=lambda x: x.quantile(0.25), q75=lambda x: x.quantile(0.75)\n)\nax = agg_dependency.plot(y=\"median\")\nax.fill_between(agg_dependency.index, agg_dependency.q25, agg_dependency.q75, alpha=0.3)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-24T11:54:54.765366Z","iopub.execute_input":"2024-12-24T11:54:54.765782Z","iopub.status.idle":"2024-12-24T11:55:51.528735Z","shell.execute_reply.started":"2024-12-24T11:54:54.765750Z","shell.execute_reply":"2024-12-24T11:55:51.527437Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dependency = tft.predict_dependency(\n    val_dataloader.dataset, \"sequential_id\", np.linspace(540144, df['sequential_id'].max(),  100), show_progress_bar=True, mode=\"dataframe\"\n)\n# plotting median and 25% and 75% percentile\nagg_dependency = dependency.groupby(\"sequential_id\").normalized_prediction.agg(\n    median=\"median\", q25=lambda x: x.quantile(0.25), q75=lambda x: x.quantile(0.75)\n)\nax = agg_dependency.plot(y=\"median\")\nax.fill_between(agg_dependency.index, agg_dependency.q25, agg_dependency.q75, alpha=0.3)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-24T11:55:51.531555Z","iopub.execute_input":"2024-12-24T11:55:51.532116Z","iopub.status.idle":"2024-12-24T11:56:45.849752Z","shell.execute_reply.started":"2024-12-24T11:55:51.532059Z","shell.execute_reply":"2024-12-24T11:56:45.848462Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dependency = tft.predict_dependency(\n    val_dataloader.dataset, \"time_id\", np.linspace(df['time_id'].min(), df['time_id'].max(),  100), show_progress_bar=True, mode=\"dataframe\"\n)\n# plotting median and 25% and 75% percentile\nagg_dependency = dependency.groupby(\"time_id\").normalized_prediction.agg(\n    median=\"median\", q25=lambda x: x.quantile(0.25), q75=lambda x: x.quantile(0.75)\n)\nax = agg_dependency.plot(y=\"median\")\nax.fill_between(agg_dependency.index, agg_dependency.q25, agg_dependency.q75, alpha=0.3)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-24T11:56:45.851410Z","iopub.execute_input":"2024-12-24T11:56:45.851820Z","iopub.status.idle":"2024-12-24T11:57:36.516545Z","shell.execute_reply.started":"2024-12-24T11:56:45.851783Z","shell.execute_reply":"2024-12-24T11:57:36.515014Z"}},"outputs":[],"execution_count":null}]}