{"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":10368326,"sourceType":"datasetVersion","datasetId":6345740,"isSourceIdPinned":true},{"sourceId":215951136,"sourceType":"kernelVersion"}],"dockerImageVersionId":30823,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"from pathlib import Path\nimport json\nimport xgboost as xgb\nimport lightgbm as lgb\nimport catboost as cb \nimport os\nimport numpy as np \nimport multiprocessing\nos.environ['OMP_NUM_THREADS'] = str(multiprocessing.cpu_count())\nimport pandas as pd\nimport polars as pl \nimport json \nfrom js_utils import *\nfrom tqdm.auto import tqdm\nimport torch\ntorch.backends.cudnn.benchmark = True\ntorch.set_float32_matmul_precision('medium')\n\nimport sys \nsys.path.append(str(Path.cwd().parent / \"data/raw\"))\nimport kaggle_evaluation.jane_street_inference_server","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-01-04T11:12:40.573597Z","iopub.execute_input":"2025-01-04T11:12:40.573891Z","iopub.status.idle":"2025-01-04T11:12:40.579650Z","shell.execute_reply.started":"2025-01-04T11:12:40.573869Z","shell.execute_reply":"2025-01-04T11:12:40.578692Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"args = {\n    'ckpt_path': '250103_091313_933085_nntree6fold', \n    'inference_folds': [0, 1, 2, 3, 4, 5],\n    'inference_fold_weights': {0: 1.0, 1: 1.5, 2: 1.4, 3: 1.3, 4: 1.2, 5: 1.1}, \n    'model_types': ['cb'], \n    'override_blend_coeffs': {'cb': 1}, \n    'check_preds': False, \n}\n# CB-RFE 6-fold\n\ndevice = 'cuda'\nbatch_size = 2048\nfeature_cols = raw_feature_cols()\nresponder = 'responder_6'\nmodel_dir = Path(f'/kaggle/input/js2024-student-checkpoints') / args['ckpt_path']\n\nwith open(model_dir / 'blend_metadata.json', 'r') as f:\n    blend_metadata = json.load(f)\nmodel_metadata = blend_metadata['model_metadata']\n\nfold_weights = {k: v / sum(args['inference_fold_weights'].values()) for k, v in args['inference_fold_weights'].items()}\nblend_coeffs = blend_metadata['blend_coeffs']\n# if args['override_model_blend'] is not None:\n#     args['override_model_blend'] = {\n#         k: v / sum(args['override_model_blend'].values()) for k, v in args['override_model_blend'].items()}\nprint('Weight assignment for folds:', fold_weights)\n# print('Override blending weights:', args['override_model_blend'])\nprint('Uniform blend r2:', blend_metadata['blend_r2']) # , 'by-symbol blend r2:', blend_metadata['symbol_blend_r2'])\nprint('Models r2', {model_type: blend_metadata['model_metadata'][model_type]['fold0']['r2'] for model_type in args['model_types']})\nblend_coeffs = {'cb': blend_coeffs[0], 'xgb': blend_coeffs[1], 'lgb': blend_coeffs[2], \n                'nn': blend_coeffs[3]} \nif args['override_blend_coeffs'] is not None: \n    blend_coeffs = args['override_blend_coeffs']\nprint('Blend coeffs', blend_coeffs)\nprint(blend_metadata['blend_args'])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-04T11:12:40.580782Z","iopub.execute_input":"2025-01-04T11:12:40.581089Z","iopub.status.idle":"2025-01-04T11:12:40.615591Z","shell.execute_reply.started":"2025-01-04T11:12:40.581060Z","shell.execute_reply":"2025-01-04T11:12:40.615003Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# # Plot the oof performances \n# import plotly.graph_objects as go\n# fig = go.Figure()\n# for model_type, fold_data in model_metadata.items():\n#     fig.add_trace(go.Bar(\n#         x=args['inference_folds'],\n#         y=[fold_data[f'fold{fold}']['r2'] for fold in args['inference_folds']],\n#         name=model_type,\n#         text=[f\"{fold_data[f'fold{fold}']['r2']:.5f}\" for fold in args['inference_folds']],\n#         textposition='auto'\n#     ))\n\n# fig.update_layout(\n#     barmode='group',\n#     xaxis_title='Fold ID',\n#     yaxis_title='Validation R2',\n#     title='Validation R2 by Fold and Model Type'\n# )\n# fig.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-04T11:12:40.617152Z","iopub.execute_input":"2025-01-04T11:12:40.617351Z","iopub.status.idle":"2025-01-04T11:12:40.620644Z","shell.execute_reply.started":"2025-01-04T11:12:40.617334Z","shell.execute_reply":"2025-01-04T11:12:40.619728Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_nntreeblendCV_models(args):\n    model_dir = Path(f'/kaggle/input/js2024-student-checkpoints') / args['ckpt_path']\n    models = {}\n    df_train = pl.scan_parquet(\n        '/kaggle/input/jane-street-real-time-market-data-forecasting/train.parquet'\n    )[:2].collect()\n    for model_type in args['model_types']: \n        models[model_type] = {}\n        for fold in args['inference_folds']:\n            if model_type == 'cb': \n                model = cb.CatBoost()\n                model.load_model(str(model_dir / model_type / f'fold{fold}' / 'cb_model.cbm'))\n                trainer = CatBoostTrainer(\n                    df_train[:2], model_metadata[model_type]['fold0']['input_args']['feature_cols'], \n                    ['responder_6'], device='cpu')\n                trainer.update_cpu_model(model)\n            elif model_type == 'xgb': \n                xgb_model = xgb.Booster()\n                xgb_model.load_model(str(model_dir / model_type / f'fold{fold}' / \"xgb_model.bin\"))\n                trainer = XGBoostTrainer(\n                    df_train[:2], model_metadata[model_type]['fold0']['input_args']['feature_cols'], \n                    fit_responders=['responder_6'], device='cpu')\n                trainer.update_cpu_model(xgb_model)\n            elif model_type == 'lgb':\n                lgb_model = lgb.Booster(model_file=str(model_dir / model_type / f'fold{fold}' / \"lgb_model.txt\"))\n                trainer = LGBTrainer(\n                    df_train[:2], model_metadata[model_type]['fold0']['input_args']['feature_cols'], \n                    device='cpu', fit_responders=['responder_6'])\n                trainer.update_cpu_model(lgb_model)\n            elif model_type == 'nn': \n                nn_model = JSModel(model_metadata['nn']['fold0']['input_args']).to(device) \n                nn_model.eval()\n                if fold == -1:\n                    nn_model.load_state_dict(\n                        torch.load(model_dir / model_type / f'total' / \"nn_model.pth\", weights_only=False))\n                else:\n                    nn_model.load_state_dict(\n                        torch.load(model_dir / model_type / f'fold{fold}' / \"nn_model.pth\", weights_only=False))\n                trainer = NNTrainer(datamodule=None)\n                trainer.model = nn_model\n            models[model_type][fold] = trainer \n    return models \n\nmodels = load_nntreeblendCV_models(args)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-04T11:12:40.621647Z","iopub.execute_input":"2025-01-04T11:12:40.621951Z","iopub.status.idle":"2025-01-04T11:12:56.875701Z","shell.execute_reply.started":"2025-01-04T11:12:40.621904Z","shell.execute_reply":"2025-01-04T11:12:56.874874Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if args['check_preds']:\n    oof_preds = {model_type: pl.scan_parquet(model_dir / model_type / 'fold0' / 'oof_preds.pq').limit(100).collect()\n                 for model_type in tqdm(args['model_types'])}\n    \n    oof_df = pl.scan_parquet('/kaggle/input/jane-street-real-time-market-data-forecasting/train.parquet') \\\n        .filter(pl.col('partition_id') == 9).limit(100).collect()\n    \n    for model_type in args['model_types']: \n        print('Checking', model_type)\n        recorded_preds = np.array(oof_preds[model_type]['pred'])\n        if model_type in ['lgb', 'cb', 'xgb']:\n            model_preds = np.array(models[model_type][0].predict(oof_df)['pred'])\n        elif model_type == 'nn': \n            model_preds = np.array(models['nn'][0].predict(\n                        oof_df, feature_cols=blend_metadata['model_metadata']['nn']['fold0']['input_args']['feature_cols'], \n                        device=device, batch_size=batch_size)['pred'])\n        print(np.abs(recorded_preds - model_preds).mean())\n        assert np.abs(recorded_preds - model_preds).mean() < 1e-3","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-04T11:12:56.876379Z","iopub.execute_input":"2025-01-04T11:12:56.876593Z","iopub.status.idle":"2025-01-04T11:12:56.882465Z","shell.execute_reply.started":"2025-01-04T11:12:56.876574Z","shell.execute_reply":"2025-01-04T11:12:56.881705Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# total_weight = (np.array(blend_metadata['coeffs_bysymbol']['cb']) + \n#         np.array(blend_metadata['coeffs_bysymbol']['lgb']) + \n#         np.array(blend_metadata['coeffs_bysymbol']['xgb']) + \n#         np.array(blend_metadata['coeffs_bysymbol']['nn']))\n\n# dual_plot(\n#     y1=blend_metadata['coeffs_bysymbol'],\n#     y2={'total': total_weight}, x=raw_symbol_ids(),\n#     title='Model Coefficients by Symbol'\n# ).show()\n\n# def nntreeblend_predict(pred_df, blend_metadata, args):\n#     \"\"\"Blend predictions using per-symbol coefficients.\n    \n#     Args:\n#         coeffs_bysymbol: Dict with keys 'cb', 'xgb', 'lgb' containing per-symbol coefficients\n#         treeblend_preds: Polars DataFrame with columns 'symbol_id', 'cb_pred', 'xgb_pred', 'lgb_pred'\n    \n#     Returns:\n#         Polars DataFrame with blended predictions in 'pred' column\n#     \"\"\"\n#     # Initialize predictions column \n#     pred_df = pred_df.with_columns(pl.lit(0.0).alias('pred'))\n#     coeffs_bysymbol = blend_metadata['coeffs_bysymbol']\n#     model_types = args['model_types']\n#     # Blend predictions by symbol\n#     for symbol in raw_symbol_ids():\n#         # Get coefficients for this symbol\n#         coefs = {model_type: coeffs_bysymbol[model_type][symbol] \n#                  for model_type in model_types} if args['override_model_blend'] is None else args['override_model_blend']\n        \n#         # Calculate blended predictions using when-then-otherwise\n#         pred_df = pred_df.with_columns(\n#             pl.when(pl.col('symbol_id') == symbol)\n#             .then(sum(pl.col(model_type) * coefs[model_type]\n#                     for model_type in model_types))\n#             .otherwise(pl.col('pred'))\n#             .alias('pred'))\n#     return pred_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-04T11:12:56.883247Z","iopub.execute_input":"2025-01-04T11:12:56.883441Z","iopub.status.idle":"2025-01-04T11:12:56.905102Z","shell.execute_reply.started":"2025-01-04T11:12:56.883424Z","shell.execute_reply":"2025-01-04T11:12:56.904354Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def predict(test, lags):\n    predictions = test.select(\n        'row_id',\n        pl.lit(0.0).alias('responder_6'))\n    test_df = test.clone()\n    # Get predictions from each model type and fold\n    model_preds = {}\n    for model_type in args['model_types']:\n        fold_preds, inferred_folds = [], []\n        for f in args['inference_folds']:\n            # Get predictions based on model type\n            if model_type == 'nn': \n                output = models[model_type][f].predict(\n                    test_df, \n                    feature_cols=blend_metadata['model_metadata']['nn']['fold0']['input_args']['feature_cols'], \n                    device=device, \n                    batch_size=batch_size\n                )['pred'].clip(-5, 5)\n            else:\n                output = models[model_type][f].predict(test_df)['pred'].clip(-5, 5)\n            fold_preds.append(output)\n            inferred_folds.append(f)\n        # Average predictions across folds using weights\n        model_preds[model_type] = np.average(\n            np.stack(fold_preds, axis=1), \n            axis=1, \n            weights=[fold_weights[f] for f in inferred_folds])\n    # Compute final weighted prediction\n    final_pred = sum(\n        model_preds[model_type] * blend_coeffs[model_type]\n        for model_type in args['model_types'])\n    # Clip predictions and update result\n    final_pred = np.clip(final_pred, -5, 5)\n    predictions = predictions.with_columns(\n        pl.Series('responder_6', final_pred))\n    return predictions","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-04T11:12:56.906336Z","iopub.execute_input":"2025-01-04T11:12:56.906532Z","iopub.status.idle":"2025-01-04T11:12:56.923766Z","shell.execute_reply.started":"2025-01-04T11:12:56.906515Z","shell.execute_reply":"2025-01-04T11:12:56.923115Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# test = pl.scan_parquet(\n#     '/kaggle/input/jane-street-real-time-market-data-forecasting/test.parquet'\n# ).collect()\n# lags = pl.scan_parquet(\n#     '/kaggle/input/jane-street-real-time-market-data-forecasting/lags.parquet'\n# ).collect()\n# out = predict(test, lags)\n# out\n# # print(blend_metadata['blend_r2'], blend_metadata['symbol_blend_r2'])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-04T11:12:56.924492Z","iopub.execute_input":"2025-01-04T11:12:56.924669Z","iopub.status.idle":"2025-01-04T11:12:56.941746Z","shell.execute_reply.started":"2025-01-04T11:12:56.924653Z","shell.execute_reply":"2025-01-04T11:12:56.941014Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"inference_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    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-04T11:12:56.942553Z","iopub.execute_input":"2025-01-04T11:12:56.942762Z","iopub.status.idle":"2025-01-04T11:12:57.184202Z","shell.execute_reply.started":"2025-01-04T11:12:56.942744Z","shell.execute_reply":"2025-01-04T11:12:57.183542Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}