{"cells":[{"cell_type":"code","execution_count":4,"metadata":{},"outputs":[],"source":"import os\nimport polars as pl\nimport sys\nfrom pathlib import Path\n\n# Set the path to the src directory (adjust the path if necessary)\nsrc_path = Path(':~/Projects/Jane Street Real-Time Market Data Forecasting/jane_street_market/src')\nsys.path.append(str(src_path))\nimport kaggle_evaluation.jane_street_inference_server\nfrom src.models.lgbm_model import LGBMModel  # Ensure this import path is correct\n\n# Global model and lags\nmodel = None\nlags_ : pl.DataFrame | None = None\n\ndef predict(test: pl.DataFrame, lags: pl.DataFrame | None) -> pl.DataFrame:\n    global model, lags_\n\n    # Load model on the first call\n    if model is None:\n        model_path = 'models/20241028_180057/model.pkl'\n        model = LGBMModel({}, is_loading=True)\n        model.load(model_path)\n\n    # Save lags\n    if lags is not None:\n        lags_ = lags\n\n    # Prediction logic\n    feature_cols = [f'feature_{i:02d}' for i in range(79)]\n    X = test.select(feature_cols).to_numpy()\n    predictions = model.predict(X)\n\n    # Format predictions\n    predictions_df = pl.DataFrame({\n        'row_id': test.get_column('row_id'),\n        'responder_6': predictions\n    })\n    return predictions_df\n\ndef main():\n    # Initialize the inference server\n    inference_server = kaggle_evaluation.jane_street_inference_server.JSInferenceServer(predict)\n\n    # Check for competition rerun or local testing\n    if os.getenv('KAGGLE_IS_COMPETITION_RERUN'):\n        inference_server.serve()\n    else:\n        test = pl.read_parquet('/kaggle/input/jane-street-real-time-market-data-forecasting/test.parquet')\n        lags = pl.read_parquet('/kaggle/input/jane-street-real-time-market-data-forecasting/lags.parquet')\n        \n        # Generate predictions\n        predictions = predict(test, lags)\n        \n        # Save as 'submission.parquet'\n        predictions.write_parquet('submission.parquet')\n\nif __name__ == \"__main__\":\n    main()\n"}],"metadata":{"kernelspec":{"display_name":".venv","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.10.12"}},"nbformat":4,"nbformat_minor":2}