{"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":"from 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.2  # サンプリング割合\nmodel = None  # モデルを初期化\nlags_ = None  # 遅延データの保存用\ntrain_features = None  # 学習時の特徴量を保存\n\n# ---------------------------------\n# 2. 高速フーリエ変換 (FFT) 関数\n# ---------------------------------\ndef apply_fft(data):\n    # 数値カラムを選択\n    numeric_columns = [col for col in data.columns if data[col].dtype in [pl.Float32, pl.Float64]]\n\n    if not numeric_columns:\n        print(\"No numeric columns found for FFT. Skipping.\")\n        return pl.DataFrame()  # 空のデータフレームを返す\n\n    fft_features = {}\n    for col in numeric_columns:\n        #print(f\"Applying FFT to column: {col}\")\n        # フーリエ変換を実行\n        fft_result = fft(data[col].to_numpy())\n        n = len(fft_result) // 2  # 正の周波数成分のみに限定\n        fft_features[f\"{col}_fft_abs\"] = np.abs(fft_result[:n])  # 振幅\n        fft_features[f\"{col}_fft_angle\"] = np.angle(fft_result[:n])  # 位相\n\n    fft_df = pd.DataFrame(fft_features)\n    print(\"FFT Columns:\", fft_df.columns)\n    return pl.from_pandas(fft_df)\n\n# ---------------------------------\n# 3. モデルの訓練\n# ---------------------------------\ndef train_model():\n    global model, train_features\n\n    # 1. データ読み込み (サンプリング)\n    print(\"Loading data (with sampling)...\")\n    train_files = [f\"{train_path}/partition_id={i}/part-0.parquet\" for i in range(10)]\n    sampled_batches = []\n    \n    for i, file in enumerate(train_files, 1):\n        print(f\"Processing file {i}/{len(train_files)}: {file}\")\n        data_batch = pl.read_parquet(file)\n        sampled_data = data_batch.sample(fraction=sample_fraction, seed=42)\n\n        # FFT 適用\n        fft_data = apply_fft(sampled_data)\n        if not fft_data.is_empty():\n            sampled_data = pl.concat([sampled_data, fft_data], how=\"horizontal\")\n        else:\n            print(\"No FFT features added for this batch.\")\n\n        sampled_batches.append(sampled_data)\n    \n    train_data = pl.concat(sampled_batches)\n    \n    # 他のデータ読み込み\n    lags_data = pl.read_parquet(lags_path)\n    test_data = pl.read_parquet(test_path)  # test_dataをここで読み込む\n    \n    # 2. 前処理\n    # 学習用データの特徴量を取得\n    exclude_columns = ['responder_6', 'date_id', 'time_id', 'symbol_id']\n    train_features = [col for col in train_data.columns if col not in exclude_columns]\n    \n    # 学習データの準備\n    X = train_data.select(train_features).to_pandas()\n    y = train_data['responder_6'].to_numpy()\n    \n    # テストデータの準備\n    test_features = [col for col in train_features if col in test_data.columns]\n    X = X[test_features]  # 学習データをテストデータのカラム順に合わせる\n    \n    # 3. 学習\n    print(\"Training model...\")\n    \n    # 訓練データと検証データの分割\n    X_train, X_valid, y_train, y_valid = train_test_split(X, y, test_size=0.2, random_state=42)\n\n    # モデルの選択と学習\n    from sklearn.model_selection import GridSearchCV\n    \n    param_grid = {\n        'max_iter': [100, 200],\n        'learning_rate': [0.01, 0.1],\n        'max_depth': [3, 5, 7],\n    }\n    \n    grid_search = GridSearchCV(\n        HistGradientBoostingRegressor(random_state=42),\n        param_grid,\n        scoring='neg_mean_squared_error',\n        cv=3\n    )\n    grid_search.fit(X_train, y_train)\n    model = grid_search.best_estimator_ \n    \n    # 検証データで評価\n    y_pred_valid = model.predict(X_valid)\n    rmse = np.sqrt(mean_squared_error(y_valid, y_pred_valid))\n    print(f\"Validation RMSE: {rmse}\")\n\n# ---------------------------------\n# 4. 予測関数\n# ---------------------------------\ndef predict(test: pl.DataFrame, lags: pl.DataFrame | None = None) -> pl.DataFrame:\n    global lags_, train_features\n\n    # 特徴量抽出（学習時の特徴量に一致させる）\n    columns = test.columns  # テストデータのカラムリスト\n    features = [col for col in train_features if col in 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\n    return result\n\n\n# ---------------------------------\n# 5. サーバーの設定と起動\n# ---------------------------------\nif __name__ == \"__main__\":\n    inference_server = kaggle_evaluation.jane_street_inference_server.JSInferenceServer(predict)\n    \n    # 環境に応じて起動\n    if os.getenv(\"KAGGLE_IS_COMPETITION_RERUN\"):\n        train_model()  # モデルを訓練\n        inference_server.serve()  # サーバー起動\n    else:\n        train_model()  # モデルを訓練\n        inference_server.run_local_gateway(\n            (\n                test_path,\n                lags_path,\n            )\n        )","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2024-12-14T16:54:42.889760Z","iopub.execute_input":"2024-12-14T16:54:42.890289Z","iopub.status.idle":"2024-12-14T16:54:57.122092Z","shell.execute_reply.started":"2024-12-14T16:54:42.890229Z","shell.execute_reply":"2024-12-14T16:54:57.121054Z"}},"outputs":[],"execution_count":null}]}