{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":84493,"databundleVersionId":9871156,"sourceType":"competition"},{"sourceId":9806342,"sourceType":"datasetVersion","datasetId":6010899},{"sourceId":10147049,"sourceType":"datasetVersion","datasetId":6263670},{"sourceId":10259189,"sourceType":"datasetVersion","datasetId":6346382},{"sourceId":214081095,"sourceType":"kernelVersion"},{"sourceId":205344,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":175117,"modelId":197477}],"dockerImageVersionId":30804,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nimport polars as pl\nimport numpy as np\nfrom numpy.typing import NDArray\n\nimport xgboost as xgb\n\nimport catboost as cbt\nfrom catboost import CatBoostRegressor\n\nimport os\nimport joblib\nimport gc\nfrom datetime import datetime\n\nimport sys\nsys.path.append(\"/kaggle/input/jane-street-real-time-market-data-forecasting\")\nimport kaggle_evaluation.jane_street_inference_server\n\nsys.path.append(\"/kaggle/input/cbt-online-config\")\nfrom config import Config\n\nsys.path.append(\"/kaggle/input/eval-metirix-cbt-online\")\nfrom eval_metrix import *","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2024-12-21T04:59:12.144212Z","iopub.execute_input":"2024-12-21T04:59:12.144574Z","iopub.status.idle":"2024-12-21T04:59:13.922312Z","shell.execute_reply.started":"2024-12-21T04:59:12.144539Z","shell.execute_reply":"2024-12-21T04:59:13.921601Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"config = Config","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-21T04:59:13.923680Z","iopub.execute_input":"2024-12-21T04:59:13.924055Z","iopub.status.idle":"2024-12-21T04:59:13.928061Z","shell.execute_reply.started":"2024-12-21T04:59:13.924027Z","shell.execute_reply":"2024-12-21T04:59:13.927119Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"config.cbt_params['iterations'] = 100\nconfig.cbt_params['depth'] = 8\nconfig.cbt_params['learning_rate'] = 0.1\nconfig.cbt_params","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-21T04:59:13.929236Z","iopub.execute_input":"2024-12-21T04:59:13.929964Z","iopub.status.idle":"2024-12-21T04:59:13.941651Z","shell.execute_reply.started":"2024-12-21T04:59:13.929938Z","shell.execute_reply":"2024-12-21T04:59:13.940901Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class JaneStreetCbtPredictor:\n    \n    def __init__(self, \n                 config: Config,\n                 cbt_model,\n                 train_data: pl.DataFrame):\n        self.config: Config = config\n        self.cbt_model = cbt_model\n        # self.lazy_train_data: pl.LazyFrame = lazy_train_data\n        self.train_data: pl.DataFrame = train_data # 是否考虑 lazyframe\n        self.lags: pl.DataFrame = None\n        \n        self.is_zero_day = True \n        self.retrain = False\n\n        self.count_date = 0\n        \n        self.max_date_id: int = (\n            self.train_data\n            .select(\"date_id\")\n            .max()\n            # .collect()\n            .to_series()\n            .to_numpy()[0]\n        )\n        self.min_date_id: int = (\n            self.train_data\n            .select(\"date_id\")\n            .min()\n            # .collect()\n            .to_series()\n            .to_numpy()[0]\n        )\n        \n        self.new_day_test: list = []\n    \n    def retrain_cbt(self):\n\n        start = datetime.now()\n        tmp = self.train_data.filter(pl.col(\"date_id\") >= self.min_date_id)\n        X_train = tmp.select(self.config.features).to_numpy()\n        y_train = tmp.select(self.config.target).to_numpy().squeeze()\n        w_train = tmp.select(self.config.sample_weight).to_numpy().squeeze()\n        print(f\"retrain data to numpy 用时：{datetime.now() - start}秒！\")\n        \n        self.cbt_model = CatBoostRegressor(\n                **self.config.cbt_params, \n                eval_metric=self.config.cbt_eval_metric,\n                random_seed=self.config.random_seeds_cbt[0]\n            )\n        self.cbt_model.fit(\n                X_train, \n                y_train, \n                sample_weight=w_train, \n                verbose=10,  # 输出详细日志\n            )\n        \n        del X_train, y_train, w_train\n        del tmp\n        gc.collect()\n    \n    def update_train_data(self, lags: pl.DataFrame):\n\n        test = self.join_lags_with_test(lags)\n        gc.collect()\n\n        self.train_data= pl.concat([self.train_data, test])\n        \n        # 数据更新完成后，将 new_day_test 重新弄设置为空列表，接收下一个 date 的数据\n        self.new_day_test = []\n        \n        gc.collect()\n        \n        return\n    \n    def join_lags_with_test(self, lags: pl.DataFrame):\n        \n        test = pl.concat(self.new_day_test)\n        test = (\n            test\n            .with_columns(\n                pl.lit(self.max_date_id, dtype=pl.Int16).alias(\"date_id\")\n            )\n        )\n        \n        # lags 2 label\n        lags2label = (\n            lags\n            .with_columns(\n                pl.lit(self.max_date_id, dtype=pl.Int16).alias(\"date_id\")\n            )\n            .rename({f\"{resp}_lag_1\": resp for resp in [f\"responder_{i}\" for i in range(9)]})\n            \n        )\n        \n        # join test with label\n        test = test.join(lags2label, on=[\"symbol_id\", \"date_id\", \"time_id\"], how=\"left\")\n        # join test with lags\n        # self.lags = (\n        #     self.lags\n        #     .with_columns(\n        #         pl.lit(self.max_date_id, dtype=pl.Int16).alias(\"date_id\")\n        #     )\n        # )\n        # test = test.join(self.lags, on=[\"symbol_id\", \"date_id\", \"time_id\"], how=\"left\")\n        \n        return test.select(self.train_data.columns)\n    \n\n    def predict(self, \n                test: pl.DataFrame,\n                lags: pl.DataFrame | None) -> pl.DataFrame | pd.DataFrame:\n        if lags is not None:\n            \n            self.count_date += 1\n            \n            if not self.is_zero_day:\n                \n                print(self.min_date_id)\n                \n                # ============================================== 更新数据 ====================================\n                start1 = datetime.now()\n                \n                # lags 和 lags2label 都在第 1 天才开始处理（date 从第 0 天开始）\n                # lags 要在前一天保存，后一天统一处理，处理完之后，新的 lags 覆盖旧的 lags\n                self.update_train_data(lags)\n\n                end1 = datetime.now()\n                print(f\"更新数据用时：{end1 - start1}秒！\")\n                # ===========================================================================================\n                \n                # ============================================== 再训练 =======================================\n                if self.count_date % 10 == 0:\n                    start2 = datetime.now()\n                    \n                    self.retrain = True\n                    self.retrain_cbt()\n\n                    end2 = datetime.now()\n                    print(f\"再训练用时：{end2 - start2}秒！\")\n                # ============================================================================================\n                \n                \n                \n            self.is_zero_day = False # 他在第 0 天之后都是 False\n            self.max_date_id += 1   # 用来替换新数据的日期\n            self.min_date_id += 1   # 用来过滤旧数据：增加一天的新数据，删除一天的旧数据\n            \n            # 前一天的数据更新完之后，保存 lags，用于下次更新\n            self.lags = None # 为了释放旧的 lags 内存\n            gc.collect() \n            self.lags = lags\n        else:\n            self.retrain = False\n\n        # ============================================== 特征工程 ====================================\n        # TODO 用到的话再把这部分逻辑写上\n        # ===========================================================================================\n\n        \n        # ============================================== cbt预测 =====================================\n        X_test = test.select(self.config.features).to_numpy()\n        cbt_y_test_pred = self.cbt_model.predict(X_test)\n        # print(y_test_pred)\n        # ===========================================================================================\n        \n        # ========================================= 保存 test data ===================================\n        # test = test.with_columns(pl.Series(cbt_y_test_pred).alias(\"pred\"))\n        # 将 test 数据添加到 new_day_test 列表中，并在 next_date 的 time_0 进行合并，然后将被重置为空列表\n        selected_test = test.select(pl.all().exclude(\"row_id\", \"is_scored\"))\n        self.new_day_test.append(selected_test)\n        # ===========================================================================================\n\n        # ============================================== nn预测 ======================================\n        # TODO\n        # ===========================================================================================\n        \n        # ============================================== 交卷 ========================================\n        y_test_pred = cbt_y_test_pred * 1.0 \n        \n        predictions = (\n            test\n            .select(\n                \"row_id\",\n                pl.lit(0.0).alias(\"responder_6\")\n            )\n            .with_columns(pl.Series(y_test_pred).alias(\"responder_6\"))\n        )\n        assert isinstance(predictions, pl.DataFrame | pd.DataFrame)\n        assert list(predictions.columns) == ['row_id', 'responder_6']\n        assert len(predictions) == len(test)\n\n        if self.retrain:\n            print(f\"总用时：{datetime.now() - start1}秒！\")\n\n        return predictions\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-21T04:59:13.942896Z","iopub.execute_input":"2024-12-21T04:59:13.943208Z","iopub.status.idle":"2024-12-21T04:59:13.961743Z","shell.execute_reply.started":"2024-12-21T04:59:13.943174Z","shell.execute_reply":"2024-12-21T04:59:13.960877Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"lazy_train_data = (\n    pl.scan_parquet(\n        r\"/kaggle/input/jane-street-real-time-market-data-forecasting/train.parquet\"\n    )\n    .filter(pl.col(\"date_id\").ge(1248))\n    # .sort(by=[\n    #     \"symbol_id\", \n    #     \"time_id\", \n    #     \"date_id\"\n    # ])\n    # .with_columns([\n    #     pl.col(f\"responder_{i}\")\n    #     .shift()\n    #     .over([\n    #         \"symbol_id\", \n    #         \"time_id\",\n    #         ])\n    #     .alias(f\"responder_{i}_lag_1\")\n    #     for i in range(9)\n    # ])\n    \n    )\ntrain_data = lazy_train_data.collect()\n\n# train_data = (\n#     train_data\n#     .with_columns(\n#         pl.lit(0.0).alias(\"pred\")\n#     )\n# )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-21T04:59:13.963384Z","iopub.execute_input":"2024-12-21T04:59:13.963666Z","iopub.status.idle":"2024-12-21T04:59:29.663962Z","shell.execute_reply.started":"2024-12-21T04:59:13.963642Z","shell.execute_reply":"2024-12-21T04:59:29.663276Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"models = joblib.load(\n    r\"/kaggle/input/2024-12-21_12_57_cbt_0.013736_1248/scikitlearn/default/1/2024-12-21_12_57_CBT_0.013736_1248.pkl\"\n)\ncbt_model = models[0]\njs_predictor = JaneStreetCbtPredictor(config, cbt_model, train_data.drop(\"partition_id\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-21T04:59:29.665184Z","iopub.execute_input":"2024-12-21T04:59:29.665850Z","iopub.status.idle":"2024-12-21T04:59:33.719416Z","shell.execute_reply.started":"2024-12-21T04:59:29.665810Z","shell.execute_reply":"2024-12-21T04:59:33.718727Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"inference_server = kaggle_evaluation.jane_street_inference_server.JSInferenceServer(js_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            # noline train 测试用例\n            # \"/kaggle/input/js24-create-simulate-data/test.parquet\",\n            # \"/kaggle/input/js24-create-simulate-data/lags.parquet\",\n        )\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-21T04:59:33.720364Z","iopub.execute_input":"2024-12-21T04:59:33.720602Z","execution_failed":"2024-12-21T05:06:34.575Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}