{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":84493,"databundleVersionId":9871156,"sourceType":"competition"}],"dockerImageVersionId":30787,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import polars as pl\nimport pandas as pd\nimport numpy as np\nfrom typing import List, Tuple\nimport gc\nfrom abc import ABC, abstractmethod\n\nimport pyarrow.parquet as pq\n\nimport kaggle_evaluation.jane_street_inference_server\n\nimport os\nimport shutil\nimport glob\n\nimport multiprocessing\nfrom multiprocessing import Pool\nimport functools\n\nimport matplotlib.pyplot as plt\n\nfrom statsmodels.tsa.arima.model import ARIMA","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2024-12-26T09:18:57.095555Z","iopub.execute_input":"2024-12-26T09:18:57.096000Z","iopub.status.idle":"2024-12-26T09:18:59.888261Z","shell.execute_reply.started":"2024-12-26T09:18:57.095947Z","shell.execute_reply":"2024-12-26T09:18:59.887343Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ArimaTrainer():\n    def __init__(\n            self,\n            train_size: int, # in days\n            data: pd.DataFrame = None,\n        ):\n        \"\"\"\n        Initialize the ArimaTrainer class with training size and data.\n\n        Args:\n            train_size (int): Number of days to be used for training.\n            data (pd.DataFrame, optional): DataFrame containing the data to be used. Defaults to None.\n        \"\"\"\n\n        self.train_size = train_size\n        self.data = data\n        self.unique_dates = []\n\n    def get_last_n_dates_per_symbol(self, data: pd.DataFrame, train_size: int) -> pd.DataFrame:\n        \"\"\"\n        Load data from a DataFrame and return the last n dates for each symbol_id\n        \n        Args:\n            data: Pandas DataFrame with training data\n            train_size: Number of distinct dates to keep per symbol\n            \n        Returns:\n            pd.DataFrame: Filtered DataFrame with last n dates per symbol\n        \"\"\"\n        \n        # Get unique dates per symbol\n        unique_dates = data[['symbol_id', 'date_id']].drop_duplicates()\n        \n        # Sort by symbol_id and date_id in descending order, then group and take top N dates\n        top_dates = (\n            unique_dates\n            .sort_values(by=['symbol_id', 'date_id'], ascending=[True, False])\n            .groupby('symbol_id')\n            .head(train_size)\n        )\n        \n        # Filter original data using these dates\n        final_result = (\n            data\n            .merge(top_dates, on=['symbol_id', 'date_id'], how='inner')\n            .sort_values(by=['symbol_id', 'date_id', 'time_id'])\n            [['date_id', 'time_id', 'symbol_id', 'responder_6', 'weight']]\n        )\n        \n        # Validate that each symbol has exactly train_size distinct dates\n        date_counts = final_result.groupby(\"symbol_id\")['date_id'].nunique().reset_index(name='distinct_date_count')\n        \n        assert (date_counts['distinct_date_count'] == train_size).all(), (\n            f\"Not all symbols have exactly {train_size} distinct dates. \"\n            f\"Found counts: {date_counts[date_counts['distinct_date_count'] != train_size]}\"\n        )\n        \n        return final_result\n\n    def _train_and_predict(self, s):\n        \"\"\"\n        Train an ARIMA model for a given symbol and make predictions.\n\n        Args:\n            s: Symbol identifier for which the model is trained.\n\n        Returns:\n            tuple: A tuple containing arrays of symbols and predictions.\n        \"\"\"\n        print(f'Starting training of the symbol {s}\\n')\n        \n        df_train = self.data.loc[\n            self.data['symbol_id'] == s\n        ]\n\n        if len(df_train) > 0:\n            df_train = self.get_last_n_dates_per_symbol(df_train, self.train_size).reset_index(drop=True)\n            model = ARIMA(df_train['responder_6'], order=(1, 0, 0))\n            model_fit = model.fit()\n\n            # After date_id 677 the time units per day stabilizes in 968\n            y_pred = model_fit.forecast(steps=968).astype(np.float32).to_numpy() \n            symbol = np.full(len(y_pred), s)\n        else:\n            y_pred = np.full(968, 0.0).astype(np.float32)\n            symbol = np.full(len(y_pred), s)\n        \n        return symbol, y_pred\n\n    def run(self):\n        \"\"\"\n        Execute the ARIMA training and prediction process for all symbols.\n\n        Returns:\n            tuple: Two concatenated arrays containing symbols and predictions.\n        \"\"\"\n        unique_symbols = self.data['symbol_id'].unique()\n        \n        # Use multiprocessing to parallelize the computation\n        with multiprocessing.Pool() as pool:\n            results = list(pool.map(self._train_and_predict, unique_symbols))\n\n        # Unpack the results into two separate lists\n        symbol, pred = zip(*results)\n        symbol = np.concatenate(symbol)\n        pred = np.concatenate(pred)\n\n        return symbol, pred\n            \n        \n            ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T09:18:59.889526Z","iopub.execute_input":"2024-12-26T09:18:59.890036Z","iopub.status.idle":"2024-12-26T09:18:59.901571Z","shell.execute_reply.started":"2024-12-26T09:18:59.889994Z","shell.execute_reply":"2024-12-26T09:18:59.900669Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class JanePredictor():\n    def __init__(\n        self,\n        initial_data: pl.DataFrame,\n        train_size: int\n    ):\n        \"\"\"\n        Initializes the JanePredictor class with initial data and training size.\n\n        Parameters:\n        initial_data (pl.DataFrame): The initial dataset to be used for predictions.\n        train_size (int): The size of the training dataset.\n\n        Attributes:\n        lags_ (None): Placeholder for lag data.\n        cached_test_data (pl.DataFrame): Stores the initial data for caching purposes.\n        train_size (int): Stores the size of the training data.\n        \"\"\"\n        self.lags_ = None\n        self.cached_test_data = initial_data\n        self.train_size = train_size\n        \n    \n    def predict(self, test: pl.DataFrame, lags: pl.DataFrame | None) -> pl.DataFrame | pd.DataFrame:\n        \"\"\"\n        Predicts future values based on the test data and optional lag data.\n\n        Parameters:\n        test (pl.DataFrame): The test dataset containing new observations.\n        lags (pl.DataFrame | None): Optional lagged data to enhance prediction accuracy.\n\n        Returns:\n        pl.DataFrame | pd.DataFrame: A DataFrame containing predictions with 'row_id' and 'responder_6' columns.\n\n        Raises:\n        TypeError: If the returned predictions are not a DataFrame.\n        \n        Notes:\n        - Aligns column types between cached data and lag data.\n        - Joins lag data with cached test data to fill missing values.\n        - Uses an ARIMA model to generate predictions.\n        - Ensures that the output DataFrame has the same number of rows as the input test data.\n        \"\"\"\n        \n        if lags is not None:\n            # Cast columns to a common data type if necessary\n            self.cached_test_data = self.cached_test_data.with_columns([\n                pl.col(\"date_id\").cast(pl.Int32),\n                pl.col(\"time_id\").cast(pl.Int32),\n                pl.col(\"symbol_id\").cast(pl.Int32)\n            ])\n            \n            lags = lags.with_columns([\n                pl.col(\"date_id\").cast(pl.Int32),\n                pl.col(\"time_id\").cast(pl.Int32),\n                pl.col(\"symbol_id\").cast(pl.Int32)\n            ])\n            \n            lags = lags.with_columns(\n                (pl.col('date_id') - 1).alias('date_id')\n            ).select(['date_id','time_id','symbol_id','responder_6_lag_1'])\n\n            self.cached_test_data = self.cached_test_data.join(\n                lags,\n                on=['date_id','time_id','symbol_id'],\n                how='left'\n            )\n\n            self.cached_test_data = self.cached_test_data.with_columns(\n                pl.when((pl.col(\"responder_6\") == 0.0) & (pl.col('responder_6_lag_1') != 0.0))\n                .then(pl.col('responder_6_lag_1'))\n                .otherwise(pl.col('responder_6'))\n                .alias(\"responder_6\")\n            ).drop(['responder_6_lag_1'])\n            \n\n            trainer = ArimaTrainer(\n                train_size=self.train_size,\n                data=self.cached_test_data.to_pandas(),\n            )\n\n            symbol, pred = trainer.run()\n\n            self.symbol_arr = np.array(symbol)\n            self.pred_arr = np.array(pred)\n\n        # Get unique symbols and their first occurrences\n        _, first_indices = np.unique(self.symbol_arr, return_index=True)\n\n        symbol = self.symbol_arr[first_indices]\n        pred = self.pred_arr[first_indices]\n\n        self.symbol_arr = np.delete(self.symbol_arr, first_indices)\n        self.pred_arr = np.delete(self.pred_arr, first_indices)\n\n        pred_df = pl.DataFrame({\n            'symbol_id': symbol,\n            'responder_6': pred\n        },\n            schema={\n                'symbol_id': pl.Int8,\n                'responder_6': pl.Float32\n            }\n        )\n        \n        self.cached_test_data = pl.concat([\n            self.cached_test_data, \n            test.with_columns(\n                pl.lit(0.0).cast(pl.Float32).alias('responder_6')\n            ).select(['date_id','time_id','symbol_id','responder_6','weight'])\n        ],how='vertical_relaxed')\n\n        predictions = test.join(pred_df, on=['symbol_id'], how='left').select(['row_id','responder_6'])\n    \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        \n        # Confirm has as many rows as the test data.\n        assert len(predictions) == len(test)\n\n        print(predictions.head())\n    \n        return predictions","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T09:18:59.904298Z","iopub.execute_input":"2024-12-26T09:18:59.904691Z","iopub.status.idle":"2024-12-26T09:18:59.926627Z","shell.execute_reply.started":"2024-12-26T09:18:59.904643Z","shell.execute_reply":"2024-12-26T09:18:59.925660Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_last_n_dates_per_symbol(file_path: str, train_size: int) -> pl.DataFrame:\n    \"\"\"\n    Load data from parquet file and return the last n dates for each symbol_id\n    \n    Args:\n        file_path: Path to the parquet file\n        train_size: Number of distinct dates to keep per symbol\n        \n    Returns:\n        pl.DataFrame: Filtered dataframe with last n dates per symbol\n    \"\"\"\n    # Load data\n    data = pl.scan_parquet(file_path)\n    \n    # Get unique dates per symbol\n    unique_dates = (\n        data\n        .select(['symbol_id', 'date_id'])\n        .unique()\n    )\n    \n    # Sort and get top N dates per symbol\n    top_dates = (\n        unique_dates\n        .sort(['symbol_id', 'date_id'], descending=[False, True])\n        .group_by('symbol_id')\n        .head(train_size)\n    )\n    \n    # Filter original data using these dates\n    final_result = (\n        data\n        .join(\n            top_dates,\n            on=['symbol_id', 'date_id'],\n            how='inner'\n        )\n        .sort(['symbol_id', 'date_id', 'time_id'])\n        .select(\n            ['date_id', 'time_id', 'symbol_id', 'responder_6', 'weight']\n        )\n        .collect()\n    )\n    \n    # Validate that each symbol has exactly train_size distinct dates\n    date_counts = final_result.group_by(\"symbol_id\").agg(\n        pl.col(\"date_id\").n_unique().alias(\"distinct_date_count\")\n    )\n    \n    assert (date_counts['distinct_date_count'] == train_size).all(), (\n        f\"Not all symbols have exactly {train_size} distinct dates. \"\n        f\"Found counts: {date_counts.filter(pl.col('distinct_date_count') != train_size)}\"\n    )\n    \n    return final_result","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T09:18:59.927686Z","iopub.execute_input":"2024-12-26T09:18:59.928002Z","iopub.status.idle":"2024-12-26T09:18:59.939688Z","shell.execute_reply.started":"2024-12-26T09:18:59.927978Z","shell.execute_reply":"2024-12-26T09:18:59.938839Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"file_path = '/kaggle/input/jane-street-real-time-market-data-forecasting/train.parquet'\ntrain_size = 32\ndata = get_last_n_dates_per_symbol(file_path, train_size)\n\njane_predictor = JanePredictor(\n    initial_data = data,\n    train_size = train_size\n)\n\ninference_server = kaggle_evaluation.jane_street_inference_server.JSInferenceServer(jane_predictor.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":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-26T09:18:59.943513Z","iopub.execute_input":"2024-12-26T09:18:59.943907Z","iopub.status.idle":"2024-12-26T09:20:28.752040Z","shell.execute_reply.started":"2024-12-26T09:18:59.943868Z","shell.execute_reply":"2024-12-26T09:20:28.750984Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}