{"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"}],"dockerImageVersionId":30804,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"\"\"\"\nfrom sklearn.experimental import enable_hist_gradient_boosting  # 必要なら有効化\nfrom sklearn.ensemble import HistGradientBoostingRegressor\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import mean_squared_error\nimport pandas as pd\nimport polars as pl\nimport numpy as np\nimport os\nimport kaggle_evaluation.jane_street_inference_server\nfrom scipy.fft import fft\n\n# ---------------------------------\n# 1. パス設定とグローバル変数\n# ---------------------------------\ntrain_path = '/kaggle/input/jane-street-real-time-market-data-forecasting/train.parquet'\nlags_path = '/kaggle/input/jane-street-real-time-market-data-forecasting/lags.parquet'\ntest_path = '/kaggle/input/jane-street-real-time-market-data-forecasting/test.parquet'\nsample_submission_path = '/kaggle/input/jane-street-real-time-market-data-forecasting/sample_submission.csv'\n\nsample_fraction = 0.1  # サンプリング割合\nmodel = None  # モデルを初期化\nlags_ = None  # 遅延データの保存用\ntrain_features = None  # 学習時の特徴量を保存\nbatch_cache = None  # バッチデータの蓄積用\nbatch_count = 0  # バッチ処理カウント\n\ndef check_data():\n    print(\"\\nChecking Train Data:\")\n    train_data = pl.read_parquet(train_path)\n    print(\"Train Data Columns:\", train_data.columns)\n    print(\"Train Data Sample:\")\n    print(train_data.head(5))\n    print(train_data.tail(5))\n\n    print(\"\\nChecking Test Data:\")\n    test_data = pl.read_parquet(test_path)\n    print(\"Test Data Columns:\", test_data.columns)\n    print(\"Test Data Sample:\")\n    print(test_data.head(5))\n    print(test_data.tail(5))\n\n    print(\"\\nChecking Lags Data:\")\n    lags_data = pl.read_parquet(lags_path)\n    print(\"Lags Data Columns:\", lags_data.columns)\n    print(\"Lags Data Sample:\")\n    print(lags_data.head(5))\n    print(lags_data.tail(5))\n    \ncheck_data()\n\"\"\"","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2024-12-25T03:21:16.489553Z","iopub.execute_input":"2024-12-25T03:21:16.490131Z","iopub.status.idle":"2024-12-25T03:22:04.053945Z","shell.execute_reply.started":"2024-12-25T03:21:16.490044Z","shell.execute_reply":"2024-12-25T03:22:04.052856Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from catboost import CatBoostRegressor, Pool\nimport polars as pl\nimport numpy as np\nimport os\nfrom sklearn.metrics import mean_squared_error\nimport kaggle_evaluation.jane_street_inference_server\n\n# ---------------------------------\n# 1. パス設定\n# ---------------------------------\ntrain_path = '/kaggle/input/jane-street-real-time-market-data-forecasting/train.parquet'\ntest_path = '/kaggle/input/jane-street-real-time-market-data-forecasting/test.parquet'\nlags_path = '/kaggle/input/jane-street-real-time-market-data-forecasting/lags.parquet'\n\nmodel = None  # モデルを初期化\ntrain_features = None  # 学習時の特徴量を保存\n\n# ---------------------------------\n# 2. モデルの学習\n# ---------------------------------\nfrom concurrent.futures import ThreadPoolExecutor\n\ndef train_model(initial_training=True, incremental_data=None, days_to_use=10):\n    \"\"\"\n    初回学習またはオンライン学習を実行する関数。\n\n    :param initial_training: 初回学習かどうか (True: 初回学習, False: オンライン学習)\n    :param incremental_data: オンライン学習に使用する追加データ (polars DataFrame)\n    :param days_to_use: 最新の日付から使用する日数 (int)\n    \"\"\"\n    global model, train_features\n\n    if initial_training:\n        print(\"Initial training started...\")\n\n        # データ読み込み\n        train_files = [os.path.join(train_path, f\"partition_id={i}/part-0.parquet\") for i in range(10)]\n        \n        # 並列でファイルを読み込む\n        print(\"Reading files in parallel...\")\n        def read_file(file):\n            return pl.read_parquet(file)\n        \n        with ThreadPoolExecutor() as executor:\n            train_data_batches = list(executor.map(read_file, [file for file in train_files if os.path.exists(file)]))\n        train_data = pl.concat(train_data_batches)\n\n        # 最新の date_id を取得\n        latest_date_id = train_data.select(\"date_id\").max()[0, 0]\n        print(f\"Using data from the latest {days_to_use} days (latest date_id: {latest_date_id})\")\n\n        # 最新の日付から指定日数分のデータを抽出\n        train_data_latest = train_data.filter(pl.col(\"date_id\") >= (latest_date_id - days_to_use + 1))\n\n        # テストデータのカラムと一致する特徴量のみを使用\n        test_data = pl.read_parquet(test_path)\n        test_columns = set(test_data.columns)\n        exclude_columns = ['responder_6', 'date_id', 'time_id', 'symbol_id']\n        train_features = [col for col in train_data_latest.columns if col not in exclude_columns and col in test_columns]\n\n        # 学習データ準備\n        X = train_data_latest.select(train_features).to_pandas()\n        y = train_data_latest['responder_6'].to_numpy()\n\n        # CatBoost モデルの初回学習\n        print(\"Training CatBoost model...\")\n        model = CatBoostRegressor(\n            iterations=500,\n            learning_rate=0.1,\n            depth=6,\n            loss_function='RMSE',\n            verbose=100\n        )\n\n        train_pool = Pool(X, y)\n        model.fit(train_pool)\n\n    else:\n        print(\"Online training started...\")\n\n        # 増分学習用データ準備\n        if incremental_data is None:\n            raise ValueError(\"No incremental data provided for online training.\")\n        \n        incremental_X = incremental_data.select(train_features).to_pandas()\n        incremental_y = incremental_data['responder_6'].to_numpy()\n\n        incremental_pool = Pool(incremental_X, incremental_y)\n\n        # 既存モデルに対して追加学習\n        model.fit(\n            incremental_pool,\n            init_model=model,\n            use_best_model=False,\n            verbose=100\n        )\n\n    print(\"Training complete.\")\n\n# ---------------------------------\n# 3. 予測関数\n# ---------------------------------\ndef predict(test: pl.DataFrame, lags: pl.DataFrame | None = None) -> pl.DataFrame:\n    global model, train_features\n\n    # 特徴量をテストデータに合わせる\n    features = [col for col in train_features if col in test.columns]\n    test_X = test.select(features).to_pandas()\n\n    # 予測実行\n    predictions = model.predict(test_X)\n\n    # 結果を整形\n    result = test.select(['row_id']).with_columns(\n        pl.Series('responder_6', predictions)\n    )\n    return result\n\n# ---------------------------------\n# 4. サーバーの設定と起動\n# ---------------------------------\nif __name__ == \"__main__\":\n    inference_server = kaggle_evaluation.jane_street_inference_server.JSInferenceServer(predict)\n\n    if os.getenv(\"KAGGLE_IS_COMPETITION_RERUN\"):\n        print(\"Kaggle competition mode detected.\")\n        train_model(initial_training=True, days_to_use=100)  # 初回学習 (最新10日分を使用)\n        inference_server.serve()  # サーバー起動\n    else:\n        print(\"Local gateway mode detected.\")\n        train_model(initial_training=True, days_to_use=100)  # 初回学習 (最新10日分を使用)\n\n        # オンライン学習用のデータ（例として一部サンプリング）\n        incremental_data = pl.read_parquet(train_path).sample(fraction=0.1, seed=42)\n        train_model(initial_training=False, incremental_data=incremental_data)  # オンライン学習\n\n        inference_server.run_local_gateway(\n            (\n                test_path,\n                lags_path,\n            )\n        )\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-25T03:44:06.342500Z","iopub.execute_input":"2024-12-25T03:44:06.344582Z","iopub.status.idle":"2024-12-25T03:52:55.829250Z","shell.execute_reply.started":"2024-12-25T03:44:06.344046Z","shell.execute_reply":"2024-12-25T03:52:55.828152Z"}},"outputs":[],"execution_count":null}]}