{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"colab":{"provenance":[],"toc_visible":true},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":105399,"databundleVersionId":12733338,"sourceType":"competition"}],"dockerImageVersionId":31236,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"widgets":{"application/vnd.jupyter.widget-state+json":{"017db7a134b24c6d80274ea639a5a545":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"HBoxModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HBoxModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HBoxView","box_style":"","children":["IPY_MODEL_bd02dc82112d45a5969ebff53ac629a5","IPY_MODEL_b184cc13c48e48d0b3dd9df36d82802b","IPY_MODEL_44339f93d47f4049ab0ee52ccb845195"],"layout":"IPY_MODEL_43fefedccc9e4f1195e6269c4eec4178"}},"05bca173a7fd46e2be9ab7a434b87ecc":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"ProgressStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"ProgressStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","bar_color":null,"description_width":""}},"074a155285504e568d2b878a3f80a77c":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"DescriptionStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"DescriptionStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","description_width":""}},"07ed8dff68784855b7b97b9220eb4e1b":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"HTMLModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HTMLModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HTMLView","description":"","description_tooltip":null,"layout":"IPY_MODEL_84738be119524c72a10ec04cf76ea8f9","placeholder":"​","style":"IPY_MODEL_5d4d85ea5af04c3991dcb2b134abcc65","value":" 100000/100000 [00:05&lt;00:00, 30391.18it/s]"}},"09c460c23c284e9c92a84e1c438348a2":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"0ead4c9562b84072be6eb1cae493fb5f":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"HTMLModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HTMLModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HTMLView","description":"","description_tooltip":null,"layout":"IPY_MODEL_99e691f478ca4f5e8480ccf6a8756421","placeholder":"​","style":"IPY_MODEL_791e20cfd0734e62a9666421dc0ac82e","value":" 10000/10000 [00:00&lt;00:00, 29066.21it/s]"}},"19f7c07df0dd4d459c580cda900f79d7":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"1d1f7700af1e42819389270ba2031488":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"HTMLModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HTMLModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HTMLView","description":"","description_tooltip":null,"layout":"IPY_MODEL_09c460c23c284e9c92a84e1c438348a2","placeholder":"​","style":"IPY_MODEL_074a155285504e568d2b878a3f80a77c","value":"Grouping: 100%"}},"248c0943af7f433a8f7b7e103e620651":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"335163839f3449ba881f5ea68c89ec40":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"3b5177a2534f445982b26940a624b642":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"HTMLModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HTMLModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HTMLView","description":"","description_tooltip":null,"layout":"IPY_MODEL_f3d15d9adbe049f4a83548f4a28ef3db","placeholder":"​","style":"IPY_MODEL_8c94fd3dca1142398c69b634d213499b","value":"Grouping: 100%"}},"40db6fbfac1943e4aaa34f8abcbc7394":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"ProgressStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"ProgressStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","bar_color":null,"description_width":""}},"4308ee9a67614aedb76ce70884f26877":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"FloatProgressModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"FloatProgressModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"ProgressView","bar_style":"success","description":"","description_tooltip":null,"layout":"IPY_MODEL_19f7c07df0dd4d459c580cda900f79d7","max":100000,"min":0,"orientation":"horizontal","style":"IPY_MODEL_d21be03ad52d42cdab28d6e65f0fa922","value":100000}},"43fefedccc9e4f1195e6269c4eec4178":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"44339f93d47f4049ab0ee52ccb845195":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"HTMLModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HTMLModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HTMLView","description":"","description_tooltip":null,"layout":"IPY_MODEL_335163839f3449ba881f5ea68c89ec40","placeholder":"​","style":"IPY_MODEL_bda7275958ac49e5b70ecf7041360253","value":" 82/82 [00:00&lt;00:00, 1034.26group/s]"}},"46e96bea47f847daa8b0171a1b390552":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"533ccd87676644a9b2601200e7539264":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"HTMLModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HTMLModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HTMLView","description":"","description_tooltip":null,"layout":"IPY_MODEL_a030eb86aa4a4a0a927eed6514a46a47","placeholder":"​","style":"IPY_MODEL_83fc1b2714704f74af21936b07941159","value":"Processing groups: 100%"}},"58d6a922b4e243379c956646891c1cb9":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"FloatProgressModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"FloatProgressModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"ProgressView","bar_style":"success","description":"","description_tooltip":null,"layout":"IPY_MODEL_248c0943af7f433a8f7b7e103e620651","max":10000,"min":0,"orientation":"horizontal","style":"IPY_MODEL_9dd582df387f410792aa4a0c6fdd6ba4","value":10000}},"5d4d85ea5af04c3991dcb2b134abcc65":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"DescriptionStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"DescriptionStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","description_width":""}},"5dac9ea202d341ada4e75945a2a5d321":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"HBoxModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HBoxModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HBoxView","box_style":"","children":["IPY_MODEL_1d1f7700af1e42819389270ba2031488","IPY_MODEL_58d6a922b4e243379c956646891c1cb9","IPY_MODEL_0ead4c9562b84072be6eb1cae493fb5f"],"layout":"IPY_MODEL_46e96bea47f847daa8b0171a1b390552"}},"6ca30c003a67418dacf4cf7279f3d9cb":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"DescriptionStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"DescriptionStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","description_width":""}},"791e20cfd0734e62a9666421dc0ac82e":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"DescriptionStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"DescriptionStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","description_width":""}},"83fc1b2714704f74af21936b07941159":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"DescriptionStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"DescriptionStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","description_width":""}},"84738be119524c72a10ec04cf76ea8f9":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"8c94fd3dca1142398c69b634d213499b":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"DescriptionStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"DescriptionStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","description_width":""}},"99e691f478ca4f5e8480ccf6a8756421":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"9dd582df387f410792aa4a0c6fdd6ba4":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"ProgressStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"ProgressStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","bar_color":null,"description_width":""}},"9df88f16c8e64e868dc9d99aa6194ca6":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"HTMLModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HTMLModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HTMLView","description":"","description_tooltip":null,"layout":"IPY_MODEL_bcc9f7f65a51452fbfd49e7a527bd382","placeholder":"​","style":"IPY_MODEL_6ca30c003a67418dacf4cf7279f3d9cb","value":" 883/883 [00:01&lt;00:00, 597.26group/s]"}},"a030eb86aa4a4a0a927eed6514a46a47":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"b11d8a31b7ad49fab80c395839b02f32":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"b184cc13c48e48d0b3dd9df36d82802b":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"FloatProgressModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"FloatProgressModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"ProgressView","bar_style":"success","description":"","description_tooltip":null,"layout":"IPY_MODEL_ecb9d959cbeb4acab26e5ec71c33c6ac","max":82,"min":0,"orientation":"horizontal","style":"IPY_MODEL_05bca173a7fd46e2be9ab7a434b87ecc","value":82}},"b30f9902b87846afad29c8a8327b281b":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"HBoxModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HBoxModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HBoxView","box_style":"","children":["IPY_MODEL_533ccd87676644a9b2601200e7539264","IPY_MODEL_be476d0395e54ab79cd25d07720a5c61","IPY_MODEL_9df88f16c8e64e868dc9d99aa6194ca6"],"layout":"IPY_MODEL_b11d8a31b7ad49fab80c395839b02f32"}},"bcc9f7f65a51452fbfd49e7a527bd382":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"bd02dc82112d45a5969ebff53ac629a5":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"HTMLModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HTMLModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HTMLView","description":"","description_tooltip":null,"layout":"IPY_MODEL_ee4f40b113c44da3a617c2db458be30e","placeholder":"​","style":"IPY_MODEL_d03a1a2a4e4646dab643844bb2834d5f","value":"Processing groups: 100%"}},"bda7275958ac49e5b70ecf7041360253":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"DescriptionStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"DescriptionStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","description_width":""}},"be476d0395e54ab79cd25d07720a5c61":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"FloatProgressModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"FloatProgressModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"ProgressView","bar_style":"success","description":"","description_tooltip":null,"layout":"IPY_MODEL_ea0db88b699f4b6f9d8b1ce3ef90985d","max":883,"min":0,"orientation":"horizontal","style":"IPY_MODEL_40db6fbfac1943e4aaa34f8abcbc7394","value":883}},"ccd45cccb5f34d2493c8c21b71f71719":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"HBoxModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HBoxModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HBoxView","box_style":"","children":["IPY_MODEL_3b5177a2534f445982b26940a624b642","IPY_MODEL_4308ee9a67614aedb76ce70884f26877","IPY_MODEL_07ed8dff68784855b7b97b9220eb4e1b"],"layout":"IPY_MODEL_e38185eff9184c3db8d2121a0dd6b67e"}},"d03a1a2a4e4646dab643844bb2834d5f":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"DescriptionStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"DescriptionStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","description_width":""}},"d21be03ad52d42cdab28d6e65f0fa922":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"ProgressStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"ProgressStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","bar_color":null,"description_width":""}},"e38185eff9184c3db8d2121a0dd6b67e":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"ea0db88b699f4b6f9d8b1ce3ef90985d":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"ecb9d959cbeb4acab26e5ec71c33c6ac":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"ee4f40b113c44da3a617c2db458be30e":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"f3d15d9adbe049f4a83548f4a28ef3db":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}}}}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Setup, Import","metadata":{}},{"cell_type":"code","source":"!pip install polars","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-10T14:34:34.363282Z","iopub.execute_input":"2026-01-10T14:34:34.364002Z","iopub.status.idle":"2026-01-10T14:34:37.432627Z","shell.execute_reply.started":"2026-01-10T14:34:34.363971Z","shell.execute_reply":"2026-01-10T14:34:37.431903Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import kagglehub\n# import os\n\n# path = kagglehub.competition_download(\"aeroclub-recsys-2025\")\n# print(\"✅ Dataset downloaded to:\", path)","metadata":{"id":"V1-piopMEHHz","outputId":"033a8ac2-4ca3-4d64-b6a1-dca02fc22584","trusted":true,"execution":{"iopub.status.busy":"2026-01-10T14:34:37.434736Z","iopub.execute_input":"2026-01-10T14:34:37.435044Z","iopub.status.idle":"2026-01-10T14:34:37.438421Z","shell.execute_reply.started":"2026-01-10T14:34:37.435017Z","shell.execute_reply":"2026-01-10T14:34:37.437817Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import polars as pl\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport time\nimport copy\nfrom huggingface_hub import hf_hub_download, login, create_repo, upload_file\nfrom pathlib import Path\n\nfrom collections import defaultdict\nfrom dataclasses import dataclass, field\nimport json\n\nimport numpy as np\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom sklearn.model_selection import train_test_split\nfrom collections import defaultdict\nimport numpy as np\nfrom sklearn.preprocessing import StandardScaler\nfrom tqdm.auto import tqdm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-10T14:34:37.439446Z","iopub.execute_input":"2026-01-10T14:34:37.439729Z","iopub.status.idle":"2026-01-10T14:34:37.455452Z","shell.execute_reply.started":"2026-01-10T14:34:37.439708Z","shell.execute_reply":"2026-01-10T14:34:37.454940Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Data Downloading","metadata":{}},{"cell_type":"code","source":"from kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\n\nlogin(user_secrets.get_secret(\"HF_TOKEN\"))\n\nrepo_id = \"dungdz123/flight-recsys\"\ncreate_repo(\n        repo_id=repo_id,\n        exist_ok=True,     # IMPORTANT: no error if already exists\n    )\n\ndef download_file(repo_id, filename, local_dir=\"\"):\n    local_path = hf_hub_download(repo_id=repo_id, filename=filename, local_dir=local_dir)\n    print(f\"Downloaded {filename} from {repo_id} to {local_path}\")\n    return local_path\n\nlocal_path = \"/kaggle/temp\"\n\ndownload_file(repo_id, \"train_samples.pt\", local_path)\n# download_file(repo_id, \"val_samples.pt\")\ndownload_file(repo_id, \"test_samples.pt\", local_path)\nvocab_path = download_file(repo_id, \"vocabularies.json\", local_path)\nlocal_path = Path(local_path)","metadata":{"id":"3hKhZsdmDxtL","outputId":"28ec1107-c4c1-4602-ba43-0756c7d9c720","trusted":true,"execution":{"iopub.status.busy":"2026-01-10T14:34:37.456369Z","iopub.execute_input":"2026-01-10T14:34:37.456605Z","iopub.status.idle":"2026-01-10T14:34:38.265672Z","shell.execute_reply.started":"2026-01-10T14:34:37.456576Z","shell.execute_reply":"2026-01-10T14:34:38.265080Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Feature Selection","metadata":{"id":"NmcPTUOvDxtO"}},{"cell_type":"code","source":"def load_vocabularies(vocab_path):\n    with open(vocab_path, \"r\", encoding=\"utf-8\") as f:\n        vocab_dict = json.load(f)\n\n    carrier_vocab2idx = vocab_dict[\"carrier_vocab2idx\"]\n    airport_vocab2idx = vocab_dict[\"airport_vocab2idx\"]\n    company_vocab2idx = vocab_dict[\"company_vocab2idx\"]\n    nationality_vocab2idx = vocab_dict[\"nationality_vocab2idx\"]\n\n    return (\n        carrier_vocab2idx,\n        airport_vocab2idx,\n        company_vocab2idx,\n        nationality_vocab2idx,\n    )\ncarrier_vocab2idx, airport_vocab2idx, company_vocab2idx, nationality_vocab2idx = load_vocabularies(vocab_path)\n\nprint(f\"\\nVocabulary sizes:\")\nprint(f\"  - Carriers: {len(carrier_vocab2idx)} (including <PAD>)\")\nprint(f\"  - Airports: {len(airport_vocab2idx)} (including <UNK>)\")\nprint(f\"  - Companies: {len(company_vocab2idx)} (including <UNK>)\")\nprint(f\"  - Nationalities: {len(nationality_vocab2idx)} (including <UNK>)\")\n\nvocab_info = {\n    'carrier': (carrier_vocab2idx, 16),\n    'airport': (airport_vocab2idx, 16),\n    'company': (company_vocab2idx, 8),\n    'nationality': (nationality_vocab2idx, 8),\n}","metadata":{"id":"IcjxjwrYDxtO","outputId":"57c7c714-d11e-473f-f82e-fa912c36fb77","trusted":true,"execution":{"iopub.status.busy":"2026-01-10T14:34:38.267280Z","iopub.execute_input":"2026-01-10T14:34:38.267867Z","iopub.status.idle":"2026-01-10T14:34:38.274644Z","shell.execute_reply.started":"2026-01-10T14:34:38.267844Z","shell.execute_reply":"2026-01-10T14:34:38.273955Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Dataset Preparation","metadata":{"id":"ahS3kplFDxtO"}},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import Dataset\n\nclass RankingDataset(Dataset):\n    def __init__(self, samples):\n        self.samples = samples\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        entry = self.samples[idx]\n\n        return entry","metadata":{"id":"MusM0QqhDxtP","trusted":true,"execution":{"iopub.status.busy":"2026-01-10T14:34:38.275555Z","iopub.execute_input":"2026-01-10T14:34:38.275901Z","iopub.status.idle":"2026-01-10T14:34:38.290017Z","shell.execute_reply.started":"2026-01-10T14:34:38.275871Z","shell.execute_reply":"2026-01-10T14:34:38.289519Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def fix_dtype_ds(ds):\n    \"\"\"\n    Fix dtypes for a dataset used with EmbeddingBag + numeric + time features.\n    \"\"\"\n\n    for i in range(len(ds)):\n        sample = ds[i]\n\n        # --------------------\n        # Numeric features\n        # --------------------\n        if \"num_features\" in sample:\n            if sample[\"num_features\"].dtype != torch.float32:\n                sample[\"num_features\"] = sample[\"num_features\"].float()\n\n        # --------------------\n        # Time features\n        # --------------------\n        if \"time_features\" in sample:\n            if sample[\"time_features\"].dtype != torch.float32:\n                sample[\"time_features\"] = sample[\"time_features\"].float()\n\n        # --------------------\n        # List features (EmbeddingBag)\n        # --------------------\n        if \"list_features\" in sample:\n            for key, v in sample[\"list_features\"].items():\n                # indices\n                if v[\"flat\"].dtype != torch.long:\n                    v[\"flat\"] = v[\"flat\"].long()\n\n                # offsets\n                if v[\"offsets\"].dtype != torch.long:\n                    v[\"offsets\"] = v[\"offsets\"].long()\n\n        # --------------------\n        # Categorical features\n        # --------------------\n        if \"cater_features\" in sample:\n            for key, t in sample[\"cater_features\"].items():\n                if t.dtype != torch.long:\n                    sample[\"cater_features\"][key] = t.long()\n\n        # --------------------\n        # Positive index\n        # --------------------\n        if \"positive_idx\" in sample:\n            if not isinstance(sample[\"positive_idx\"], int):\n                sample[\"positive_idx\"] = int(sample[\"positive_idx\"])\n\n    return ds\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-10T14:34:38.290729Z","iopub.execute_input":"2026-01-10T14:34:38.290950Z","iopub.status.idle":"2026-01-10T14:34:38.313810Z","shell.execute_reply.started":"2026-01-10T14:34:38.290926Z","shell.execute_reply":"2026-01-10T14:34:38.313325Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Loading Train Data\")\ntrain_samples = torch.load(local_path / \"train_samples.pt\", map_location=\"cpu\")\nprint(\"Loading Test Data\")\ntest_s = torch.load(local_path/ \"test_samples.pt\", map_location=\"cpu\")\nprint(\"Complete Loading\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-10T14:34:38.314538Z","iopub.execute_input":"2026-01-10T14:34:38.314802Z","iopub.status.idle":"2026-01-10T14:37:32.142747Z","shell.execute_reply.started":"2026-01-10T14:34:38.314770Z","shell.execute_reply":"2026-01-10T14:37:32.141949Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_samples = fix_dtype_ds(train_samples)\ntrain_s, val_s = train_test_split(train_samples, test_size=0.2)\ntest_s = fix_dtype_ds(test_s)\n\ntorch.save(val_s, \"val_s.pt\")\n\ntrain_ds = RankingDataset(train_s)\nval_ds = RankingDataset(val_s)\ntest_ds = RankingDataset(test_s)\n\nprint(f\"Train: {len(train_ds)} samples, Val: {len(val_ds)} samples, Test: {len(test_ds)} samples\")\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-10T14:37:32.143788Z","iopub.execute_input":"2026-01-10T14:37:32.144032Z","iopub.status.idle":"2026-01-10T14:37:53.183787Z","shell.execute_reply.started":"2026-01-10T14:37:32.144010Z","shell.execute_reply":"2026-01-10T14:37:53.176581Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training Setup","metadata":{}},{"cell_type":"markdown","source":"## Model Configuration","metadata":{}},{"cell_type":"code","source":"class MLPEncoder(nn.Module):\n    \"\"\"Flexible MLP with options for residual connections, layer norm, etc.\"\"\"\n\n    def __init__(self, input_dim, output_dim, hidden_dims=512, dropout=0.1):\n        super().__init__()\n        layers = []\n        prev_dim = input_dim\n\n        for i, hidden_dim in enumerate(hidden_dims):\n            # Linear\n            layers.append(nn.Linear(prev_dim, hidden_dim))\n            layers.append(nn.LayerNorm(hidden_dim))\n            layers.append(nn.ReLU())\n\n            # Dropout\n            if dropout > 0:\n                layers.append(nn.Dropout(dropout))\n\n            prev_dim = hidden_dim\n        self.residual_proj = (\n            nn.Identity() if input_dim == output_dim\n            else nn.Linear(input_dim, output_dim)\n        )\n        # Final layer\n        layers.append(nn.Linear(prev_dim, output_dim))\n\n        self.layers = nn.Sequential(*layers)\n\n    def forward(self, x):\n        out = self.layers(x)\n        out = out + self.residual_proj(x) # Residual Connection\n\n        return out\n\nclass EmbeddingEncoder(nn.Module):  # Fixed typo: EmeddingEncoder -> EmbeddingEncoder\n    \"\"\"Vocab-based encoder using EmbeddingBag for efficiency\"\"\"\n\n    def __init__(self, vocab_info, output_dim=16, mlp_hidden_dims=None, dropout=0.1):\n        \"\"\"\n        Args:\n            vocab_info: Dict with tuple structure:\n                {\n                    'carrier': (carrier_vocab2idx, 16),\n                    'airport': (airport_vocab2idx, 16),\n                    'company': (company_vocab2idx, 8),\n                    ...\n                }\n                where each value is (vocab2idx_dict, embed_dim)\n            output_dim: Final output dimension after MLP\n            mlp_hidden_dims: List of hidden dimensions for MLP, e.g., [128, 64]\n            dropout: Dropout rate for MLP\n        \"\"\"\n        super().__init__()\n\n        self.vocab_info = vocab_info\n\n        # Extract vocab2idx and embed_dims from tuples\n        self.vocab2idx_dict = {\n            name: vocab_tuple[0]\n            for name, vocab_tuple in vocab_info.items()\n        }\n        self.embed_dims = {\n            name: vocab_tuple[1]\n            for name, vocab_tuple in vocab_info.items()\n        }\n\n        # Map pooling name to EmbeddingBag mode\n        self.mode = 'mean'\n\n        # Create embeddings - use EmbeddingBag for multi-item, Embedding for single-item\n        self.single_embeddings = nn.ModuleDict()  # For single-item features\n        self.multi_embeddings = nn.ModuleDict()   # For multi-item features\n\n        # Build feature registry\n        self.feature_registry = self._build_feature_registry()\n\n        # Initialize embeddings based on feature types\n        self._initialize_embeddings()\n\n        # Calculate concatenated dimension\n        self.concat_dim = sum(\n            self.embed_dims[vocab_name]\n            for _, vocab_name, _ in self.feature_registry.values()\n        )\n\n        # Build MLP\n        self.mlp = self._build_mlp(\n            input_dim=self.concat_dim,\n            output_dim=output_dim,\n            hidden_dims=mlp_hidden_dims,\n            dropout=dropout\n        )\n        self.output_dim = output_dim\n\n    def _build_feature_registry(self):\n        \"\"\"\n        Registry: feature_name -> (is_multi, vocab_name, extractor)\n        Note: vocab_name should match keys in vocab_info (e.g., 'carrier')\n        \"\"\"\n        return {\n            'carriers_used': (\n                True,  # multi-item -> use EmbeddingBag\n                'carrier',\n                lambda batch: (\n                    batch[\"list_features\"][\"carriers_used\"][\"flat\"],  # List of tensors [B]\n                    batch[\"list_features\"][\"carriers_used\"][\"offsets\"]  # [B, padded_len]\n                )\n            ),\n            'ff_carriers': (\n                True,  # multi-item -> use EmbeddingBag\n                'carrier',\n                lambda batch: (\n                    batch[\"list_features\"][\"ff_carriers\"][\"flat\"],  # List of tensors [B]\n                    batch[\"list_features\"][\"ff_carriers\"][\"offsets\"]  # [B, padded_len]\n                )\n            ),\n            'departure_airport': (\n                False,  # single-item -> use regular Embedding\n                'airport',\n                lambda batch: batch[\"cater_features\"][\"departure_airport\"]\n            ),\n            'destination_airport': (\n                False,  # single-item -> use regular Embedding\n                'airport',\n                lambda batch: batch[\"cater_features\"][\"destination_airport\"]\n            ),\n            'companyID': (\n                False,\n                'company',\n                lambda batch: batch[\"cater_features\"][\"companyID\"]\n            ),\n            'nationality': (\n                False,\n                'nationality',\n                lambda batch: batch[\"cater_features\"][\"nationality\"]\n            ),\n        }\n\n    def _initialize_embeddings(self):\n        \"\"\"Create Embedding or EmbeddingBag based on feature types\"\"\"\n        vocab_usage = {}  # vocab_name -> needs_multi\n\n        for is_multi, vocab_name, _ in self.feature_registry.values():\n            if vocab_name not in vocab_usage:\n                vocab_usage[vocab_name] = {'multi': False, 'single': False}\n\n            if is_multi:\n                vocab_usage[vocab_name]['multi'] = True\n            else:\n                vocab_usage[vocab_name]['single'] = True\n\n        # Create embeddings\n        for vocab_name, usage in vocab_usage.items():\n            vocab_size = len(self.vocab2idx_dict[vocab_name])\n            embed_dim = self.embed_dims[vocab_name]\n\n            if usage['single']:\n                self.single_embeddings[vocab_name] = nn.Embedding(vocab_size, embed_dim)\n\n            if usage['multi']:\n                self.multi_embeddings[vocab_name] = nn.EmbeddingBag(\n                    vocab_size,\n                    embed_dim,\n                    mode=self.mode\n                )\n\n    def _build_mlp(self, input_dim, output_dim, hidden_dims=None, dropout=0.1):\n        \"\"\"\n        Build MLP: input_dim -> hidden_dims -> output_dim\n        \"\"\"\n        if hidden_dims is None:\n            # Simple linear projection\n            return nn.Linear(input_dim, output_dim)\n\n        # Multi-layer MLP\n        layers = []\n        prev_dim = input_dim\n\n        for hidden_dim in hidden_dims:\n            layers.extend([\n                nn.Linear(prev_dim, hidden_dim),\n                nn.ReLU(),\n                nn.Dropout(dropout),\n            ])\n            prev_dim = hidden_dim\n\n        # Final projection\n        layers.append(nn.Linear(prev_dim, output_dim))\n\n        return nn.Sequential(*layers)\n\n    def encode_single_item(self, data_lists, vocab_name):\n        \"\"\"\n        Encode single-item features using regular Embedding\n        Args:\n            data_lists: List of indices, e.g., [3, 7, 1]\n        Returns:\n            Tensor [B, embed_dim]\n        \"\"\"\n        return self.single_embeddings[vocab_name](data_lists)\n\n    def encode_multi_item(self, data_lists, vocab_name):\n        \"\"\"\n        Encode multi-item features using EmbeddingBag\n        Args:\n            data_lists: List of index lists, e.g., [[1, 2], [3], []]\n        Returns:\n            Tensor [B, embed_dim]\n        \"\"\"\n\n        indices, offsets = data_lists\n        \n        # EmbeddingBag directly returns pooled embeddings [B, embed_dim]\n        return self.multi_embeddings[vocab_name](indices, offsets)\n\n    def forward(self, batch):\n        \"\"\"\n        Args:\n            batch: Dict containing features\n        Returns:\n            Tensor [B, N, output_dim]\n        \"\"\"\n        num = batch[\"num_features\"]  # [B, N, num_dim]\n        B, N, _ = num.shape\n\n        encoded_features = []\n\n        for feature_name, (is_multi, vocab_name, extractor) in self.feature_registry.items():\n            # Extract data\n            data = extractor(batch)\n            # print(f\"Encoding feature '{feature_name}', data_sample={data[:5] if isinstance(data, list) else data}\")\n            # Encode\n            if is_multi:\n                feature = self.encode_multi_item(data, vocab_name)\n            else:\n                feature = self.encode_single_item(data, vocab_name)\n            # Reshape to [B*N, dim]\n            # print(f\"Shape: {feature.shape}\")\n            feature = feature.view(B * N, -1)\n            encoded_features.append(feature)\n\n        # Concatenate all embeddings\n        concat_features = torch.cat(encoded_features, dim=1)  # [B*N, concat_dim]\n        # Apply MLP\n        output = self.mlp(concat_features)  # [B*N, output_dim]\n\n        return output.view(B, N, -1)  # [B, N, output_dim]","metadata":{"id":"GLFWev7BDxtP","trusted":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2026-01-10T14:37:53.201901Z","iopub.execute_input":"2026-01-10T14:37:53.202120Z","iopub.status.idle":"2026-01-10T14:37:53.220857Z","shell.execute_reply.started":"2026-01-10T14:37:53.202100Z","shell.execute_reply":"2026-01-10T14:37:53.220199Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class SharedFeatureEncoder(nn.Module):\n    def __init__(self, num_dim, embed_dim, vocab_info, time_feat_dim=5, \n        time_hidden_dim=2, num_hidden_dim=None, embed_mlp_hidden_dims=None,dropout=0.1):\n        super().__init__()\n\n        # --------------------\n        # Numeric encoder\n        # --------------------\n        if num_hidden_dim is not None:\n            self.num_encoder = nn.Sequential(\n                nn.Linear(num_dim, num_hidden_dim),\n                nn.ReLU(),\n                nn.LayerNorm(num_hidden_dim),\n                nn.Dropout(dropout),\n            )\n            self.num_output_dim = num_hidden_dim\n        else:\n            self.num_encoder = nn.Identity()\n            self.num_output_dim = num_dim\n\n        # --------------------\n        # Time encoder (shared)\n        # --------------------\n        if time_hidden_dim:\n            self.time_encoder = nn.Sequential(\n                nn.Linear(time_feat_dim, time_hidden_dim),\n                nn.ReLU(),\n                nn.LayerNorm(time_hidden_dim),\n                nn.Dropout(dropout),\n            )\n        self.time_output_dim = time_hidden_dim\n\n        # --------------------\n        # Embedding encoder\n        # --------------------\n        if vocab_info is not None and embed_dim > 0:\n            self.embed_encoder = EmbeddingEncoder(\n                vocab_info=vocab_info,\n                output_dim=embed_dim,\n                mlp_hidden_dims=embed_mlp_hidden_dims,\n                dropout=dropout,\n            )\n            self.embed_output_dim = embed_dim\n        else:\n            self.embed_encoder = None\n            self.embed_output_dim = 0\n\n        # --------------------\n        # Total output dim\n        # --------------------\n        self.output_dim = (\n            self.num_output_dim\n            + self.time_output_dim\n            + self.embed_output_dim\n        )\n\n    def forward(self, batch):\n        \"\"\"\n        batch:\n            num_features:  [B, N, num_dim]\n            time_features: [B, N, T, 5]\n        \"\"\"\n        num = batch[\"num_features\"]\n        time = batch[\"time_features\"]\n\n        B, N, _ = num.shape\n        _, _, T, F = time.shape  # F should be 5\n\n        # ---- numeric ----\n        num = num.view(B * N, -1)\n        num_encoded = self.num_encoder(num)\n\n        features = [num_encoded]\n\n        # ---- time ----\n        if self.time_output_dim: \n            time = time.view(B * N * T, F)\n            time_encoded = self.time_encoder(time)\n            time_encoded = time_encoded.view(B * N, T, -1)\n    \n            # Aggregate time slots (mean pooling – stable default)\n            time_encoded = time_encoded.mean(dim=1)  # [B*N, time_hidden_dim]\n            features.append(time_encoded)\n\n        # ---- embeddings ----\n        if self.embed_encoder is not None:\n            embed_encoded = self.embed_encoder(batch)\n            embed_encoded = embed_encoded.view(B * N, -1)\n            features.append(embed_encoded)\n\n        out = torch.cat(features, dim=1)\n        return out.view(B, N, -1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-10T14:37:53.221859Z","iopub.execute_input":"2026-01-10T14:37:53.222166Z","iopub.status.idle":"2026-01-10T14:37:53.246096Z","shell.execute_reply.started":"2026-01-10T14:37:53.222114Z","shell.execute_reply":"2026-01-10T14:37:53.245365Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MLPScoreModel(nn.Module):\n    def __init__(self, input_dim, hidden_dims=[256, 128], dropout=0.1):\n        super().__init__()\n\n        layers = []\n        prev_dim = input_dim\n\n        # Build hidden layers\n        for hidden_dim in hidden_dims:\n            layers.extend([\n                nn.Linear(prev_dim, hidden_dim),\n                nn.ReLU(),\n                nn.Dropout(dropout)\n            ])\n            prev_dim = hidden_dim\n\n        # Final scoring layer (outputs single score per item)\n        layers.append(nn.Linear(prev_dim, 1))\n\n        self.mlp = nn.Sequential(*layers)\n\n    def forward(self, x):\n        # x: [B*N, D] -> [B*N, 1] -> [B*N]\n        return self.mlp(x).squeeze(-1)\n\nclass MLPRanker(nn.Module):\n    def __init__(self, feature_encoder, hidden_dims):\n        super().__init__()\n        self.encoder = feature_encoder\n        self.scorer = MLPScoreModel(feature_encoder.output_dim, hidden_dims)\n\n    def forward(self, batch):\n        x = self.encoder(batch)     # [B, N, D]\n        B, N, D = x.shape\n        scores = self.scorer(x.view(B * N, D))\n        return scores.view(B, N)\n","metadata":{"id":"MAhfW5vu2hRC","trusted":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2026-01-10T14:37:53.247015Z","iopub.execute_input":"2026-01-10T14:37:53.247283Z","iopub.status.idle":"2026-01-10T14:37:53.267378Z","shell.execute_reply.started":"2026-01-10T14:37:53.247262Z","shell.execute_reply":"2026-01-10T14:37:53.266698Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training Loss and Ulti","metadata":{}},{"cell_type":"markdown","source":"### Losses","metadata":{}},{"cell_type":"code","source":"def batch_cross_entropy_loss(scores, pos_idxs, mask):\n    \"\"\"\n    scores:   [B, N]\n    pos_idxs: [B]\n    mask:     [B, N]\n    \"\"\"\n\n    scores = scores.masked_fill(~mask, float('-inf'))\n    loss = F.cross_entropy(scores, pos_idxs)\n    \n    return loss\n\n# --- 2. Vectorized Loss Function ---\ndef batch_binary_cross_entropy_loss(scores, pos_idxs, mask, pos_weight=5.0):\n    \"\"\"\n    scores:   [B, N]\n    pos_idxs: [B]\n    mask:     [B, N]\n    \"\"\"\n    scores = scores.masked_fill(~mask, 0.0)   \n    labels = torch.zeros_like(scores)\n    labels.scatter_(1, pos_idxs.unsqueeze(1), 1.0)\n    loss = F.binary_cross_entropy_with_logits(\n        scores,\n        labels,\n        pos_weight=torch.tensor(pos_weight, device=scores.device)\n    )\n    \n    return loss\n\n# --- 2. Vectorized Loss Function ---\ndef batch_pairwise_hinge_loss(scores, pos_idx, mask, margin=1.0):\n    \"\"\"\n    Pairwise hinge loss - positive should score higher than negatives\n    \"\"\"\n    batch_size, max_len = scores.shape\n\n    # Get positive scores\n    pos_scores = scores[torch.arange(batch_size), pos_idx].unsqueeze(1)  # [B, 1]\n\n    # Create negative mask (all except positive)\n    neg_mask = mask.clone()\n    neg_mask[torch.arange(batch_size), pos_idx] = False\n\n    # Compute pairwise differences\n    differences = margin - (pos_scores - scores)  # [B, Max_Len]\n    differences = differences.masked_fill(~neg_mask, 0)\n\n    # Hinge loss\n    loss = torch.clamp(differences, min=0).sum() / neg_mask.sum()\n    return loss","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2026-01-10T14:37:53.268151Z","iopub.execute_input":"2026-01-10T14:37:53.268408Z","iopub.status.idle":"2026-01-10T14:37:53.289384Z","shell.execute_reply.started":"2026-01-10T14:37:53.268387Z","shell.execute_reply":"2026-01-10T14:37:53.288826Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Helper Ulti","metadata":{}},{"cell_type":"code","source":"# --- 1. Updated Collate Function for EmbeddingBag Architecture ---\ndef bucketed_collate_fn(batch):\n    \"\"\"\n    Collate function for RankingDataset that returns dicts with EmbeddingBag format.\n\n    Args:\n        batch: List of dicts from RankingDataset.__getitem__\n               Each dict has: num_features, list_features (with 'flat' and 'offsets'),\n                             cater_features, positive_idx, Id, ranker_id\n\n    Returns:\n        batch_dict: Dict with padded features for model input\n        pos_idx: Tensor of positive indices [B]\n        attention_masks: Tensor of valid positions [B, padded_len]\n    \"\"\"\n    bucket_boundaries = [5, 20, 50, 160, 256, 620, 1500, 2500, 6400, 8300]\n    batch_size = len(batch)\n\n    # Extract positive indices\n    pos_idx = torch.tensor([sample['positive_idx'] for sample in batch])\n\n    # Get all features and their lengths\n    num_features_list = [sample['num_features'] for sample in batch]\n    lengths = [f.shape[0] for f in num_features_list]\n    max_len_in_batch = max(lengths)\n    \n    # Determine bucket size\n    padded_len = max_len_in_batch\n    for boundary in bucket_boundaries:\n        if max_len_in_batch <= boundary:\n            padded_len = boundary\n            break\n    \n    # Pad numerical features\n    padded_num_features = torch.nn.utils.rnn.pad_sequence(\n        num_features_list,\n        batch_first=True,\n        padding_value=0.0\n    )\n    \n    # If bucket size > max_len_in_batch, add extra padding\n    if padded_len > max_len_in_batch:\n        extra_padding = padded_len - max_len_in_batch\n        padded_num_features = torch.nn.functional.pad(\n            padded_num_features,\n            (0, 0, 0, extra_padding),  # (left, right, top, bottom)\n            value=0.0\n        )\n    # Pad time features [N, T, 5]\n    time_features_list = [sample[\"time_features\"] for sample in batch]\n    \n    # This pads on dim=0 (N) only\n    padded_time_features = torch.nn.utils.rnn.pad_sequence(\n        time_features_list,\n        batch_first=True,\n        padding_value=0.0  # safe: zero = neutral for sin/cos + flags\n    )\n    \n    # If bucket size > max_len_in_batch, add extra padding\n    if padded_len > max_len_in_batch:\n        extra_padding = padded_len - max_len_in_batch\n        padded_time_features = torch.nn.functional.pad(\n            padded_time_features,\n            (0, 0, 0, 0, 0, extra_padding),  # pad N dimension only\n            value=0.0\n        )\n\n    # Pad list features in EmbeddingBag format (carriers_used, ff_carriers)\n    all_list_features = {}\n    for key in ['carriers_used', 'ff_carriers']:\n        batch_flat = []\n        batch_offsets = []\n\n        current_base = 0\n        for sample in batch:\n            flat = sample['list_features'][key]['flat']\n            offsets = sample['list_features'][key]['offsets']\n\n            batch_flat.append(flat)\n            batch_offsets.append(offsets + current_base)\n\n            current_base += flat.numel()\n\n        # Concatenate once\n        batch_flat = torch.cat(batch_flat)\n        batch_offsets = torch.cat(batch_offsets)\n\n        # Add padding offsets\n        padding_offsets = []\n        for sample in batch:\n            actual_len = sample['list_features'][key]['offsets'].numel()\n            padding_needed = padded_len - actual_len\n            if padding_needed > 0:\n                padding_offsets.append(\n                    torch.full((padding_needed,), batch_flat.numel(), dtype=torch.long)\n                )\n\n        if padding_offsets:\n            batch_offsets = torch.cat([batch_offsets] + padding_offsets)\n\n        # Convert to tensors\n        all_list_features[key] = {\n            'flat': batch_flat,  # [total_items]\n            'offsets': batch_offsets  # [B * padded_len]\n        }\n\n    # Pad categorical features (departure_airport, destination_airport, companyID, nationality)\n    all_cater_features = {}\n    for key in ['departure_airport', 'destination_airport', 'companyID', 'nationality']:\n        B = len(batch)\n\n        # Allocate once (already padded with zeros on the right)\n        batch_cater = torch.zeros(\n            (B, padded_len),\n            dtype=torch.long\n        )\n\n        for i, sample in enumerate(batch):\n            t = torch.as_tensor(sample['cater_features'][key], dtype=torch.long)\n            n = min(t.numel(), padded_len)\n            batch_cater[i, :n] = t[:n]  # right-padding is implicit\n\n        all_cater_features[key] = batch_cater\n        \n    # Create attention mask vectorized\n    lengths_tensor = torch.tensor(lengths, dtype=torch.long)\n    attention_masks = torch.arange(padded_len).unsqueeze(0) < lengths_tensor.unsqueeze(1)  # [B, padded_len]\n\n    # Build batch dict for model\n    batch_dict = {\n        'num_features': padded_num_features,      # [B, padded_len, num_dim]\n        'time_features': padded_time_features,\n        'list_features': all_list_features,       # Dict with 'flat' and 'offsets' tensors\n        'cater_features': all_cater_features,     # Dict of [B, padded_len] tensors\n    }\n\n    return batch_dict, pos_idx, attention_masks","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2026-01-10T14:37:53.291710Z","iopub.execute_input":"2026-01-10T14:37:53.292074Z","iopub.status.idle":"2026-01-10T14:37:53.310661Z","shell.execute_reply.started":"2026-01-10T14:37:53.292054Z","shell.execute_reply":"2026-01-10T14:37:53.310082Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def recall_at_k(scores, pos_idxs, mask, k=200):\n    \"\"\"\n    scores:   [B, L]\n    pos_idxs: [B]\n    mask:     [B, L]\n    \"\"\"\n    scores = scores.masked_fill(~mask, -1e9)\n    topk = scores.topk(k=min(k, scores.size(1)), dim=1).indices\n    hits = (topk == pos_idxs.unsqueeze(1)).any(dim=1)\n    return hits.float().mean().item()\n\ndef hitrate_at_3(scores, pos_idx, mask):\n    \"\"\"\n    Calculate Hit Rate @ 3: proportion of samples where positive item is in top 3.\n\n    Args:\n        scores: (batch_size, max_len) - predicted scores\n        pos_idx: (batch_size,) - index of positive item\n        mask: (batch_size, max_len) - attention mask\n\n    Returns:\n        float: hit rate (0.0 to 1.0)\n    \"\"\"\n    # Mask out padding\n    masked_scores = scores.clone()\n    masked_scores[~mask] = -1e9\n\n    # Get top 3 indices\n    _, top3_indices = torch.topk(masked_scores, k=3, dim=1)\n\n    # Check if positive is in top 3\n    hits = (top3_indices == pos_idx.unsqueeze(1)).any(dim=1)\n\n    return hits.float().mean().item()\n\nimport torch\n\ndef to_device(batch_dict, device, dtype=None):\n    \"\"\"\n    Recursively move all tensors in a nested structure to the specified device and dtype.\n    \n    Args:\n        batch_dict: Can be:\n            - Dict: Recursively processes all values\n            - List/Tuple: Recursively processes all elements\n            - Tensor: Moves to device and optionally converts dtype\n            - Other: Returns as-is\n        device: torch.device to move tensors to\n        dtype: Optional torch.dtype to convert tensors to (e.g., torch.float32)\n    \n    Returns:\n        The input structure with all tensors moved to the specified device and dtype\n    \"\"\"\n    if isinstance(batch_dict, torch.Tensor):\n        # Move tensor to device\n        tensor = batch_dict.to(device)\n        # Optionally convert dtype\n        if dtype is not None:\n            tensor = tensor.to(dtype)\n        return tensor\n    \n    elif isinstance(batch_dict, dict):\n        # Recursively process dictionary values\n        return {key: to_device(value, device, dtype) for key, value in batch_dict.items()}\n    \n    elif isinstance(batch_dict, (list, tuple)):\n        # Recursively process list/tuple elements, preserving the type\n        processed = [to_device(item, device, dtype) for item in batch_dict]\n        return type(batch_dict)(processed)\n    \n    else:\n        # Return non-tensor types as-is (int, str, None, etc.)\n        return batch_dict\n\n\ndef plot_training_metrics(epochs, train_losses, val_losses=None, val_recalls=None, save_path='training_plot.png'):\n    \"\"\"\n    Plot training and validation metrics and save to file.\n    \n    Args:\n        epochs: List of epoch numbers\n        train_losses: List of training losses\n        val_losses: Optional list of validation losses\n        val_recalls: Optional list of validation recall@k scores\n        save_path: Path to save the plot image\n    \"\"\"\n    fig, axes = plt.subplots(1, 2, figsize=(15, 5))\n    \n    # Plot 1: Training and Validation Loss\n    axes[0].plot(epochs, train_losses, 'b-o', label='Train Loss', linewidth=2, markersize=6)\n    if val_losses:\n        axes[0].plot(epochs, val_losses, 'r-s', label='Val Loss', linewidth=2, markersize=6)\n    \n    axes[0].set_xlabel('Epoch', fontsize=12)\n    axes[0].set_ylabel('Loss', fontsize=12)\n    axes[0].set_title('Training and Validation Loss', fontsize=14, fontweight='bold')\n    axes[0].legend(fontsize=10)\n    axes[0].grid(True, alpha=0.3)\n    \n    # Plot 2: Validation Recall@k\n    if val_recalls:\n        axes[1].plot(epochs, val_recalls, 'g-^', label='Val Recall@3', linewidth=2, markersize=6)\n        axes[1].set_xlabel('Epoch', fontsize=12)\n        axes[1].set_ylabel('Recall@3', fontsize=12)\n        axes[1].set_title('Validation Recall@3', fontsize=14, fontweight='bold')\n        axes[1].legend(fontsize=10)\n        axes[1].grid(True, alpha=0.3)\n    else:\n        axes[1].text(0.5, 0.5, 'No validation data', \n                    ha='center', va='center', fontsize=14, transform=axes[1].transAxes)\n        axes[1].set_title('Validation Recall@3', fontsize=14, fontweight='bold')\n    \n    plt.tight_layout()\n    plt.savefig(save_path, dpi=300, bbox_inches='tight')\n    plt.close()\n    print(f\"Training plot saved to: {save_path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-10T14:37:53.311489Z","iopub.execute_input":"2026-01-10T14:37:53.311838Z","iopub.status.idle":"2026-01-10T14:37:53.332475Z","shell.execute_reply.started":"2026-01-10T14:37:53.311805Z","shell.execute_reply":"2026-01-10T14:37:53.331840Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Make Validate Data","metadata":{}},{"cell_type":"code","source":"def val_collate_fn(batch):\n    batch_dict, pos_idx, attention_masks = bucketed_collate_fn(batch)\n    id_list = [sample['Id'] for sample in batch]\n    ranker_id_list = [sample['ranker_id'] for sample in batch]\n    return batch_dict, pos_idx, attention_masks, id_list, ranker_id_list\n\ndef make_val_df(val_ds, reranker_model, mlp_model, feature_encoder):\n    \"\"\"\n    Perform inference on val_ds using mlp_model for prefiltering and reranker_model for final scoring.\n    If reranker_model is None, only MLP scores are used.\n\n    Args:\n        val_ds: RankingDataset for validation (unprocessed, raw features)\n        reranker_model: Trained TransformerReranker model (can be None)\n        mlp_model: Trained MLPScoreModel for prefiltering\n        feature_encoder: SharedFeatureEncoder for encoding features\n\n    Returns:\n        val_df: Polars DataFrame with inference results\n    \"\"\"\n    val_loader = DataLoader(val_ds, batch_size=32, collate_fn=val_collate_fn, shuffle=False)\n\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n    # Set models to eval mode\n    if reranker_model is not None:\n        reranker_model.eval()\n        reranker_model.to(device)\n    \n    mlp_model.eval()\n    mlp_model.to(device)\n    \n    feature_encoder.eval()\n    feature_encoder.to(device)\n\n    # Store results per sample (not concatenated)\n    all_preds = []  # List of 1D tensors (variable length)\n    all_ids = []\n    all_ranker_ids = []\n    all_pos_idx = []\n\n    total_batches = len(val_loader)\n    \n    with torch.no_grad():\n        pbar = tqdm(enumerate(val_loader), total=total_batches, desc=\"Processing batches\")\n        for batch_idx, (batch_dict, pos_idx, masks, ids, ranker_ids) in pbar:\n            batch_dict = to_device(batch_dict, device)\n            masks = masks.to(device)\n            \n            actual_batch_size = batch_dict['num_features'].shape[0]\n            seq_len = batch_dict['num_features'].shape[1]\n            \n            # Get MLP prefetch scores\n            prefetch_scores = mlp_model(batch_dict)\n            \n            # If no reranker, use MLP scores directly\n            preds = prefetch_scores\n            pbar.set_postfix({\n                'batch_size': actual_batch_size,\n                'seq_len': seq_len,\n                'action': 'mlp_only'\n            })\n            \n            # Store per-sample results (move to CPU immediately)\n            preds_cpu = preds.cpu()\n            masks_cpu = masks.cpu()\n            \n            for i in range(actual_batch_size):\n                seq_len_i = masks_cpu[i].sum().item()\n                all_preds.append(preds_cpu[i, :seq_len_i])  # Store only valid predictions\n            \n            all_ids.extend(ids)\n            all_ranker_ids.extend(ranker_ids)\n            all_pos_idx.extend(pos_idx.tolist())\n\n    # Process all results\n    all_pred_scores = []\n    all_selected = []\n    expanded_ids = []\n    expanded_ranker_ids = []\n    \n    for i in tqdm(range(len(all_preds)), desc=\"Building final results\"):\n        preds_np = all_preds[i].numpy()\n        seq_len = len(preds_np)\n        \n        # Build labels\n        selected = np.zeros(seq_len, dtype=int)\n        selected[all_pos_idx[i]] = 1\n        \n        # Extend lists\n        all_pred_scores.extend(preds_np.tolist())\n        all_selected.extend(selected.tolist())\n        expanded_ids.extend(all_ids[i])\n        expanded_ranker_ids.extend([all_ranker_ids[i]] * seq_len)\n\n    # Convert to Polars DataFrame\n    val_df = pl.DataFrame({\n        \"Id\": expanded_ids,\n        \"ranker_id\": expanded_ranker_ids,\n        \"pred_score\": all_pred_scores,\n        \"selected\": all_selected\n    })\n\n    # Add group_size\n    val_df = val_df.join(\n        val_df.group_by(\"ranker_id\").agg(pl.len().alias(\"group_size\")),\n        on=\"ranker_id\"\n    )\n\n    print(f\"\\nValidation DataFrame created successfully! Total rows: {len(val_df)}\")\n    \n    return val_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-10T14:37:53.333418Z","iopub.execute_input":"2026-01-10T14:37:53.333642Z","iopub.status.idle":"2026-01-10T14:37:53.355418Z","shell.execute_reply.started":"2026-01-10T14:37:53.333622Z","shell.execute_reply":"2026-01-10T14:37:53.354750Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def print_hitrates(val_df, k_values=[1, 3, 5, 10]):\n    \"\"\"\n    Calculate and print HitRate@k for specified k values.\n    \n    Args:\n        val_df: Polars DataFrame with columns: ranker_id, pred_score, selected, group_size\n        k_values: List of k values to calculate HitRate for (default: [1, 3, 5, 10])\n    \"\"\"\n    # Filter groups with more than 10 items\n    va_df = val_df.filter(pl.col(\"group_size\") > 10)\n    \n    # Sort by ranker_id and pred_score (descending)\n    sorted_df = va_df.sort([\"ranker_id\", \"pred_score\"], descending=[False, True])\n    \n    print(\"\\nHitRate Results:\")\n    print(\"-\" * 40)\n    \n    for k in k_values:\n        hitrate = (\n            sorted_df.group_by(\"ranker_id\", maintain_order=True)\n            .head(k)\n            .group_by(\"ranker_id\")\n            .agg(pl.col(\"selected\").max().alias(\"hit\"))\n            .select(pl.col(\"hit\").mean())\n            .item()\n        )\n        print(f\"HitRate@{k:2d}: {hitrate:.4f} ({hitrate*100:.2f}%)\")\n    \n    print(\"-\" * 40)\n    print(f\"Total groups evaluated: {va_df['ranker_id'].n_unique()}\")\n    print(f\"Total items: {len(va_df)}\")\n\ndef quick_validate_model(val_ds, reranker_model, mlp_model, feature_encoder, k_values=[1, 3, 5, 10]):\n    val_df = make_val_df(val_ds, reranker_model, mlp_model, feature_encoder)\n    print_hitrates(val_df, k_values=[1, 3, 5, 10])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-10T14:37:53.356319Z","iopub.execute_input":"2026-01-10T14:37:53.356597Z","iopub.status.idle":"2026-01-10T14:37:53.376006Z","shell.execute_reply.started":"2026-01-10T14:37:53.356567Z","shell.execute_reply":"2026-01-10T14:37:53.375501Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom dataclasses import dataclass, field\nimport numpy as np\n\n\n# ============================================================\n# Training Configuration (Dataclass)\n# ============================================================\n\n@dataclass\nclass TrainingConfig:\n    mlp_train_batch_size: int = 32\n\n    mlp_hidden_dim: tuple = (512, 256, 128)\n    lr: float = 1e-4\n    num_epochs: int = 50\n    # Dynamically set later\n    input_dim: int = None\n    num_dim: int = train_ds[0]['num_features'].shape[1]        # REQUIRED: numeric feature dimension\n    time_dim: int = train_ds[0]['time_features'].shape[2]\n\nconfig = TrainingConfig()\n\n# Save vocabularies for later use\nvocab_info = {\n    'carrier': (carrier_vocab2idx, 16),\n    'airport': (airport_vocab2idx, 16),\n    'company': (company_vocab2idx, 8),\n    'nationality': (nationality_vocab2idx, 8),\n}","metadata":{"id":"ZEqxjaFUBbR-","trusted":true,"execution":{"iopub.status.busy":"2026-01-10T14:37:53.376765Z","iopub.execute_input":"2026-01-10T14:37:53.376982Z","iopub.status.idle":"2026-01-10T14:37:53.397992Z","shell.execute_reply.started":"2026-01-10T14:37:53.376939Z","shell.execute_reply":"2026-01-10T14:37:53.397455Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## MLP Training","metadata":{"id":"w0O0daIADxtP"}},{"cell_type":"code","source":"print(config.num_dim)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-10T14:37:53.398659Z","iopub.execute_input":"2026-01-10T14:37:53.398831Z","iopub.status.idle":"2026-01-10T14:37:53.416427Z","shell.execute_reply.started":"2026-01-10T14:37:53.398814Z","shell.execute_reply":"2026-01-10T14:37:53.415691Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"feature_encoder = SharedFeatureEncoder(\n    num_dim=config.num_dim,\n    embed_dim=0,\n    time_feat_dim=config.time_dim,\n    time_hidden_dim=64,\n    vocab_info=vocab_info,\n    num_hidden_dim=256,\n    embed_mlp_hidden_dims=[256, 128, 64],\n    dropout=0.1\n)\nmanual_batch = [train_ds[i] for i in range(3)]\ncollate_output = bucketed_collate_fn(manual_batch)\n# print(collate_output)\n\n# print(\"Feature shape:\", collate_output[0]['num_features'].shape)\nprint(feature_encoder(collate_output[0]).shape)","metadata":{"id":"c2WvysN4PnwQ","outputId":"fb927732-f289-4d1a-8174-84a336231791","trusted":true,"execution":{"iopub.status.busy":"2026-01-10T14:37:53.417413Z","iopub.execute_input":"2026-01-10T14:37:53.417793Z","iopub.status.idle":"2026-01-10T14:37:53.564292Z","shell.execute_reply.started":"2026-01-10T14:37:53.417761Z","shell.execute_reply":"2026-01-10T14:37:53.563655Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader\nimport matplotlib.pyplot as plt\nfrom pathlib import Path\n\n# --- 3. Updated CPU/GPU Training Function ---\ndef train_model_mlp(config, feature_encoder, train_dataset, val_dataset=None):\n    \"\"\"\n    Updated training function for SharedFeatureEncoder architecture on CPU/GPU.\n\n    Args:\n        config: Training configuration object with attributes:\n            - mlp_train_batch_size: Batch size for training\n            - mlp_hidden_dim: Hidden dimensions for MLP scorer\n            - lr: Learning rate\n            - use_bfloat16: Whether to use bfloat16 (default False for CPU/GPU)\n            - device: Optional device specification ('cuda', 'cpu', or torch.device)\n        feature_encoder: SharedFeatureEncoder instance\n        train_dataset: RankingDataset for training\n        val_dataset: Optional RankingDataset for validation\n    \"\"\"\n    # Initialize lists to track losses\n    train_losses = []\n    val_losses = []\n    val_recalls = []\n    epochs_list = []\n    \n    # 1. Create DataLoaders FIRST (before any CUDA operations)\n    print(\"Creating DataLoaders...\")\n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=config.mlp_train_batch_size,\n        shuffle=True,\n        collate_fn=bucketed_collate_fn,\n        num_workers=4,\n        persistent_workers=True,\n        pin_memory=True,\n        prefetch_factor=4,\n    )\n    \n    val_loader = None\n    if val_dataset:\n        val_loader = DataLoader(\n            val_dataset,\n            batch_size=config.mlp_train_batch_size,\n            shuffle=False,\n            collate_fn=bucketed_collate_fn,\n            num_workers=4,\n            persistent_workers=True,\n            pin_memory=True,\n            prefetch_factor=4,\n        )\n    path = Path(\"mlp_model\")\n    path.mkdir(exist_ok=True)\n    Path(\"training_plot\").mkdir(exist_ok=True)\n    \n    # 2. NOW set up device and move model to device\n    if hasattr(config, 'device') and config.device is not None:\n        device = config.device if isinstance(config.device, torch.device) else torch.device(config.device)\n    else:\n        device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n    print(f\"Training on device: {device}\")\n    \n    # 3. Build model and move to device AFTER DataLoader creation\n    model = MLPRanker(\n        feature_encoder=feature_encoder,\n        hidden_dims=config.mlp_hidden_dim,\n    ).to(device)\n\n    # Enable mixed precision training for GPU\n    use_amp = getattr(config, 'use_amp', False) and device.type == 'cuda'\n\n    if use_amp:\n        print(\"Using automatic mixed precision (AMP) for faster training...\")\n\n    optimizer = torch.optim.AdamW(\n        model.parameters(),\n        lr=config.lr,\n        weight_decay=0.01\n    )\n    \n    scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n        optimizer, mode='min', factor=0.5, patience=2\n    )\n    \n    print(\"Setup complete. Start training...\")\n\n    try:\n        for epoch in range(config.num_epochs):\n            model.train()\n            epoch_loss_sum = 0.0\n            step_count = 0\n\n            for step, (batch_dict, batch_pos, batch_masks) in tqdm(enumerate(train_loader), desc=f\"Epoch {epoch+1}\", total=len(train_loader)):\n                # Move data to device\n                batch_dict = to_device(batch_dict, device)               \n                batch_pos = batch_pos.to(device)\n                batch_masks = batch_masks.to(device)\n\n                optimizer.zero_grad()\n\n                # Forward pass\n                scores = model(batch_dict)  # [B, N]\n                loss = batch_cross_entropy_loss(scores, batch_pos, batch_masks)\n\n                # Backward pass\n                loss.backward()\n                torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n                optimizer.step()\n\n                # Accumulate loss\n                epoch_loss_sum += loss.item()\n                step_count += 1\n\n            final_avg_loss = epoch_loss_sum / step_count\n            train_losses.append(final_avg_loss)\n            epochs_list.append(epoch + 1)\n            \n            print(f\"[Epoch {epoch+1}] Train Loss: {final_avg_loss:.4f} | Lr: {optimizer.param_groups[0]['lr']}\")\n\n            # ---- Validation ----\n            if val_dataset and val_loader is not None:\n                model.eval()\n                val_loss_sum = 0.0\n                val_recall_sum_3 = 0.0\n                val_steps = 0\n\n                print(\"Start Validating\")\n\n                with torch.no_grad():\n                    for step, (batch_dict, batch_pos, batch_masks) in tqdm(enumerate(val_loader), desc=\"Validating\", total=len(val_loader)):\n                        # Move data to device\n                        batch_dict = to_device(batch_dict, device)\n                        batch_pos = batch_pos.to(device)\n                        batch_masks = batch_masks.to(device)\n                        scores = model(batch_dict)  # [B, N]\n\n                        # Calculate metrics\n                        val_recall_sum_3 += recall_at_k(scores, batch_pos, batch_masks, k=3)\n                        loss = batch_cross_entropy_loss(scores, batch_pos, batch_masks)\n\n                        val_loss_sum += loss.item()\n                        val_steps += 1\n\n                avg_val_loss = val_loss_sum / val_steps\n                avg_val_recall_3 = val_recall_sum_3 / val_steps\n                \n                val_losses.append(avg_val_loss)\n                val_recalls.append(avg_val_recall_3)\n                \n                print(f\"            Val Loss: {avg_val_loss:.4f} | Recall@3: {avg_val_recall_3:.4f}\")\n                \n                scheduler.step(avg_val_loss)\n\n            # Save checkpoint\n            if (epoch + 1) % 5 == 0 or epoch == config.num_epochs - 1:\n                print(f\"Saving checkpoint at epoch {epoch+1}...\")\n                torch.save(model.state_dict(), f'mlp_model/mlp_epoch_{epoch+1}.pth')\n                \n                # Save loss plot\n                plot_training_metrics(epochs_list, train_losses, val_losses, val_recalls, \n                                    save_path=f'training_plot/training_plot_epoch_{epoch+1}.png')\n\n    except Exception as e:\n        print(f\"Training interrupted: {e}\")\n        raise\n    finally:\n        print(\"Cleaning up...\")\n        if device.type == 'cuda':\n            torch.cuda.empty_cache()\n\n    print(\"Training complete!\")\n    \n    # Final plot\n    plot_training_metrics(epochs_list, train_losses, val_losses, val_recalls, \n                        save_path='training_plot/final_training_plot.png')\n\n    return model\n","metadata":{"id":"jvTHXXAaDxtP","trusted":true,"execution":{"iopub.status.busy":"2026-01-10T14:37:53.565145Z","iopub.execute_input":"2026-01-10T14:37:53.565400Z","iopub.status.idle":"2026-01-10T14:37:53.582366Z","shell.execute_reply.started":"2026-01-10T14:37:53.565378Z","shell.execute_reply":"2026-01-10T14:37:53.581684Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## MLP Train Execute","metadata":{}},{"cell_type":"code","source":"# Train model\nmlp_model = train_model_mlp(\n    config=config,\n    feature_encoder=feature_encoder,\n    train_dataset=train_ds,\n    val_dataset=val_ds,\n)\n# Initialize model architecture\n# mlp_model = MLPScoreModel(\n#     input_dim=33,\n#     hidden_dim=config.mlp_hidden_dim,\n#     layers=config.mlp_layers\n# )\n\n# # Load trained weights\n# checkpoint_path = 'mlp_epoch_30.pth'\n# state_dict = torch.load(checkpoint_path, map_location='cpu')\n# model.load_state_dict(state_dict)\n# print(f\"✓ Loaded model from {checkpoint_path}\")","metadata":{"id":"J9FOAeftsYoL","outputId":"7e906118-b214-4d3b-fcc4-a26058f0dcb9","trusted":true,"execution":{"iopub.status.busy":"2026-01-10T14:37:53.583163Z","iopub.execute_input":"2026-01-10T14:37:53.583516Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Evaluation and visualization","metadata":{"id":"_tuvBCqTDxtP"}},{"cell_type":"markdown","source":"## Validation and visualize","metadata":{}},{"cell_type":"code","source":"quick_validate_model(val_ds, None, mlp_model, feature_encoder, k_values=[1, 3, 5, 10])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Submission","metadata":{"id":"1zfnXgFxDxtP"}},{"cell_type":"code","source":"test_df = make_val_df(test_ds, None, mlp_model, feature_encoder)","metadata":{"id":"f_6IZuWFJ5Fa","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# Build submission from test_df + model predictions\n# ============================================================\n\n# test_df must contain:\n#   - Id\n#   - ranker_id\n#   - pred_score  (already added when building test_df)\n\nsubmission_df = (\n    test_df\n    .select([\"Id\", \"ranker_id\", \"pred_score\"])\n    .with_columns(\n        # Rank pred_score within each ranker_id group\n        # Higher pred_score → Lower rank number (rank 1 = best)\n        pl.col(\"pred_score\")\n        .rank(method=\"ordinal\", descending=True)\n        .over(\"ranker_id\")\n        .cast(pl.Int32)\n        .alias(\"selected\")\n    )\n    .select([\"Id\", \"ranker_id\", \"selected\"])\n    # CRITICAL: Sort by Id to preserve original test.csv row order\n    .sort(\"Id\")\n)\n\nprint(submission_df)\n# Save to CSV\nsubmission_df.write_csv(\"submission.csv\")\n\nprint(\"Submission saved to submission.csv\")\nprint(f\"Total rows: {len(submission_df)}\")\n\n# Validation checks\nprint(\"\\n=== Validation Checks ===\")\nvalidation = submission_df.group_by(\"ranker_id\").agg([\n    pl.col(\"selected\").min().alias(\"min_rank\"),\n    pl.col(\"selected\").max().alias(\"max_rank\"),\n    pl.col(\"selected\").n_unique().alias(\"unique_ranks\"),\n    pl.col(\"selected\").count().alias(\"n_flights\")\n])\n\n# Check if ranks form valid permutations (1, 2, 3, ..., N)\ninvalid = validation.filter(\n    (pl.col(\"min_rank\") != 1) |\n    (pl.col(\"max_rank\") != pl.col(\"n_flights\")) |\n    (pl.col(\"unique_ranks\") != pl.col(\"n_flights\"))\n)\n\nif len(invalid) > 0:\n    print(f\"⚠️  WARNING: {len(invalid)} ranker_ids have invalid rank permutations!\")\n    print(invalid)\nelse:\n    print(\"✓ All ranker_ids have valid rank permutations (1, 2, 3, ..., N)\")\n\nprint(f\"✓ Total unique ranker_ids: {submission_df['ranker_id'].n_unique()}\")","metadata":{"id":"9e54gTPiDxtP","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# !kaggle competitions submit -c aeroclub-recsys-2025 -f submission.csv -m \"Message\"","metadata":{"id":"bPDoRFGBn8Z6","trusted":true},"outputs":[],"execution_count":null}]}