{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":84493,"databundleVersionId":9871156,"sourceType":"competition"},{"sourceId":10448669,"sourceType":"datasetVersion","datasetId":6467544},{"sourceId":10460978,"sourceType":"datasetVersion","datasetId":6471753}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install zarr -q --no-index --find-links=/kaggle/input/zarr-package/","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-13T20:50:35.261813Z","iopub.execute_input":"2025-01-13T20:50:35.262115Z","iopub.status.idle":"2025-01-13T20:50:38.470919Z","shell.execute_reply.started":"2025-01-13T20:50:35.262094Z","shell.execute_reply":"2025-01-13T20:50:38.469834Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!cp -r /kaggle/input/jane-street-market-library/data/symbol_data.zarr /kaggle/working/sysymbol_data.zarr","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-13T20:50:38.472486Z","iopub.execute_input":"2025-01-13T20:50:38.472771Z","iopub.status.idle":"2025-01-13T20:50:39.049431Z","shell.execute_reply.started":"2025-01-13T20:50:38.472735Z","shell.execute_reply":"2025-01-13T20:50:39.048301Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport polars as pl\nimport zarr\nfrom zarr import ProcessSynchronizer\nfrom zarr import DirectoryStore\nfrom sklearn.preprocessing import StandardScaler, OrdinalEncoder\nfrom sklearn.compose import ColumnTransformer\nfrom sklearn import set_config\nimport torch\nfrom tqdm import tqdm, trange\nfrom numpy.typing import NDArray\nfrom functools import lru_cache\nfrom omegaconf import OmegaConf\nimport json\nimport sys\nsys.path.append('/kaggle/input/jane-street-market-library')\n\nset_config(transform_output = \"default\")\n\nstore = DirectoryStore('/kaggle/working/symbol_data.zarr')\nsynchronizer = ProcessSynchronizer('/kaggle/working/sync.lock')\nroot = zarr.group(store, overwrite=False, synchronizer=synchronizer)\nsymb_last_dates = {k:0 for k in range(39)}\nlatest_time_idx = 0 \nwith open('/kaggle/input/jane-street-market-library/data/categories.json') as f:\n    categories = json.load(f)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-01-13T20:50:40.547974Z","iopub.execute_input":"2025-01-13T20:50:40.548325Z","iopub.status.idle":"2025-01-13T20:50:40.557139Z","shell.execute_reply.started":"2025-01-13T20:50:40.548294Z","shell.execute_reply":"2025-01-13T20:50:40.556171Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_data = pl.scan_parquet('/kaggle/input/jane-street-real-time-market-data-forecasting/train.parquet')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-13T20:50:43.377886Z","iopub.execute_input":"2025-01-13T20:50:43.378191Z","iopub.status.idle":"2025-01-13T20:50:43.382905Z","shell.execute_reply.started":"2025-01-13T20:50:43.378169Z","shell.execute_reply":"2025-01-13T20:50:43.381872Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# DataStore","metadata":{}},{"cell_type":"code","source":"from time import perf_counter\n\n\n# @torch.compile\ndef format_dec_data(dec_tensor: list[torch.Tensor]):\n    dec_tensor_ = torch.stack(dec_tensor)\n    return {\n        \"decoder_lengths\": torch.tensor(dec_tensor_.shape[0]).to(dtype=torch.int64),\n        \"decoder_categoricals\": dec_tensor_[:, -5:].to(dtype=torch.int64).long(),\n        \"decoder_reals\": dec_tensor_[:, :-5].to(dtype=torch.float32),\n    }\n\n\n# @torch.compile\ndef merge_current_data(\n    curr_day_dec_store: dict[int, list[torch.Tensor]], timestep: torch.Tensor\n):\n    for data_idx in range(timestep.size(0)):\n        symbol_id = int(timestep[data_idx, -1].item())\n        symbol_data = timestep[data_idx, :]\n        if symbol_id not in curr_day_dec_store:\n            curr_day_dec_store[symbol_id] = [symbol_data]\n        else:\n            curr_day_dec_store[symbol_id].append(symbol_data)\n    # print({k:len(v) for k,v in curr_day_dec_store.items()})\n    return curr_day_dec_store\n\n# @torch.compile\ndef format_enc_data(data: torch.Tensor):\n    device = data.device\n    return {\n        \"encoder_lengths\": torch.tensor(data.shape[0]).to(dtype=torch.int64, device=device),\n        \"encoder_categoricals\": data[:, :-9][:, -5:].to(dtype=torch.int64, device=device).long(),\n        \"encoder_reals\": data[:, :79].to(dtype=torch.float32, device=device),\n        \"encoder_targets\": data[:, -9:].to(dtype=torch.float32, device=device),\n    }\n\nclass DataStore:\n    def __init__(self, *args, **kwargs):\n        self.curr_day_dec_store: dict[str, list[torch.Tensor]] = {}\n        self.store = DirectoryStore(\"/kaggle/working/symbol_data.zarr\")\n        self.synchronizer = ProcessSynchronizer(\"/kaggle/working/sync.lock\")\n        self.root = zarr.group(store, overwrite=False, synchronizer=synchronizer)\n        self.symb_last_dates = {k: 1695 for k in range(39)}\n        self.latest_time_idx = 0\n        self.col_order = ['row_id', 'date_id', 'time_id', 'symbol_id', 'weight', 'feature_00', 'feature_01', 'feature_02', 'feature_03', 'feature_04', 'feature_05', 'feature_06', 'feature_07', 'feature_08', 'feature_09', 'feature_10', 'feature_11', 'feature_12', 'feature_13', 'feature_14', 'feature_15', 'feature_16', 'feature_17', 'feature_18', 'feature_19', 'feature_20', 'feature_21', 'feature_22', 'feature_23', 'feature_24', 'feature_25', 'feature_26', 'feature_27', 'feature_28', 'feature_29', 'feature_30', 'feature_31', 'feature_32', 'feature_33', 'feature_34', 'feature_35', 'feature_36', 'feature_37', 'feature_38', 'feature_39', 'feature_40', 'feature_41', 'feature_42', 'feature_43', 'feature_44', 'feature_45', 'feature_46', 'feature_47', 'feature_48', 'feature_49', 'feature_50', 'feature_51', 'feature_52', 'feature_53', 'feature_54', 'feature_55', 'feature_56', 'feature_57', 'feature_58', 'feature_59', 'feature_60', 'feature_61', 'feature_62', 'feature_63', 'feature_64', 'feature_65', 'feature_66', 'feature_67', 'feature_68', 'feature_69', 'feature_70', 'feature_71', 'feature_72', 'feature_73', 'feature_74', 'feature_75', 'feature_76', 'feature_77', 'feature_78', 'is_scored']\n        self.categories = {\n            \"feature_09\": [\n                2,\n                4,\n                9,\n                11,\n                12,\n                14,\n                15,\n                25,\n                26,\n                30,\n                34,\n                42,\n                44,\n                46,\n                49,\n                50,\n                57,\n                64,\n                68,\n                70,\n                # 81,\n                # 82,\n            ],\n            \"feature_10\": [1, 2, 3, 4, 5, 6, 7,\n                           # 12\n                          ],\n            \"feature_11\": [\n                9,\n                11,\n                13,\n                16,\n                24,\n                25,\n                34,\n                40,\n                48,\n                50,\n                59,\n                62,\n                63,\n                66,\n                76,\n                150,\n                158,\n                159,\n                171,\n                195,\n                214,\n                230,\n                261,\n                297,\n                336,\n                376,\n                388,\n                410,\n                522,\n                # 534,\n                # 539,\n            ],\n            \"time_id\": [i for i in range(968)],\n            \"symbol_id\":[i for i in range(38)]\n        }\n        self.symbol_idx = 3\n        self.time_id_idx = 2\n        self.weight_idx = 4\n        self.date_id_scaler = StandardScaler()\n        self.date_id_scaler.fit(np.arange(0, 1900).reshape(-1, 1))\n        self.time_idx_scaler = StandardScaler()\n        self.time_idx_scaler.fit(np.arange(0, 1854468).reshape(-1, 1))\n        self.real_scaler = StandardScaler()\n        self.feat_idx = [\n            4,\n            5,\n            6,\n            7,\n            8,\n            9,\n            10,\n            11,\n            12,\n            16,\n            17,\n            18,\n            19,\n            20,\n            21,\n            22,\n            23,\n            24,\n            25,\n            26,\n            27,\n            28,\n            29,\n            30,\n            31,\n            32,\n            33,\n            34,\n            35,\n            36,\n            37,\n            38,\n            39,\n            40,\n            41,\n            42,\n            43,\n            44,\n            45,\n            46,\n            47,\n            48,\n            49,\n            50,\n            51,\n            52,\n            53,\n            54,\n            55,\n            56,\n            57,\n            58,\n            59,\n            60,\n            61,\n            62,\n            63,\n            64,\n            65,\n            66,\n            67,\n            68,\n            69,\n            70,\n            71,\n            72,\n            73,\n            74,\n            75,\n            76,\n            77,\n            78,\n            79,\n            80,\n            81,\n            82,\n        ]\n        self.dec_ct = ColumnTransformer(\n            transformers=[\n                (\"real_time_indices\", \"passthrough\", [0, 1]),\n                (\"weight_idx\", \"passthrough\", [self.weight_idx]),\n                (\"features\", StandardScaler(), self.feat_idx),\n                (\n                    \"cat_encoder_1\",\n                    OrdinalEncoder(\n                        categories=[\n                            self.categories[\"feature_09\"],\n                        ],\n                        handle_unknown=\"use_encoded_value\",\n                        unknown_value=len(categories[\"feature_09\"])-1\n                    ),\n                    [14],\n                ),\n                (\n                    \"cat_encoder_2\",\n                    OrdinalEncoder(\n                        categories=[\n                            self.categories[\"feature_10\"],\n                        ],\n                        handle_unknown=\"use_encoded_value\",\n                        unknown_value=len(categories[\"feature_10\"])-1,\n                    ),\n                    [15],\n                ),\n                (\n                    \"cat_encoder_3\",\n                    OrdinalEncoder(\n                        categories=[\n                            self.categories[\"feature_11\"],\n                        ],\n                        handle_unknown=\"use_encoded_value\",\n                        unknown_value=len(categories[\"feature_11\"])-1,\n                    ),\n                    [16],\n                ),\n                (\"static_real\", \"passthrough\", [self.time_id_idx, self.symbol_idx]),\n            ],\n            remainder=\"drop\",\n        )\n        self.tensor_device = \"cpu\"\n        self.old_symb_order = []\n        self.symb_order_changed = False\n        self.curr_symb_order = []\n        self.new_symb_order = []\n        self.curr_day_enc_data = None\n        self.pred_device = \"cuda\"\n        self.is_scored = None\n\n    def prep_dec_df(self, timestep: pl.DataFrame) -> torch.Tensor:\n        symb_scaled = self.dec_ct.fit_transform(timestep.to_numpy())\n        symb_scaled[:, 0] = self.time_idx_scaler.transform(\n            symb_scaled[:, 0].reshape(-1, 1)\n        ).ravel()\n        symb_scaled[:, 1] = self.date_id_scaler.transform(\n            symb_scaled[:, 1].reshape(-1, 1)\n        ).ravel()\n        return torch.tensor(symb_scaled, dtype=torch.float32, device=self.pred_device)  ## Changed\n\n    def add_time_idx(self, timestep: pl.DataFrame) -> pl.DataFrame:\n        curr_date = timestep[\"date_id\"].max()\n        if curr_date < 677:\n            timestep = timestep.with_columns(pl.col(\"date_id\") + 1699)\n        return (\n            timestep.with_columns(\n                (\n                    (pl.col(\"date_id\").cast(pl.Int32) - 677) * 968\n                    + pl.col(\"time_id\").cast(pl.Int32)\n                    + (677 * 849)\n                ).alias(\"time_idx\")\n            )\n            .fill_null(0)\n            .select(\n                pl.col(\"time_idx\"),\n                pl.col(\n                    [\n                        x\n                        for x in timestep.columns\n                        if x not in [\"time_idx\", \"partition_id\"]\n                    ]\n                ),\n            )\n        )\n\n    def flush_curr_day_data(self, lags: pl.DataFrame):\n        symb_lags = lags.sort(\"time_id\").partition_by(\"symbol_id\", maintain_order=True)\n        curr_date = None\n        for symb_lag in symb_lags:\n            self.symb_last_dates[symb_lag[\"symbol_id\"].min()] = symb_lag[\n                \"date_id\"\n            ].min()\n            curr_symbol = symb_lag[\"symbol_id\"].min()\n            if curr_date is None:\n                curr_date = symb_lag[\"date_id\"].min()\n            symb_lag = (\n                symb_lag.select(\"^responder_.*$\")\n                .to_torch()\n                .to(dtype=torch.float32, device=self.tensor_device)\n            )\n            combined_data = torch.cat(\n                [torch.stack(self.curr_day_dec_store[curr_symbol],dim=0).detach().cpu(), symb_lag], dim=-1\n            )\n            self.root[f\"{curr_symbol}/{curr_date}\"] = zarr.array(\n                combined_data.detach().cpu().numpy(), chunks=(968, 93), dtype=np.float32\n            )\n        self.curr_day_dec_store = {}\n        self.curr_symb_order = []\n        self.get_prev_data.cache_clear()\n        self.curr_day_enc_data = None\n        self.curr_batch_enc_data = None\n\n    def update_curr_day_timestep(self, timestep: pl.DataFrame):\n        # Since Everything is arriving in order no need for sorting\n        # timestep = timestep.select(self.col_order)\n        self.is_scored = timestep['is_scored'].to_numpy().ravel().astype('bool')\n        self.is_scored  = torch.tensor(self.is_scored, device=self.pred_device, dtype=torch.bool)\n        self.row_id = timestep['row_id'].to_numpy().ravel()\n        timestep = timestep.drop(['row_id', 'is_scored'])\n        timestep = self.add_time_idx(timestep)\n        # print(timestep)\n        timestep: torch.Tensor = self.prep_dec_df(timestep)\n        \n        curr_symb_order = timestep[:, -1].cpu().numpy().tolist()\n        # print(timestep[:,-1])\n        \n        # curr_symb_order = curr_symb_order[is_scored]\n        # curr_symb_order = curr_symb_order.tolist()\n        \n        if self.curr_symb_order != curr_symb_order:\n            self.symb_order_changed = True\n            self.new_symb_order = curr_symb_order\n            self.old_symb_order = self.curr_symb_order\n            self.curr_symb_order = curr_symb_order\n        else:\n            self.symb_order_changed = False\n        self.curr_day_dec_store = merge_current_data(self.curr_day_dec_store, timestep)\n\n    @lru_cache(maxsize=None)\n    def get_prev_data(self, symbol_id: int, sampling=8):\n        # print(f'Symbol ID {symbol_id}')\n        data = torch.tensor(\n            self.root[f\"{int(symbol_id)}/{self.symb_last_dates[int(symbol_id)]}\"].oindex[::sampling],\n            dtype=torch.float32, device=self.pred_device\n        )\n        data_dict = format_enc_data(data)\n        return data_dict\n    \n    def get_curr_batch(self):\n        if self.curr_day_enc_data is None:\n            self.curr_day_enc_data = {x: self.get_prev_data(x) for x in self.curr_symb_order}\n            self.curr_batch_enc_data = {\n                \"encoder_reals\": torch.stack([self.curr_day_enc_data[x][\"encoder_reals\"] for x in self.curr_symb_order]).to(self.pred_device),\n                \"encoder_lengths\": torch.stack([self.curr_day_enc_data[x][\"encoder_lengths\"] for x in self.curr_symb_order]).to(self.pred_device),\n                \"encoder_targets\": torch.stack([self.curr_day_enc_data[x][\"encoder_targets\"] for x in self.curr_symb_order]).to(self.pred_device),\n                \"encoder_categoricals\": torch.stack([self.curr_day_enc_data[x][\"encoder_categoricals\"] for x in self.curr_symb_order]).to(self.pred_device),\n            }\n        if self.symb_order_changed:\n            self.curr_batch_enc_data = {\n            \"encoder_reals\": torch.stack([self.curr_day_enc_data[x][\"encoder_reals\"] for x in self.curr_symb_order]).to(self.pred_device),\n            \"encoder_lengths\": torch.stack([self.curr_day_enc_data[x][\"encoder_lengths\"] for x in self.curr_symb_order]).to(self.pred_device),\n            \"encoder_targets\": torch.stack([self.curr_day_enc_data[x][\"encoder_targets\"] for x in self.curr_symb_order]).to(self.pred_device),\n            \"encoder_categoricals\": torch.stack([self.curr_day_enc_data[x][\"encoder_categoricals\"] for x in self.curr_symb_order]).to(self.pred_device),\n            }\n        dec_data = {x:format_dec_data(self.curr_day_dec_store[x]) for x in self.curr_symb_order}\n        self.curr_batch_enc_data.update({\n            \"decoder_reals\": torch.stack([dec_data[x][\"decoder_reals\"] for x in self.curr_symb_order]).to(self.pred_device),\n            \"decoder_lengths\": torch.stack([dec_data[x][\"decoder_lengths\"] for x in self.curr_symb_order]).to(self.pred_device),\n            \"decoder_categoricals\": torch.stack([dec_data[x][\"decoder_categoricals\"] for x in self.curr_symb_order]).to(self.pred_device),\n        })\n        return self.curr_batch_enc_data, self.is_scored","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-13T20:50:45.440160Z","iopub.execute_input":"2025-01-13T20:50:45.440504Z","iopub.status.idle":"2025-01-13T20:50:45.682691Z","shell.execute_reply.started":"2025-01-13T20:50:45.440475Z","shell.execute_reply":"2025-01-13T20:50:45.682004Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"#### ColdSTart","metadata":{}},{"cell_type":"code","source":"def initialise_store(train_data=train_data):\n    DS = DataStore()\n    last_day = train_data.filter(pl.col(\"date_id\") == 1698).collect()\n    lags = last_day.select(\n        [\n            \"date_id\",\n            \"time_id\",\n            \"symbol_id\",\n            \"^responder_.*$\",\n        ]\n    )\n    test_data_timesteps = last_day.drop('^responder_.*$').with_row_index('row_id').with_columns(pl.lit(False).alias('is_scored')).partition_by('time_id')\n    print('Loading the Initial Data')\n    for ts in tqdm(test_data_timesteps):\n        DS.update_curr_day_timestep(ts)\n        pred = np.zeros(39)\n        if sum(DS.is_scored) > 0:\n            p, s = DS.get_curr_batch()\n            p = {k:v[s] for k, v in p.items()}\n            B,_,_ = p['encoder_reals'].size()\n            if B!=0:\n                o = model_1(p)\n                pred[s] = o.detach().cpu().numpy().ravel()\n        \n        # break\n        \n        \n    DS.flush_curr_day_data(lags)\n\n    return DS","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-13T20:50:52.371291Z","iopub.execute_input":"2025-01-13T20:50:52.371586Z","iopub.status.idle":"2025-01-13T20:50:52.377980Z","shell.execute_reply.started":"2025-01-13T20:50:52.371563Z","shell.execute_reply":"2025-01-13T20:50:52.376870Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Import Main Config File\nfrom jane_street.layers.tft import TFT\nfrom jane_street.layers.embeddings import get_embedding_size\nfrom jane_street.models.temporal_ft import TemporalFT\nfrom jane_street.models.tsmixer import  TSMixerExtModel\n# with open('/kaggle/input/jane-street-market-library/configs/master_config.yaml', 'r') as fp:\n#     config = OmegaConf.load(fp)\n\ndef make_tft_config():\n    with open('/kaggle/input/jane-street-market-library/data/categories.json','r') as fp:\n        categories = json.load(fp)\n    conf = OmegaConf.load(\"/kaggle/input/jane-street-market-library/configs/master_config.yaml\")\n    ## Data config\n    # conf.limit_val_batches = 800\n    # conf.val_dataset_size = 256*800*n_devices\n    conf.x_categoricals = list(categories.keys())\n    \n    conf.hidden_size = 8\n    cat_sizes = {\n        k: (len(v), get_embedding_size(len(v), 16))\n        for k, v in categories.items()\n    }\n    conf.embedding_sizes = cat_sizes\n    conf.real_hidden_size = conf.hidden_size\n    conf.log_interval = 2\n    conf.log_val_interval = 3\n    conf.lstm_layers = 2\n    conf.n_heads = 4\n    conf.dropout = 0.1\n    conf.learning_rate = 0.0001\n    conf.quantiles = [0.3, 0.5, 0.7]\n    conf.weight_decay = 0.001\n    conf.max_encoder_length = 968\n    conf.output_size = 1\n    conf.causal_attention = True\n    return conf\n\n# config = make_tft_config()\n# config.x_categoricals\ntft = TemporalFT.load_from_checkpoint('/kaggle/input/jane-street-market-library/main_ckpts/epoch=20-val_score=0.43-temporal-ft.ckpt')\n# tft = TFT(config)\n# tft = TFT(config)\ntft = tft.to('cuda').eval()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-13T20:50:54.953313Z","iopub.execute_input":"2025-01-13T20:50:54.953616Z","iopub.status.idle":"2025-01-13T20:50:56.034404Z","shell.execute_reply.started":"2025-01-13T20:50:54.953590Z","shell.execute_reply":"2025-01-13T20:50:56.033676Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"tft","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-13T20:51:06.585877Z","iopub.execute_input":"2025-01-13T20:51:06.586180Z","iopub.status.idle":"2025-01-13T20:51:06.627319Z","shell.execute_reply.started":"2025-01-13T20:51:06.586157Z","shell.execute_reply":"2025-01-13T20:51:06.626559Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Replace this function with your inference code.\n# You can return either a Pandas or Polars dataframe, though Polars is recommended.\n# Each batch of predictions (except the very first) must be returned within 1 minute of the batch features being provided.\nstarted = False\nlags_ = None\ndays_passed = 0\nDS = initialise_store()\ndef predict(test: pl.DataFrame, lags: pl.DataFrame | None) -> pl.DataFrame | pd.DataFrame:\n    \"\"\"Make a prediction.\"\"\"\n    # All the responders from the previous day are passed in at time_id == 0. We save them in a global variable for access at every time_id.\n    # Use them as extra features, if you like.\n    global lags_, DS, tft, started, days_passed\n    \n    curr_time_id = test['time_id'].min()\n    if curr_time_id == 0 and started and lags is not None:\n        DS.flush_curr_day_data(lags)\n        lags_ = lags\n        days_passed+=1\n    started=True\n    # if days_passed % 30 ==0:\n        \n    # if lags is not None:\n    #     lags_ = lags\n    with torch.no_grad():\n        DS.update_curr_day_timestep(test)\n        # Replace this section with your own predictions\n        predictions = test.select(\n            'row_id',\n            pl.lit(0.0).alias('responder_6'),\n        )\n        # pred = np.zeros_like(s)\n        pred = np.zeros_like(predictions['responder_6'].to_numpy().ravel())\n        # print(DS.curr_symb_order)\n        if sum(DS.is_scored) > 0:\n            p, s = DS.get_curr_batch()\n            # print(p['decoder_categoricals'][0])\n            p = {k:v[s] for k, v in p.items()}\n            B,_,_ = p['encoder_reals'].size()\n            if B!=0: # Only Predict if scored\n                o = tft(p)\n                s = s.detach().cpu().numpy().ravel()\n                pred[s] = o.detach().cpu().numpy().ravel()\n        predictions = predictions.with_columns(pl.Series('responder_6', pred.ravel()))\n        if isinstance(predictions, pl.DataFrame):\n            assert predictions.columns == ['row_id', 'responder_6']\n        elif isinstance(predictions, pd.DataFrame):\n            assert (predictions.columns == ['row_id', 'responder_6']).all()\n        else:\n            raise TypeError('The predict function must return a DataFrame')\n        # Confirm has as many rows as the test data.\n        assert len(predictions) == len(test)\n        # print(predictions)\n    return predictions","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-13T20:51:25.247071Z","iopub.execute_input":"2025-01-13T20:51:25.247357Z","iopub.status.idle":"2025-01-13T20:51:32.782608Z","shell.execute_reply.started":"2025-01-13T20:51:25.247336Z","shell.execute_reply":"2025-01-13T20:51:32.781699Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import kaggle_evaluation.jane_street_inference_server\n\nimport os\n\nimport pandas as pd\nimport polars as pl\n\ntest_data = pl.read_parquet('/kaggle/input/jane-street-real-time-market-data-forecasting/test.parquet')\nlags =  pl.read_parquet('/kaggle/input/jane-street-real-time-market-data-forecasting/lags.parquet')\n# predict(test_data, lags)\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    )\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-13T20:51:32.933818Z","iopub.execute_input":"2025-01-13T20:51:32.934026Z","iopub.status.idle":"2025-01-13T20:51:33.084834Z","shell.execute_reply.started":"2025-01-13T20:51:32.934008Z","shell.execute_reply":"2025-01-13T20:51:33.083474Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}