{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.10.18","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"tpu1vmV38","dataSources":[{"sourceId":105399,"databundleVersionId":12733338,"sourceType":"competition"},{"sourceId":2799445,"sourceType":"datasetVersion","datasetId":1710071},{"sourceId":12811424,"sourceType":"datasetVersion","datasetId":8101014},{"sourceId":12811466,"sourceType":"datasetVersion","datasetId":8101032},{"sourceId":12811493,"sourceType":"datasetVersion","datasetId":8101054},{"sourceId":12811512,"sourceType":"datasetVersion","datasetId":8101067},{"sourceId":12813058,"sourceType":"datasetVersion","datasetId":8100976},{"sourceId":12813166,"sourceType":"datasetVersion","datasetId":8102059},{"sourceId":12813548,"sourceType":"datasetVersion","datasetId":8102308}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# README  \nThis is an exmaple with FLAGS.fast = True set which means only use 1/100 data for train/eval, incase you want to run all data you might need about 300-500g memrory, so 512G mem is recommended, which exceeds TPU V3 300G mem.    \nNOTICE as I used random seed due to randomness, there might be slightly diff for LB/PB score, especially LB score.  \nTo get my best single model score just set   \nFLAGS.mode = 'infer'  \nFLAGS.fast = False   \nFLAGS.use_ext = True  \nFLAGS.n_models = 1  \nFLAGS.history_avg = False   \nFLAGS.model_dir = '../input/aeroclub-recsys-2025-model1'  \nTo get ensemble moels score just set FLAGS.n_models = 0  \nTo get even better score you could set:  \nFLAGS.history_avg = True   \nFLAGS.model_dir = '../input/aeroclub-recsys-2025-model2'   \nIn case you want to run training just set FLAGS.mode = 'train' but you need to train on local machine with 512G+ mem recommened.  \n\n","metadata":{}},{"cell_type":"markdown","source":"# Dependences\n!pip install icecream  \n!pip install airportsdata  \n!pip install timezonefinder  ","metadata":{}},{"cell_type":"code","source":"!pip install icecream --no-index --find-links=file:///kaggle/input/icecream/ \n!pip install airportsdata --no-index --find-links=file:///kaggle/input/airportsdata/ \n!pip install timezonefinder --no-index --find-links=file:///kaggle/input/timezonefinder/ \n!pip install polars --no-index --find-links=file:///kaggle/input/polars/ \n!pip install xgboost --no-index --find-links=file:///kaggle/input/xgboost/ ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T02:46:52.002339Z","iopub.execute_input":"2025-08-25T02:46:52.004181Z","iopub.status.idle":"2025-08-25T02:47:17.091946Z","shell.execute_reply.started":"2025-08-25T02:46:52.004145Z","shell.execute_reply":"2025-08-25T02:47:17.087128Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from icecream import ic\nimport sys\nimport os\nimport pickle\nimport numpy as np\nimport polars as pl\nimport pandas as pd\nfrom tqdm.auto import tqdm\nimport json\nimport glob\nimport logging\nimport time\nimport math\nimport polars.selectors as cs\nfrom collections import OrderedDict\nfrom itertools import chain\nfrom IPython.display import display","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T02:47:17.094026Z","iopub.execute_input":"2025-08-25T02:47:17.094265Z","iopub.status.idle":"2025-08-25T02:47:18.270740Z","shell.execute_reply.started":"2025-08-25T02:47:17.094224Z","shell.execute_reply":"2025-08-25T02:47:18.266066Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import airportsdata\nfrom timezonefinder import TimezoneFinder\nimport pytz\nfrom datetime import datetime\nfrom zoneinfo import ZoneInfo","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T02:47:18.273074Z","iopub.execute_input":"2025-08-25T02:47:18.273450Z","iopub.status.idle":"2025-08-25T02:47:18.420281Z","shell.execute_reply.started":"2025-08-25T02:47:18.273407Z","shell.execute_reply":"2025-08-25T02:47:18.416781Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"logger = logging.getLogger('aeroclub')\n#handler = logging.StreamHandler()\nfrom rich.logging import RichHandler\nhandler = RichHandler(rich_tracebacks=True)\nhandler.setLevel(logging.INFO)\nformatter = logging.Formatter(\"%(asctime)s [%(levelname)s] %(message)s\")\nhandler.setFormatter(formatter)\n\nlogger.addHandler(handler)\nlogger.setLevel(logging.INFO)\nlogger.propagate = False   # 如果不想传到 root\nlogger.info('start!')\nic.configureOutput(prefix='', outputFunction=logger.info)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T02:47:18.423534Z","iopub.execute_input":"2025-08-25T02:47:18.423836Z","iopub.status.idle":"2025-08-25T02:47:18.615945Z","shell.execute_reply.started":"2025-08-25T02:47:18.423812Z","shell.execute_reply":"2025-08-25T02:47:18.611571Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# FLAGS","metadata":{}},{"cell_type":"code","source":"class FLAGS:\n  root = '../input/aeroclub-recsys-2025'\n  # tree model here can be 'xgb' or 'lgb' \n  # not used here as I will present 4 xgb models only which coul get similar score offline and online\n  tm = 'xgb'\n  # objective here can be ndcg,map,pariwse for xgb and lambdarank,xendcg for lgb\n  obj = 'ndcg'\n  # task could be ranking or classification\n  task = 'ranking'\n  # this is a bit hack as I created 5 folds and only use fold 4 wich is train on day 1-103 valid on day 104-166\n  folds = 5\n  fold = 4\n  \n  cat_method = 'count'\n  # remove cast means all cat cols to be treated as numer cols \n  remove_cats = True\n  reserve_cats = False\n  # weather to use train/test for stats or only use train\n  stats_all = True\n  trees = 1000\n  seed = 42\n  \n  # if false only use original csv existed cols as feats\n  add_feats = True\n  \n  # weather to use history of selected data \n  # notice adding this could improve a lot on LB/PB but it is added after the game finished\n  # and code copy from https://www.kaggle.com/code/mikhailgolubchik/sm-xgboost-single\n  # also the feat generate using more time likely about 20-30 mintues and more memory needed about 400-500g\n  history_avg = True\n  # history_avg = False\n\n  # model_dir = '../input/aeroclub-recsys-2025-model1'\n  model_dir = '../input/aeroclub-recsys-2025-model2'\n  \n  external_dir = '../input/aeroclub-recsys-2025-external'\n  out_dir = '../working'\n    \n  # wether use json files, not affect much\n  use_ext = True\n  # use_ext = False\n  \n  # online = False means offline train/valid mode, online = True means train on all train data for submission\n  # for better LB/PB you need to set online = True, but if set online = True local valid is overly optimistic as we valid on data which also in train\n  # online = False\n  online = True\n  \n  # fast = True means for debug only which will run pipline using 0.01 ratio data and train model using 100 trees only\n  # fast = True\n  fast = False\n  \n  # n_models = 0 means not limit using all 4 models, n_models=1 means the best single model only with objective rank:ndcg\n  # n_models = 1\n  n_models = 0\n  \n  # mode = 'train' means train + eval + test this need CPU MEM more then 200-300g without his_avg\n  # mode = 'infer'/'test' means only load from pretrained model and do infer online\n  # mode = 'train'\n  mode = 'infer'\n  \n  device = 'gpu'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T02:47:18.617048Z","iopub.execute_input":"2025-08-25T02:47:18.617293Z","iopub.status.idle":"2025-08-25T02:47:18.629376Z","shell.execute_reply.started":"2025-08-25T02:47:18.617263Z","shell.execute_reply":"2025-08-25T02:47:18.624375Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def in_notebook():\n  try:\n    from IPython import get_ipython\n    if 'IPKernelApp' not in get_ipython().config:  # pragma: no cover\n      return False\n  except Exception as e:\n    return False\n  return True\n\nif not in_notebook():\n  if len(sys.argv) > 1:\n    FLAGS.fast = bool(int(sys.argv[1]))\n    FLAGS.online = bool(int(sys.argv[2]))\n    FLAGS.use_ext = bool(int(sys.argv[3]))\n    FLAGS.history_avg = bool(int(sys.argv[4]))\n  else:\n    if 'fast' in os.environ:\n      FLAGS.fast = bool(int(os.environ['fast']))\n    if 'online' in os.environ:\n      FLAGS.online = bool(int(os.environ['online']))\n    if 'use_ext' in os.environ:\n      FLAGS.use_ext = bool(int(os.environ['use_ext']))\n    if 'history_avg' in os.environ:\n      FLAGS.history_avg = bool(int(os.environ['history_avg']))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T02:47:18.631815Z","iopub.execute_input":"2025-08-25T02:47:18.632060Z","iopub.status.idle":"2025-08-25T02:47:19.000146Z","shell.execute_reply.started":"2025-08-25T02:47:18.632036Z","shell.execute_reply":"2025-08-25T02:47:18.995462Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# online + not use_ext + not use history_avg about 250G memory needed after preprocess and neeed 360G before xgb train\nFLAGS.out_dir = f'{FLAGS.out_dir}/fast{int(FLAGS.fast)}-online{int(FLAGS.online)}-use_ext{int(FLAGS.use_ext)}-history_avg{int(FLAGS.history_avg)}'\nic(FLAGS.fast)\nic(FLAGS.online)\nic(FLAGS.use_ext)\nic(FLAGS.history_avg)\nic(FLAGS.out_dir)\nos.system(f'mkdir -p {FLAGS.out_dir}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T02:47:19.001940Z","iopub.execute_input":"2025-08-25T02:47:19.002149Z","iopub.status.idle":"2025-08-25T02:47:19.397141Z","shell.execute_reply.started":"2025-08-25T02:47:19.002129Z","shell.execute_reply":"2025-08-25T02:47:19.392793Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"params_xgb = {\n    'objective': 'rank:ndcg',\n    'eval_metric': 'ndcg@3',\n    'max_depth': 12,\n    'min_child_weight': 10,\n    'subsample': 0.8,\n    'colsample_bytree': 0.8,\n    'lambda': 100,\n    'learning_rate': 0.05,\n    'n_estimators': 1000,\n    'seed': 42,\n}\n\n# this is what I used for lgb but it is fine to only train and ensemble xgb models\nparams_lgb = {\n      'objective': 'lambdarank',\n      'eval_metric': ['ndcg@3'],\n      'ndcg_eval_at': [3], \n      'eval_at': [3],  \n      'boosting_type': 'gbdt',\n      'num_leaves': 63,\n      'max_depth': 12,\n      'min_data_in_leaf': 50,\n      'feature_fraction': 0.8,\n      'bagging_fraction': 0.8,\n      'bagging_freq': 5,\n      'lambda_l1': 0.1,\n      'lambda_l2': 100,\n      'learning_rate': 0.05,\n      'n_estimators': 1000,\n      'seed': 42,\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T02:47:19.399141Z","iopub.execute_input":"2025-08-25T02:47:19.400502Z","iopub.status.idle":"2025-08-25T02:47:19.412543Z","shell.execute_reply.started":"2025-08-25T02:47:19.400475Z","shell.execute_reply":"2025-08-25T02:47:19.406640Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if FLAGS.fast:\n  FLAGS.trees = 100","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T02:47:19.413287Z","iopub.execute_input":"2025-08-25T02:47:19.413497Z","iopub.status.idle":"2025-08-25T02:47:19.650486Z","shell.execute_reply.started":"2025-08-25T02:47:19.413475Z","shell.execute_reply":"2025-08-25T02:47:19.646118Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Predefined cols","metadata":{}},{"cell_type":"code","source":"COLS_TO_COMPARE = [\n  \"legs0_departureAt\", \n  \"legs0_arrivalAt\", \n  \"legs1_departureAt\",\n  \"legs1_arrivalAt\", \n  \"legs0_segments0_flightNumber\",\n  \"legs1_segments0_flightNumber\"\n]\n\nIGNORE_COLS = [\n  'Id', \n  'ranker_id',\n  'selected',\n  'fold',\n  'uid',\n  'day',\n  'group_day',\n  'sample_weight',\n  'requestDate', #DateTime\n  'bySelf', #train only 1 unique value, test half 0 half 1\n  'pricingInfo_passengerCount', # nunique==1\n  'legs0_departureAt',  # original time will be converted\n  'legs0_arrivalAt',\n  'legs1_departureAt',\n  'legs1_arrivalAt',\n  'legs0_segments3_baggage_count',\n  'legs1_segments3_baggage_count',\n  'legs0_segments3_baggage_weight',\n  'legs1_segments3_baggage_weight',\n  'requestReturnDate',\n  'requestDepartureDate',\n  'legs0_segments3_baggageAllowance_weightMeasurementType',\n  'legs0_segments3_cabinClass',\n  'legs1_segments3_baggageAllowance_quantity',\n  'legs1_segments3_baggageAllowance_weightMeasurementType',\n  'legs1_segments3_cabinClass',\n  'legs1_segments3_seatsAvailable',\n  'legs1_seg3_dep_offset', \n  'legs0_departureAirport',\n  'legs1_departureAirport',\n  'legs1_seg3_arr_offset',\n  'isGlobal',\n  'flight_hash',\n]\n\nrank_order = {\n    'totalPrice': 'asc',\n    'flight_duration_total': 'asc',\n    'book_lead_time_hours': 'desc',\n    'flight_duration_travel_ratio': 'asc',\n    'seg_legs_all_count': 'asc',\n    'avg_cabin_legs_all': 'desc',\n    'avg_baggage_count_legs_all': 'desc',\n    'avg_baggage_weight_legs_all': 'desc',\n    'direct_price_per_km': 'asc',\n}\n\nsource_cols = [\n        'time_legs0_departureAt_hour',\n        'time_legs1_departureAt_hour',\n        'time_legs0_arrivalAt_hour',\n        'time_legs1_arrivalAt_hour',\n        'rank_totalPrice',\n        'rank_flight_duration_total',\n        'avg_cabin_legs_all',\n        'avg_baggage_count_legs_all',\n        'avg_baggage_weight_legs_all',\n        'direct_price_per_km',\n        'miniRules1_statusInfos',\n        'miniRules0_statusInfos',\n]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T02:47:19.652661Z","iopub.execute_input":"2025-08-25T02:47:19.652898Z","iopub.status.idle":"2025-08-25T02:47:19.812543Z","shell.execute_reply.started":"2025-08-25T02:47:19.652876Z","shell.execute_reply":"2025-08-25T02:47:19.807839Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# preprocess for raw data\nParse out requestDepartureDate requestReturnDate age from json files\nThough this not affect much, need 15-25 minutes","metadata":{}},{"cell_type":"code","source":"def load_external_data(n_files=0):\n  datas = []\n  json_files = glob.glob(f'{FLAGS.root}/raw/*.json')\n  \n  if FLAGS.fast:\n    ic('fast mode just use 10 json files')\n    json_files = json_files[:10]\n  \n  if n_files:\n    json_files = json_files[:n_files]\n\n  for json_file in tqdm(json_files, desc='json_files'):\n    with open(json_file) as fh:\n      data = json.load(fh)\n      datas.append(data)\n  \n  l = []\n  for data in tqdm(datas, desc='datas'):\n    m = {\n      'ranker_id': data['ranker_id']\n    }\n    routeData = data['routeData']\n    m.update({\n      'requestDepartureDate': routeData.get('requestDepartureDate', None),\n      'requestReturnDate': routeData.get('requestReturnDate', None),\n    })\n    personalData = data['personalData']\n    m['hasAssistant'] = personalData.get('hasAssistant', None)\n    m['isGlobal'] = personalData.get('isGlobal', None)\n    m['age'] = None\n    if 'yearOfBirth' in personalData:\n      try:\n        m['age'] = 2024 - personalData['yearOfBirth']\n      except Exception as e:\n        m['age'] = None\n    l.append(m)\n\n  ic('to df_ext')\n  df_ext = pl.DataFrame(l)\n  ic(df_ext['requestReturnDate'].n_unique())\n  return df_ext","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T02:47:19.814083Z","iopub.execute_input":"2025-08-25T02:47:19.814298Z","iopub.status.idle":"2025-08-25T02:47:19.999483Z","shell.execute_reply.started":"2025-08-25T02:47:19.814278Z","shell.execute_reply":"2025-08-25T02:47:19.993940Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def create_days(df):\n  if df[\"requestDate\"].dtype != pl.Datetime:\n    df = df.with_columns(\n        [pl.col(\"requestDate\").str.to_datetime().alias(\"requestDate\")])\n\n  min_date = df[\"requestDate\"].min()\n  max_date = df[\"requestDate\"].max()\n\n  ic(f\"Date range: {min_date} to {max_date}\")\n\n  df = df.with_columns([\n      # 计算相对天数（从1开始）\n      ((pl.col(\"requestDate\") - min_date).dt.total_days() + 1\n      ).cast(pl.Int32).alias(\"day\")\n  ])\n\n  max_day = df[\"day\"].max()\n  min_day = df[\"day\"].min()\n  ic(f\"Day range: {min_day} to {max_day}\")\n\n  group_day_mapping = (df.group_by(\"ranker_id\", maintain_order=True).agg(\n      pl.mean(\"day\").alias(\"group_day\")))\n\n  df = df.join(group_day_mapping, on=\"ranker_id\", how=\"left\")\n  return df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T02:47:20.002058Z","iopub.execute_input":"2025-08-25T02:47:20.002295Z","iopub.status.idle":"2025-08-25T02:47:20.249625Z","shell.execute_reply.started":"2025-08-25T02:47:20.002274Z","shell.execute_reply":"2025-08-25T02:47:20.245398Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_bool_cols(df: pl.DataFrame) -> list[str]:\n  return df.select(cs.boolean()).columns\n\ndef get_cat_cols(df: pl.DataFrame) -> list[str]:\n  cat_cols = df.select(cs.string() | cs.categorical()).columns\n  return cat_cols\n\ndef get_numer_cols(df: pl.DataFrame) -> list[str]:\n  num_cols = df.select(cs.numeric() | cs.boolean()).columns\n  return num_cols\n\ndef bool2int(df):\n  return df.with_columns([pl.col(c).cast(pl.Int8) for c in get_bool_cols(df)])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T02:47:20.251877Z","iopub.execute_input":"2025-08-25T02:47:20.252105Z","iopub.status.idle":"2025-08-25T02:47:20.276729Z","shell.execute_reply.started":"2025-08-25T02:47:20.252083Z","shell.execute_reply":"2025-08-25T02:47:20.271897Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def icl(lst, n=5):\n  if not ic.enabled:\n    return\n\n  import inspect\n\n  frame = inspect.currentframe().f_back\n  vars_dict = frame.f_locals.items()\n  # ic(vars_dict)\n\n  var_names = [name for name, value in vars_dict if value is lst]\n  var_name = var_names[0] if var_names else \"unknown\"\n\n  if isinstance(lst, list):\n    if len(lst) > n * 2:\n      logger.info(f'{var_name} first {n}: {lst[:n]}')\n      logger.info(f'{var_name} last {n}: {lst[-n:]}')\n    else:\n      logger.info(f'{var_name}: {lst}')\n    logger.info(f'len({var_name}): {len(lst)}')\n  else:\n    logger.info(f'{var_name}: {lst}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T02:47:20.278416Z","iopub.execute_input":"2025-08-25T02:47:20.278681Z","iopub.status.idle":"2025-08-25T02:47:20.308900Z","shell.execute_reply.started":"2025-08-25T02:47:20.278654Z","shell.execute_reply":"2025-08-25T02:47:20.303775Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def timeit(info=''):\n\n  def decorator(func):\n    @functools.wraps(func)\n    def wrapper(*args, **kwargs):\n      logger.info(f'{info} ---------------- {func.__name__} start')\n      start = time.time()\n      result = func(*args, **kwargs)\n      end = time.time()\n      logger.info(f'{info} ################ {func.__name__} elapsed: {end - start:.4f} seconds')\n      return result\n\n    return wrapper\n\n  return decorator\n\nimport functools\ndef monitor_feats(info='', out_index=0):\n  def decorator(func):\n    @functools.wraps(func)\n    def wrapper(df, *args, **kwargs):\n      original_columns = set(getattr(df, \"columns\", []))\n      ret = func(df, *args, **kwargs)\n\n      if isinstance(ret, (tuple, list)):\n        if out_index >= len(ret):\n          return ret\n        new_df = ret[out_index]\n      else:\n        new_df = ret\n\n      cols = getattr(new_df, \"columns\", None)\n      if cols is None:\n        return ret\n\n      new_cols = [col for col in cols if col not in original_columns]\n      logger.info(f\"{info} {func.__name__} added:\")\n      icl(new_cols, 10)\n\n      return ret\n    return wrapper\n  return decorator\n\n\ndef time_feats(info=''):\n  def decorator(func):\n    decorated_func = monitor_feats(info)(func)\n    decorated_func = timeit(info)(decorated_func)\n    return decorated_func\n  return decorator","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T02:47:20.311038Z","iopub.execute_input":"2025-08-25T02:47:20.311310Z","iopub.status.idle":"2025-08-25T02:47:20.332617Z","shell.execute_reply.started":"2025-08-25T02:47:20.311286Z","shell.execute_reply":"2025-08-25T02:47:20.328045Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Using hour not minutes","metadata":{}},{"cell_type":"code","source":"def dur_to_min(col: pl.Expr) -> pl.Expr:\n  # extract day（ '3.05:10:00' -> 3）\n  days = col.str.extract(r\"^(\\d+)\\.\", 1).cast(pl.Int64).fill_null(0) * 24 * 60\n\n  # extract time part\n  time_str = pl.when(col.str.contains(r\"^\\d+\\.\")) \\\n                .then(col.str.replace(r\"^\\d+\\.\", \"\")) \\\n                .otherwise(col)\n\n  hours = time_str.str.extract(r\"^(\\d+):\", 1).cast(pl.Int64).fill_null(0) * 60\n  minutes = time_str.str.extract(r\":(\\d+):\", 1).cast(pl.Int64).fill_null(0)\n\n  return (days + hours + minutes).fill_null(0)\n\n\n@timeit()\ndef durs_to_unit(df):\n  exprs = []\n\n  for leg in (0, 1):\n    col = f\"legs{leg}_duration\"\n    assert col in df.columns\n    exprs.append((dur_to_min(pl.col(col)) / 60))\n\n    for s in range(4):\n      col = f\"legs{leg}_segments{s}_duration\"\n      assert col in df.columns\n      exprs.append((dur_to_min(pl.col(col)) / 60))\n\n  df = df.with_columns(exprs)\n  return df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T02:47:20.335989Z","iopub.execute_input":"2025-08-25T02:47:20.336207Z","iopub.status.idle":"2025-08-25T02:47:20.368885Z","shell.execute_reply.started":"2025-08-25T02:47:20.336186Z","shell.execute_reply":"2025-08-25T02:47:20.364944Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_df(use_ext=True):\n  df_train = pl.read_parquet(f'{FLAGS.root}/train.parquet').drop('__index_level_0__')    \n  df_test = pl.read_parquet(f'{FLAGS.root}/test.parquet').drop('__index_level_0__')\n  df_test = df_test.with_columns(pl.lit(-1, dtype=pl.Int64).alias('selected'))\n  df = pl.concat([df_train, df_test], how='vertical')\n  df = create_days(df)\n  \n  df = df.with_columns(\n    pl.col('profileId').alias('uid')\n  )\n  \n  if use_ext:\n    if os.path.exists(FLAGS.external_dir):\n      logger.info('load external data from pre dumped external.parquet')\n      df_ext = pl.read_parquet(f'{FLAGS.external_dir}/external.parquet')\n    else:\n      logger.info('load external data from json files')\n      df_ext = load_external_data()\n    \n    ic(df_ext['requestReturnDate'].n_unique())\n    display(df_ext)\n    df = df.join(df_ext, on='ranker_id', how='left')\n    assert 'age' in df.columns\n\n  df = df.with_columns(\n    pl.col('profileId').cast(pl.Utf8)\n  )\n  \n  cat_cols = [\n    'profileId',\n    'companyID', \n    'corporateTariffCode', \n    'nationality',\n  ]\n    \n  df = df.with_columns(\n    pl.col(col).cast(pl.Utf8) for col in cat_cols \n  )\n \n  df = df.with_columns([\n      pl.concat_str([\n          pl.col(c).cast(str).fill_null(\"NULL\") \n          for c in COLS_TO_COMPARE\n      ]).alias(\"flight_hash\")\n  ])\n\n  df = bool2int(df)\n  df = durs_to_unit(df)\n  return df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T02:47:20.370956Z","iopub.execute_input":"2025-08-25T02:47:20.371186Z","iopub.status.idle":"2025-08-25T02:47:20.423115Z","shell.execute_reply.started":"2025-08-25T02:47:20.371164Z","shell.execute_reply":"2025-08-25T02:47:20.418760Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_valid(df):\n  df_valid = df.filter(pl.col('fold') <= FLAGS.fold)\n  return df_valid\n\ndef get_train(df):\n  if not FLAGS.online:\n    df_train = df.filter(pl.col('fold') > FLAGS.fold)\n  else:\n    df_train = df\n    \n  return df_train\n\ndef get_test(df):\n  return df.filter(pl.col('selected') == -1)\n\ndef get_nontest(df):\n  # Keep only non-test data\n  return df.filter(pl.col('selected') != -1)\n\ndef get_train_valid(df):\n  # leave out test\n  df = df.filter(pl.col('selected') != -1)\n\n  df_train = get_train(df)\n  df_valid = get_valid(df)\n\n  return df_train, df_valid","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T02:47:20.424102Z","iopub.execute_input":"2025-08-25T02:47:20.424322Z","iopub.status.idle":"2025-08-25T02:47:20.466717Z","shell.execute_reply.started":"2025-08-25T02:47:20.424303Z","shell.execute_reply":"2025-08-25T02:47:20.460954Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def set_fold(df):\n  df_train = df.filter(pl.col('selected') != -1)\n  df_test = df.filter(pl.col('selected') == -1)\n  # Group by ranker_id and find the latest requestDate for each\n  # Original code that sorts by request date\n  ranker_dates = (df_train\n            .group_by('ranker_id', maintain_order=True)\n            .agg(pl.col('requestDate').max().alias('last_request'))\n            .sort('last_request'))\n  \n  # Modified code that preserves original ranker_id order\n  unique_rankers = ranker_dates.select('ranker_id').unique(maintain_order=True)\n\n  # Total number of unique ranker_ids\n  n_rankers = unique_rankers.height\n  \n  # Initialize fold column with the last fold number\n  unique_rankers = unique_rankers.with_columns(pl.lit(FLAGS.folds).alias('fold'))\n  \n  # Calculate how many ranker_ids for each fold (10% per fold)\n  fold_size = int(n_rankers * 0.1)\n  ic(fold_size)\n  \n  # Update fold values for the newest (FLAGS.folds - 1) groups\n  for i in range(FLAGS.folds):\n    # Get ranker_ids for current fold (10% of data)\n    start_idx = n_rankers - (i + 1) * fold_size\n    end_idx = n_rankers - i * fold_size\n    len_ = end_idx - start_idx\n    # Get the list of ranker_ids for this fold\n    fold_rankers = unique_rankers.slice(start_idx, len_).get_column('ranker_id')\n    \n    # Update the fold value for these ranker_ids\n    unique_rankers = unique_rankers.with_columns(\n      pl.when(pl.col('ranker_id').is_in(fold_rankers))\n      .then(i)\n      .otherwise(pl.col('fold'))\n      .alias('fold')\n    )\n  \n  # Join fold assignments back to the training data\n  df_train = df_train.join(\n    unique_rankers.select('ranker_id', 'fold'),\n    on='ranker_id'\n  )\n\n  df_train = df_train.with_columns(pl.col('fold').cast(pl.Int32))\n  ic(len(df_train))\n  df_test = df_test.with_columns(pl.lit(-1, dtype=pl.Int32).alias('fold'))\n  df = pl.concat([df_train, df_test], how='vertical')\n\n  ic(df.group_by('fold').agg(pl.len()).sort('fold'))\n  \n  df_valid = get_valid(df_train)\n  ic(df_valid['day'].min(), df_valid['day'].max(), df_valid['day'].max() - df_valid['day'].min())\n  \n  return df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T02:47:20.468452Z","iopub.execute_input":"2025-08-25T02:47:20.468684Z","iopub.status.idle":"2025-08-25T02:47:20.492150Z","shell.execute_reply.started":"2025-08-25T02:47:20.468656Z","shell.execute_reply":"2025-08-25T02:47:20.487425Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def filter(df, ratio=0.01, seed=42):\n  unique_ids = df.select(\"ranker_id\").unique()\n  keep_ids = np.random.default_rng(seed).choice(\n    unique_ids[\"ranker_id\"].to_list(),\n    size=int(len(unique_ids) * ratio),  \n    replace=False\n  )\n  df = df.filter(pl.col(\"ranker_id\").is_in(keep_ids))\n  return df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T02:47:20.494463Z","iopub.execute_input":"2025-08-25T02:47:20.494719Z","iopub.status.idle":"2025-08-25T02:47:20.507560Z","shell.execute_reply.started":"2025-08-25T02:47:20.494681Z","shell.execute_reply":"2025-08-25T02:47:20.502892Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@timeit()\ndef smart_fillnull(df: pl.DataFrame, numer_cols: list[str], cat_cols: list[str]) -> pl.DataFrame:\n  zero_flags = df.select([\n      (pl.col(col) == 0).any().alias(col)\n      for col in numer_cols\n  ])\n\n  fill_map = {\n      col: -1 if zero_flags[col][0] else 0\n      for col in numer_cols\n  }\n\n  exprs = [\n      pl.col(col).fill_null(fill_map[col]) for col in numer_cols\n  ] + [\n      pl.col(col).fill_null(\"missing\") for col in cat_cols\n  ]\n\n  return df.with_columns(exprs)\n\ndef promote_dtype(dtypes):  \n  unique_types = set(dtypes)\n  is_float = lambda dt: pl.datatypes.is_float_dtype(dt)\n  is_int = lambda dt: pl.datatypes.is_integer_dtype(dt)\n\n  # int + float\n  if any(is_float(dt) for dt in unique_types) and any(is_int(dt) for dt in unique_types):\n    max_float_bits = max(dt.bit_width for dt in unique_types if is_float(dt))\n    return pl.Float64 if max_float_bits > 32 else pl.Float32\n\n  # all float\n  if all(is_float(dt) for dt in unique_types):\n    max_bits = max(dt.bit_width for dt in unique_types)\n    return pl.Float64 if max_bits > 32 else pl.Float32\n\n  # all int\n  if all(is_int(dt) for dt in unique_types):\n    max_bits = max(dt.bit_width for dt in unique_types)\n    return {8: pl.Int8, 16: pl.Int16, 32: pl.Int32, 64: pl.Int64}[max_bits]\n\n  # fallback\n  return list(unique_types)[0]\n\n\ndef align_and_concat(dfs, verbose=True):\n  from builtins import set\n\n  if not dfs:\n    raise ValueError(\"DataFrame empty\")\n\n  common_cols = set.intersection(*(set(df.columns) for df in dfs))\n\n  for col in sorted(common_cols):\n    dtypes = [df[col].dtype for df in dfs]\n    if len(set(dtypes)) > 1:\n      if verbose:\n        print(f\"[col type not same] {col}\")\n        for i, dt in enumerate(dtypes):\n          print(f\"  DF[{i}] dtype: {dt}\")\n      target_type = promote_dtype(dtypes)\n      if verbose:\n        print(f\"  → convert to: {target_type}\\n\")\n      dfs = [df.with_columns(pl.col(col).cast(target_type)) for df in dfs]\n\n  return pl.concat(dfs, how=\"vertical\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T02:47:20.509925Z","iopub.execute_input":"2025-08-25T02:47:20.510155Z","iopub.status.idle":"2025-08-25T02:47:20.528399Z","shell.execute_reply.started":"2025-08-25T02:47:20.510133Z","shell.execute_reply":"2025-08-25T02:47:20.522788Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Counting based category encoding","metadata":{}},{"cell_type":"code","source":"unified_features = {\n    'airport_iata': [],\n    'airport_city_iata': [],\n    'marketingCarrier_code': [],\n    'operatingCarrier_code': [],\n    'aircraft_code': [],\n    'flightNumber': [],\n    'searchRoute': [],\n    'statusInfos': [],\n    # 'baggageAllowance_weightMeasurementType': [],\n    # # 'baggageAllowance_quantity': [],\n    # 'cabinClass': [],\n}\n\ndef get_unified_cat(col):\n  for key in unified_features:\n    if key in col:\n      return key\n\n  return col\n\n\ndef get_unified_cat_columns(cat_cols):\n  for col in cat_cols:\n    for key in unified_features:\n      if key in col:\n        unified_features[key].append(col)\n\n  return unified_features\n\n@timeit()\ndef encode_unified_cats(df_train,\n                        unified_features,\n                        id_col=None,\n                        method='count',\n                        num_workers=1):\n  unified_cats = OrderedDict()\n\n  for feature_type, columns in tqdm(unified_features.items(), desc='encode_unified_cats'):\n    if not columns:\n      continue\n\n    all_values = []\n    for col in columns:\n      if col in df_train.columns:\n        all_values.extend(\n            df_train.select(pl.col(col)).drop_nulls().to_series().to_list())\n\n    if method == 'count':\n      from collections import Counter\n      value_counts = Counter(all_values)\n      sorted_values = [val for val, count in value_counts.most_common()]\n      feature_dict = {val: idx for idx, val in enumerate(sorted_values)}\n    elif method == 'val':\n      unique_values = sorted(list(set(all_values)))\n      feature_dict = {val: idx for idx, val in enumerate(unique_values)}\n    elif method == 'seq':\n      seen = set()\n      unique_values = []\n      for val in all_values:\n        if val not in seen:\n          unique_values.append(val)\n          seen.add(val)\n      feature_dict = {val: idx for idx, val in enumerate(unique_values)}\n\n    unified_cats[feature_type] = feature_dict\n\n  return unified_cats\n\n@timeit()\ndef encode_cat_byseq(df_train, cat_cols, id_col=None):\n  cats = OrderedDict()\n  if id_col:\n    df_train = sort_dataframe(df_train, id_col)\n\n  for col in tqdm(cat_cols, desc='encode_cat_byseq'):\n    categories = df_train[col].unique(maintain_order=True).to_list()\n    category_dict = {category: idx for idx, category in enumerate(categories)}\n    cats.update({col: category_dict})\n\n  return cats\n\n@timeit()\ndef encode_cat_bycount(df_train, cat_cols, id_col=None, num_workers=1):\n  cats = OrderedDict()\n\n  def process_column(col):\n    dg = df_train.group_by(col).agg(pl.len().alias('count')).sort(\n        'count', descending=True)\n\n    return {col: {val: idx for idx, val in enumerate(dg[col].to_list())}}\n\n  results = []\n  for col in tqdm(cat_cols, desc='encode_cat_bycount'):\n    results.append(process_column(col))\n\n  for result in results:\n    cats.update(result)\n\n  return cats\n\n@timeit()\ndef encode_cat_byval(df_train, cat_cols, num_workers=1):\n  cats = OrderedDict()\n\n  def process_column(col):\n    categories = df_train.select(\n        pl.col(col)).unique().sort(col).to_series().to_list()\n    return {col: {val: idx for idx, val in enumerate(categories)}}\n\n  for cat in tqdm(cat_cols, desc='encode_cat_byval'):\n    results.append(process_column(cat))\n\n  for result in results:\n    cats.update(result)\n\n  return cats\n\n@timeit()\ndef encode_cat(df_train, cat_cols, id_col=None, method='count', num_workers=1):\n  cats = None\n  if method == 'seq':\n    cats = encode_cat_byseq(df_train, cat_cols, id_col)\n  elif method == 'count':\n    cats = encode_cat_bycount(df_train, cat_cols, id_col, num_workers=num_workers)\n  elif method == 'val':\n    # same as rank('dense') in polars\n    cats = encode_cat_byval(df_train, cat_cols, num_workers=num_workers)\n  else:\n    raise ValueError(f\"Unsupported method: {method}\")\n  \n  return cats\n\n@timeit()\ndef encode_cat_unified(df_train,\n                       cat_cols,\n                       id_col=None,\n                       method='count',\n                       num_workers=1):\n\n  unified_features = get_unified_cat_columns(cat_cols)\n\n  unified_cols = set()\n  for columns in unified_features.values():\n    unified_cols.update(columns)\n\n  unified_cat_cols = [col for col in cat_cols if col in unified_cols]\n  regular_cat_cols = [col for col in cat_cols if col not in unified_cols]\n\n  unified_cats = encode_unified_cats(df_train, unified_features, id_col, method, num_workers)\n\n  regular_cats = encode_cat(df_train, regular_cat_cols, id_col, method,num_workers)\n\n  all_cats = OrderedDict()\n  all_cats.update(unified_cats)\n  all_cats.update(regular_cats)\n\n  return all_cats","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T02:47:20.530851Z","iopub.execute_input":"2025-08-25T02:47:20.531083Z","iopub.status.idle":"2025-08-25T02:47:20.558804Z","shell.execute_reply.started":"2025-08-25T02:47:20.531062Z","shell.execute_reply":"2025-08-25T02:47:20.554249Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Manual feats","metadata":{}},{"cell_type":"code","source":"@time_feats()\ndef add_group_feats(df: pl.DataFrame) -> pl.DataFrame:\n  df = df.with_columns([\n      pl.col(\"Id\").count().over(\"ranker_id\").alias(\n          \"group_size\"),  \n  ])\n  return df\n\n@time_feats()\ndef add_user_feats(df: pl.DataFrame) -> pl.DataFrame:\n  df = df.with_columns(\n      [pl.col(\"frequentFlyer\").fill_null(\"\").alias(\"frequentFlyer\")])\n\n  unique_combos = (\n      df.select(\"frequentFlyer\").unique().get_column(\"frequentFlyer\").to_list())\n\n  all_codes = sorted(\n      set(chain.from_iterable(\n          code.split(\"/\") for code in unique_combos if code)))\n\n  for code in all_codes:\n    df = df.with_columns([\n        pl.col(\"frequentFlyer\").str.split(\"/\").list.contains(code).cast(\n            pl.Int8).alias(f\"ff_{code}\")\n    ])\n\n  df = df.with_columns(\n      [pl.col(\"frequentFlyer\").str.split(\"/\").list.get(0).alias(\"ff_primary\")])\n\n  df = df.with_columns(\n      [pl.col(\"frequentFlyer\").str.count_matches(\"/\").add(1).alias(\"ff_count\")])\n\n  df = df.with_columns([\n      pl.when(pl.col(\"frequentFlyer\") == \"\").then(0).otherwise(\n          pl.col(\"ff_count\")).alias(\"ff_count\")\n  ])\n\n  df = df.drop(\"frequentFlyer\")\n  return df\n\n@time_feats()\ndef add_searchRoute_feats(df):\n  df = df.with_columns([\n      pl.col(\"searchRoute\").str.contains(\"/\").not_().cast(\n          pl.Int8).alias(\"isDirect\")\n  ])\n\n  df = df.with_columns([\n      pl.col(\"searchRoute\").str.split(\"/\").list.first().str.slice(\n          0, 6).alias(\"base_searchRoute\")\n  ]).with_columns([\n      pl.col(\"base_searchRoute\").str.slice(0, 3).alias(\"p1\"),\n      pl.col(\"base_searchRoute\").str.slice(3, 3).alias(\"p2\"),\n  ]).with_columns([\n      pl.when(pl.col(\"p1\") <= pl.col(\"p2\")).then(\n          pl.concat_str([pl.col(\"p1\"), pl.col(\"p2\")])).otherwise(\n              pl.concat_str([pl.col(\"p2\"),\n                             pl.col(\"p1\")])).alias(\"normed_searchRoute\")\n  ]).drop([\"p1\", \"p2\"])\n\n  return df\n\n@time_feats()\ndef add_segment_feats(df: pl.DataFrame) -> pl.DataFrame:\n  exprs = []\n  for leg in (0, 1):\n    seg_cols = [\n        f\"legs{leg}_segments{s}_duration\" for s in range(4)\n        if f\"legs{leg}_segments{s}_duration\" in df.columns\n    ]\n    assert seg_cols, f\"legs{leg} seg_cols is empty\"\n    exprs.append(\n        pl.sum_horizontal([\n            (pl.col(c) > 0).cast(pl.UInt8) for c in seg_cols\n        ]).cast(pl.Int32).alias(f\"seg_legs{leg}_count\"))\n\n  df = df.with_columns(exprs)\n\n  df = df.with_columns([\n      pl.sum_horizontal([\n          pl.col(c) for c in [\"seg_legs0_count\", \"seg_legs1_count\"]\n      ]).alias(\"seg_legs_all_count\"),\n  ])\n\n  df = df.with_columns([\n      pl.col(\"legs0_segments0_departureFrom_airport_iata\").alias(\n          \"legs0_departureAirport\"),\n      pl.col(\"legs1_segments0_departureFrom_airport_iata\").alias(\n          \"legs1_departureAirport\"),\n      pl.when(pl.col(\"seg_legs0_count\") == 1\n             ).then(pl.col(\"legs0_segments0_arrivalTo_airport_iata\")\n                   ).when(pl.col(\"seg_legs0_count\") == 2).then(\n                       pl.col(\"legs0_segments1_arrivalTo_airport_iata\")\n                   ).when(pl.col(\"seg_legs0_count\") == 3).then(\n                       pl.col(\"legs0_segments2_arrivalTo_airport_iata\")\n                   ).when(pl.col(\"seg_legs0_count\") == 4).then(\n                       pl.col(\"legs0_segments3_arrivalTo_airport_iata\")\n                   ).otherwise(None).alias(\"legs0_arrival_airport_iata\"),\n\n      pl.when(pl.col(\"seg_legs1_count\") == 1\n             ).then(pl.col(\"legs1_segments0_arrivalTo_airport_iata\")\n                   ).when(pl.col(\"seg_legs1_count\") == 2).then(\n                       pl.col(\"legs1_segments1_arrivalTo_airport_iata\")\n                   ).when(pl.col(\"seg_legs1_count\") == 3).then(\n                       pl.col(\"legs1_segments2_arrivalTo_airport_iata\")\n                   ).when(pl.col(\"seg_legs1_count\") == 4).then(\n                       pl.col(\"legs1_segments3_arrivalTo_airport_iata\")\n                   ).otherwise(None).alias(\"legs1_arrival_airport_iata\"),\n  ])\n\n  return df\n\n@time_feats()\ndef add_flight_duration_feats(df: pl.DataFrame) -> pl.DataFrame:\n  df = df.with_columns([\n      (pl.col(\"legs0_duration\") +\n       pl.col(\"legs1_duration\")).alias(\"flight_duration_total\"),\n  ])\n  df = df.with_columns([\n      (pl.col(\"legs0_duration\") /\n       (pl.col(\"flight_duration_total\") + 1e-5)).alias(\"legs0_duration_ratio\"),\n      (pl.col(\"legs1_duration\") /\n       (pl.col(\"flight_duration_total\") + 1e-5)).alias(\"legs1_duration_ratio\"),\n  ])\n  return df\n\ndef get_utc_offset(timezone_str):\n  try:\n    tz = ZoneInfo(timezone_str)\n    now = datetime.now(tz)\n    offset = now.utcoffset().total_seconds() / 3600\n    return offset\n  except Exception as e:\n    print(f\"Invalid timezone: {timezone_str} → {e}\")\n    return None\n\n\nAIRPORTS_DB = None\nTIMEZONE_FINDER = None\n\n@timeit()\ndef get_airport_timezone_mapping():\n  global AIRPORTS_DB, TIMEZONE_FINDER\n\n  if AIRPORTS_DB is None:\n    AIRPORTS_DB = airportsdata.load('IATA')\n    TIMEZONE_FINDER = TimezoneFinder()\n\n  airport_tz_map = {}\n  for iata, info in AIRPORTS_DB.items():\n    if info.get('lat') and info.get('lon'):\n      timezone = TIMEZONE_FINDER.timezone_at(lat=info['lat'], lng=info['lon'])\n      if timezone:\n        airport_tz_map[iata] = timezone\n      else:\n        airport_tz_map[iata] = 'UTC'\n    else:\n      airport_tz_map[iata] = 'UTC'\n\n  return airport_tz_map\n\ndef haversine_distance_expr(lat1, lon1, lat2, lon2):\n  lat1_rad = lat1 * (math.pi / 180)\n  lon1_rad = lon1 * (math.pi / 180)\n  lat2_rad = lat2 * (math.pi / 180)\n  lon2_rad = lon2 * (math.pi / 180)\n\n  dlat = lat2_rad - lat1_rad\n  dlon = lon2_rad - lon1_rad\n\n  a = (dlat / 2\n      ).sin().pow(2) + lat1_rad.cos() * lat2_rad.cos() * (dlon / 2).sin().pow(2)\n  c = 2 * a.sqrt().arcsin()\n\n  return c * 6371  \n\n@time_feats()\ndef add_segment_geography_time_feats(df: pl.DataFrame) -> pl.DataFrame:\n  logger.info('Adding segment geography and time features - merged version')\n\n  global AIRPORTS_DB\n  if AIRPORTS_DB is None:\n    AIRPORTS_DB = airportsdata.load('IATA')\n\n  airport_info_data = []\n  for iata, info in AIRPORTS_DB.items():\n    timezone_str = info.get('tz', 'UTC')\n    airport_info_data.append({\n        'airport_code': iata,\n        # 'country': info.get('country', ''),\n        # 'city': info.get('city', ''),\n        'lat': info.get('lat'),\n        'lon': info.get('lon'),\n        # 'elevation': info.get('elevation'),\n        'timezone': timezone_str,\n        'utc_offset_hours': get_utc_offset(timezone_str)\n    })\n\n  airport_info_df = pl.DataFrame(airport_info_data)\n\n  geo_time_exprs = []\n\n  for leg in [0, 1]:\n    for seg in range(4):\n      dep_airport_col = f\"legs{leg}_segments{seg}_departureFrom_airport_iata\"\n      arr_airport_col = f\"legs{leg}_segments{seg}_arrivalTo_airport_iata\"\n      duration_col = f\"legs{leg}_segments{seg}_duration\"\n\n      if all(col in df.columns\n             for col in [dep_airport_col, arr_airport_col, duration_col]):\n        df = df.join(\n            airport_info_df.select([\n                'airport_code',\n                # 'country',\n                # 'city',\n                'lat',\n                'lon',\n                'utc_offset_hours'\n            ]).rename({\n                'airport_code': dep_airport_col,\n                # 'country': f\"legs{leg}_seg{seg}_dep_country\",\n                # 'city': f\"legs{leg}_seg{seg}_dep_city\",\n                'lat': f\"legs{leg}_seg{seg}_dep_lat\",\n                'lon': f\"legs{leg}_seg{seg}_dep_lon\",\n                'utc_offset_hours': f\"legs{leg}_seg{seg}_dep_offset\"\n            }),\n            on=dep_airport_col,\n            how='left'\n        ).with_columns([\n            pl.when((pl.col(dep_airport_col).is_not_null()) &\n                    (pl.col(duration_col) > 0)).then(\n                        pl.col(f\"legs{leg}_seg{seg}_dep_offset\").fill_null(0)\n                    ).otherwise(None).alias(f\"legs{leg}_seg{seg}_dep_offset\"),\n            # pl.when((pl.col(dep_airport_col).is_not_null()) & (pl.col(duration_col) > 0))\n            #   .then(pl.col(f\"legs{leg}_seg{seg}_dep_country\"))\n            #   .otherwise(None)\n            #   .alias(f\"legs{leg}_seg{seg}_dep_country\"),\n            # pl.when((pl.col(dep_airport_col).is_not_null()) & (pl.col(duration_col) > 0))\n            #   .then(pl.col(f\"legs{leg}_seg{seg}_dep_city\"))\n            #   .otherwise(None)\n            #   .alias(f\"legs{leg}_seg{seg}_dep_city\"),\n            pl.when((pl.col(dep_airport_col).is_not_null()) &\n                    (pl.col(duration_col) > 0)).then(\n                        pl.col(f\"legs{leg}_seg{seg}_dep_lat\")\n                    ).otherwise(None).alias(f\"legs{leg}_seg{seg}_dep_lat\"),\n            pl.when((pl.col(dep_airport_col).is_not_null()) &\n                    (pl.col(duration_col) > 0)).then(\n                        pl.col(f\"legs{leg}_seg{seg}_dep_lon\")).otherwise(\n                            None).alias(f\"legs{leg}_seg{seg}_dep_lon\"),\n        ])\n\n        df = df.join(\n            airport_info_df.select([\n                'airport_code',\n                # 'country',\n                # 'city',\n                'lat',\n                'lon',\n                'utc_offset_hours'\n            ]).rename({\n                'airport_code': arr_airport_col,\n                # 'country': f\"legs{leg}_seg{seg}_arr_country\",\n                # 'city': f\"legs{leg}_seg{seg}_arr_city\",\n                'lat': f\"legs{leg}_seg{seg}_arr_lat\",\n                'lon': f\"legs{leg}_seg{seg}_arr_lon\",\n                'utc_offset_hours': f\"legs{leg}_seg{seg}_arr_offset\"\n            }),\n            on=arr_airport_col,\n            how='left'\n        ).with_columns([\n            pl.when((pl.col(arr_airport_col).is_not_null()) &\n                    (pl.col(duration_col) > 0)).then(\n                        pl.col(f\"legs{leg}_seg{seg}_arr_offset\").fill_null(0)).\n            otherwise(None).alias(f\"legs{leg}_seg{seg}_arr_offset\"),\n            # pl.when((pl.col(arr_airport_col).is_not_null()) & (pl.col(duration_col) > 0))\n            #   .then(pl.col(f\"legs{leg}_seg{seg}_arr_country\"))\n            #   .otherwise(None)\n            #   .alias(f\"legs{leg}_seg{seg}_arr_country\"),\n            # pl.when((pl.col(arr_airport_col).is_not_null()) & (pl.col(duration_col) > 0))\n            #   .then(pl.col(f\"legs{leg}_seg{seg}_arr_city\"))\n            #   .otherwise(None)\n            #   .alias(f\"legs{leg}_seg{seg}_arr_city\"),\n            pl.when((pl.col(arr_airport_col).is_not_null()) &\n                    (pl.col(duration_col) > 0)).then(\n                        pl.col(f\"legs{leg}_seg{seg}_arr_lat\")\n                    ).otherwise(None).alias(f\"legs{leg}_seg{seg}_arr_lat\"),\n            pl.when((pl.col(arr_airport_col).is_not_null()) &\n                    (pl.col(duration_col) > 0)).then(\n                        pl.col(f\"legs{leg}_seg{seg}_arr_lon\")).otherwise(\n                            None).alias(f\"legs{leg}_seg{seg}_arr_lon\"),\n        ])\n\n        valid_coords_condition = (\n            pl.col(f\"legs{leg}_seg{seg}_dep_lat\").is_not_null() &\n            pl.col(f\"legs{leg}_seg{seg}_dep_lon\").is_not_null() &\n            pl.col(f\"legs{leg}_seg{seg}_arr_lat\").is_not_null() &\n            pl.col(f\"legs{leg}_seg{seg}_arr_lon\").is_not_null() &\n            (pl.col(duration_col) > 0))\n\n        valid_offset_condition = (\n            (pl.col(duration_col) > 0) &\n            (pl.col(f\"legs{leg}_seg{seg}_dep_offset\").is_not_null()) &\n            (pl.col(f\"legs{leg}_seg{seg}_arr_offset\").is_not_null()))\n\n        geo_time_exprs.extend([\n            pl.when(valid_coords_condition).then(\n                haversine_distance_expr(pl.col(f\"legs{leg}_seg{seg}_dep_lat\"),\n                                        pl.col(f\"legs{leg}_seg{seg}_dep_lon\"),\n                                        pl.col(f\"legs{leg}_seg{seg}_arr_lat\"),\n                                        pl.col(f\"legs{leg}_seg{seg}_arr_lon\"))\n            ).otherwise(0.0).alias(f\"legs{leg}_seg{seg}_distance_km\"),\n\n            # pl.when(valid_coords_condition)\n            #   .then((pl.col(f\"legs{leg}_seg{seg}_dep_country\") != pl.col(f\"legs{leg}_seg{seg}_arr_country\")).cast(pl.Int8))\n            #   .otherwise(0)\n            #   .alias(f\"legs{leg}_seg{seg}_is_international\"),\n\n            # pl.when(valid_coords_condition)\n            #   .then((pl.col(f\"legs{leg}_seg{seg}_dep_city\") == pl.col(f\"legs{leg}_seg{seg}_arr_city\")).cast(pl.Int8))\n            #   .otherwise(0)\n            #   .alias(f\"legs{leg}_seg{seg}_is_same_city\"),\n        ])\n\n  legs_airport_pairs = [\n      ('legs0_departureAirport', 'legs0_dep_country', 'legs0_dep_city',\n       'legs0_dep_lat', 'legs0_dep_lon', 'legs0_dep_offset'),\n      ('legs0_arrival_airport_iata', 'legs0_arr_country', 'legs0_arr_city',\n       'legs0_arr_lat', 'legs0_arr_lon', 'legs0_arr_offset'),\n      ('legs1_departureAirport', 'legs1_dep_country', 'legs1_dep_city',\n       'legs1_dep_lat', 'legs1_dep_lon', 'legs1_dep_offset'),\n      ('legs1_arrival_airport_iata', 'legs1_arr_country', 'legs1_arr_city',\n       'legs1_arr_lat', 'legs1_arr_lon', 'legs1_arr_offset'),\n  ]\n\n  for airport_col, country_col, city_col, lat_col, lon_col, offset_col in legs_airport_pairs:\n    if airport_col in df.columns:\n      df = df.join(\n          airport_info_df.select([\n              'airport_code',\n              # 'country',\n              # 'city',\n              'lat',\n              'lon',\n              'utc_offset_hours'\n          ]).rename({\n              'airport_code': airport_col,\n              # 'country': country_col,\n              # 'city': city_col,\n              'lat': lat_col,\n              'lon': lon_col,\n              'utc_offset_hours': offset_col,\n          }),\n          on=airport_col,\n          how='left').with_columns([\n              pl.when(pl.col(airport_col).is_not_null()\n                     ).then(pl.col(offset_col).fill_null(0)\n                           ).otherwise(0).alias(offset_col)\n          ])\n\n  if geo_time_exprs:\n    df = df.with_columns(geo_time_exprs)\n\n  legs_geo_exprs = []\n\n  if all(col in df.columns for col in\n         ['legs0_dep_lat', 'legs0_dep_lon', 'legs0_arr_lat', 'legs0_arr_lon']):\n    legs_geo_exprs.extend([\n        pl.when(\n            pl.col('legs0_dep_lat').is_not_null() &\n            pl.col('legs0_dep_lon').is_not_null() &\n            pl.col('legs0_arr_lat').is_not_null() &\n            pl.col('legs0_arr_lon').is_not_null()).then(\n                haversine_distance_expr(pl.col('legs0_dep_lat'),\n                                        pl.col('legs0_dep_lon'),\n                                        pl.col('legs0_arr_lat'),\n                                        pl.col('legs0_arr_lon'))\n            ).otherwise(None).alias('legs0_direct_distance_km'),\n\n        # (pl.col('legs0_dep_country') != pl.col('legs0_arr_country')).cast(pl.Int8).alias('legs0_is_international'),\n\n        # (pl.col('legs0_dep_city') == pl.col('legs0_arr_city')).cast(pl.Int8).alias('legs0_is_same_city'),\n    ])\n\n  if all(col in df.columns for col in\n         ['legs1_dep_lat', 'legs1_dep_lon', 'legs1_arr_lat', 'legs1_arr_lon']):\n    legs_geo_exprs.extend([\n        # legs1直线距离\n        pl.when(\n            pl.col('legs1_dep_lat').is_not_null() &\n            pl.col('legs1_dep_lon').is_not_null() &\n            pl.col('legs1_arr_lat').is_not_null() &\n            pl.col('legs1_arr_lon').is_not_null()).then(\n                haversine_distance_expr(pl.col('legs1_dep_lat'),\n                                        pl.col('legs1_dep_lon'),\n                                        pl.col('legs1_arr_lat'),\n                                        pl.col('legs1_arr_lon'))\n            ).otherwise(0.0).alias('legs1_direct_distance_km'),\n\n        # (pl.col('legs1_dep_country') != pl.col('legs1_arr_country')).cast(pl.Int8).alias('legs1_is_international'),\n\n        # (pl.col('legs1_dep_city') == pl.col('legs1_arr_city')).cast(pl.Int8).alias('legs1_is_same_city'),\n    ])\n\n  if legs_geo_exprs:\n    df = df.with_columns(legs_geo_exprs)\n\n  assert 'legs0_seg0_distance_km' in df.columns\n\n  summary_exprs = []\n\n  for leg in [0, 1]:\n    seg_distance_cols = [\n        f\"legs{leg}_seg{seg}_distance_km\" for seg in range(4)\n        if f\"legs{leg}_seg{seg}_distance_km\" in df.columns\n    ]\n    assert seg_distance_cols\n    summary_exprs.append(\n        pl.sum_horizontal([pl.col(col) for col in seg_distance_cols\n                          ]).alias(f\"legs{leg}_total_segment_distance_km\"))\n\n  df = df.with_columns(summary_exprs)\n  assert 'legs0_total_segment_distance_km' in df.columns\n\n  summary_exprs = []\n  for leg in [0, 1]:\n    summary_exprs.append((pl.col(f\"legs{leg}_total_segment_distance_km\") /\n                          (pl.col(f\"legs{leg}_direct_distance_km\") +\n                           1)).alias(f\"legs{leg}_detour_ratio\"))\n\n    # seg_distance_cols = [f\"legs{leg}_seg{seg}_distance_km\" for seg in range(4)\n    #                       if f\"legs{leg}_seg{seg}_distance_km\" in df.columns]\n\n    # # if len(seg_distance_cols) > 1:\n    # summary_exprs.append(\n    #   pl.concat_list([pl.col(col) for col in seg_distance_cols])\n    #     .list.eval(pl.element().filter(pl.element() > 0))  # 过滤掉0距离\n    #     .list.std()\n    #     .alias(f\"legs{leg}_segment_distance_std\")\n    # )\n\n    # international_cols = [f\"legs{leg}_seg{seg}_is_international\" for seg in range(4)]\n    # # if international_cols:\n    # summary_exprs.append(\n    #   pl.sum_horizontal([pl.col(col) for col in international_cols])\n    #     .alias(f\"legs{leg}_international_segments_count\")\n    # )\n\n  df = df.with_columns(summary_exprs)\n\n  assert 'legs0_total_segment_distance_km' in df.columns\n\n  df = df.with_columns([\n      (pl.col(\"legs0_total_segment_distance_km\") +\n       pl.col(\"legs1_total_segment_distance_km\").fill_null(0)\n      ).alias(\"total_flight_distance_km\"),\n      (pl.col(\"legs0_direct_distance_km\") +\n       pl.col(\"legs1_direct_distance_km\").fill_null(0)\n      ).alias(\"direct_flight_distance_km\"),\n  ])\n\n  df = df.with_columns([\n      (pl.col(\"total_flight_distance_km\") /\n       (pl.col(\"flight_duration_total\") + 1e-5)).alias(\"avg_flight_speed_kmh\"),\n\n      (pl.col(\"totalPrice\") / (pl.col(\"total_flight_distance_km\") + 1)\n      ).alias(\"flight_price_per_km\"),\n\n      (pl.col(\"direct_flight_distance_km\") /\n       (pl.col(\"flight_duration_total\") + 1e-5)).alias(\"avg_direct_speed_kmh\"),\n\n      (pl.col(\"totalPrice\") / (pl.col(\"direct_flight_distance_km\") + 1)\n      ).alias(\"direct_price_per_km\"),\n  ])\n\n  # region_exprs = []\n  # country_cols = []\n  # for leg in [0, 1]:\n  #   country_cols.extend([f\"legs{leg}_dep_country\", f\"legs{leg}_arr_country\"])\n  #   for seg in range(4):\n  #     country_cols.extend([f\"legs{leg}_seg{seg}_dep_country\", f\"legs{leg}_seg{seg}_arr_country\"])\n\n  # region_exprs.append(\n  #   pl.concat_list([pl.col(col) for col in country_cols if col in df.columns])\n  #     .list.drop_nulls()\n  #     .list.unique()\n  #     .list.len()\n  #     .alias(\"total_unique_countries\")\n  # )\n\n  # df = df.with_columns(region_exprs)\n\n  return df\n\n@time_feats()\ndef add_travel_duration_feats(df: pl.DataFrame) -> pl.DataFrame:\n  utc_time_exprs = []\n  time_offset_pairs = [\n      ('legs0_departureAt', 'legs0_dep_offset', 'legs0_departure_utc'),\n      ('legs0_arrivalAt', 'legs0_arr_offset', 'legs0_arrival_utc'),\n      ('legs1_departureAt', 'legs1_dep_offset', 'legs1_departure_utc'),\n      ('legs1_arrivalAt', 'legs1_arr_offset', 'legs1_arrival_utc'),\n  ]\n\n  for time_col, offset_col, utc_col in time_offset_pairs:\n    if time_col in df.columns and offset_col in df.columns:\n      utc_time_exprs.append(\n          (pl.col(time_col).str.to_datetime() -\n           pl.duration(hours=pl.col(offset_col))).alias(utc_col))\n\n  df = df.with_columns(utc_time_exprs)\n\n  exprs = []\n\n  exprs.extend([\n      ((pl.col(\"legs0_arrival_utc\") -\n        pl.col(\"legs0_departure_utc\")).dt.total_seconds() /\n       3600).alias(\"travel_duration_legs0\"),\n      ((pl.col(\"legs1_arrival_utc\") -\n        pl.col(\"legs1_departure_utc\")).dt.total_seconds() /\n       3600).alias(\"travel_duration_legs1\"),\n      ((pl.col(\"legs1_departure_utc\") -\n        pl.col(\"legs0_arrival_utc\")).dt.total_seconds() / 3600\n      ).alias(\"travel_connection_duration\"),\n      ((pl.col(\"legs1_arrival_utc\") -\n        pl.col(\"legs0_departure_utc\")).dt.total_seconds() / 3600\n      ).alias(\"travel_duration_total\"),\n  ])\n\n  df = df.with_columns(exprs)\n\n  df = df.with_columns([\n      pl.max_horizontal(pl.col(\"travel_connection_duration\"),\n                        1).alias(\"travel_connection_duration\"),\n      pl.max_horizontal(pl.col(\"travel_duration_total\"),\n                        1).alias(\"travel_duration_total\"),\n  ])\n\n  df = df.with_columns([\n      (pl.col(\"travel_connection_duration\") /\n       (pl.col(\"travel_duration_total\") +\n        1e-5)).alias(\"travel_connection_duration_ratio\"),\n      (pl.col(\"flight_duration_total\") /\n       (pl.col(\"travel_duration_total\") +\n        1e-5)).alias(\"flight_duration_travel_ratio\"),\n  ])\n\n  exprs = [\n      (((pl.col(\"legs0_departure_utc\") -\n         pl.col(\"requestDate\")).dt.total_seconds()) / 3600\n      ).alias(\"book_lead_time_hours\"),\n      (((pl.col(\"legs1_arrival_utc\") -\n         pl.col(\"requestDate\")).dt.total_seconds()) /\n       3600).alias(\"book_after_time_hours\"),\n  ]\n  if 'requestDepartureDate' in df.columns:\n    exprs += [\n        (((pl.col(\"legs0_departureAt\").str.to_datetime() -\n           pl.col(\"requestDepartureDate\").str.to_datetime()).dt.total_seconds())\n         / 3600).alias(\"requestDepartureDate_diff_hours\"),\n        (((pl.col(\"legs1_departureAt\").str.to_datetime() -\n           pl.col(\"requestReturnDate\").str.to_datetime()).dt.total_seconds()) /\n         3600).alias(\"requestReturnDate_diff_hours\"),\n    ]\n  df = df.with_columns(exprs)\n\n  df = df.with_columns(\n      (pl.col('travel_duration_total') / 24).alias('travel_total_days'),\n      (pl.col('book_lead_time_hours') / 24).alias(\"book_lead_time_days\"),\n      (pl.col('book_after_time_hours') / 24).alias(\"book_after_time_days\"),\n  )\n\n  temp_cols = [\n      'legs0_departure_utc',\n      'legs0_arrival_utc',\n      'legs1_departure_utc',\n      'legs1_arrival_utc',\n      'legs0_dep_tz',\n      'legs0_arr_tz',\n      'legs1_dep_tz',\n      'legs1_arr_tz',\n      'legs0_departureAirport',\n      # 'legs0_segments0_departureFrom_airport_iata',\n      #  'legs0_arrivalAirport',\n      'legs1_departureAirport',\n      # 'legs1_segments0_departureFrom_airport_iata',\n      #  'legs1_arrivalAirport'\n  ]\n  df = df.drop([col for col in temp_cols if col in df.columns])\n\n  return df\n\n@time_feats()\ndef add_time_feats(df: pl.DataFrame) -> pl.DataFrame:\n  exprs = []\n  time_cols = [\n      \"legs0_departureAt\",\n      \"legs0_arrivalAt\",\n      \"legs1_departureAt\",\n      \"legs1_arrivalAt\",\n      \"requestDepartureDate\",\n      \"requestReturnDate\",\n  ]\n  # for leg in [0, 1]:\n  #   for seg in range(2):\n  #     time_cols.extend([\n  #       f'legs{leg}_seg{seg}_departure_local',\n  #       f'legs{leg}_seg{seg}_arrival_local'\n  #     ])\n  ic(time_cols)\n\n  for c in time_cols:\n    # assert c in df.columns\n    if c not in df.columns:\n      logger.info(f\"Column {c} is missing from DataFrame\")\n      continue\n    if not c.endswith('_local'):\n      dt = pl.col(c).str.to_datetime()\n    else:\n      dt = pl.col(c)\n    hour_col = dt.dt.hour()\n    weekday_col = dt.dt.weekday()\n    month_col = dt.dt.month()\n    exprs += [\n        hour_col.alias(f\"time_{c}_hour\"),\n        weekday_col.alias(f\"time_{c}_weekday\"),\n        month_col.alias(f\"time_{c}_month\"),\n        (dt.dt.weekday() >= 5).cast(pl.Int32).alias(f\"time_{c}_is_weekend\"),\n        (dt.dt.hour().is_between(6, 9) | dt.dt.hour().is_between(17, 20)\n        ).cast(pl.Int32).alias(f\"time_{c}_is_peak\"),\n        (dt.dt.hour().is_between(0, 5)\n        ).cast(pl.Int32).alias(f\"time_{c}_is_red_eye\"),\n        pl.when(hour_col.is_between(5, 8)).then(0)  \n        .when(hour_col.is_between(9, 11)).then(1)  \n        .when(hour_col.is_between(12, 17)).then(2)  \n        .when(hour_col.is_between(18, 22)).then(3)  \n        .otherwise(4)  \n        .alias(f\"time_{c}_period\"),\n\n        pl.when(month_col.is_in([12, 1, 2])).then(0)\n        .when(month_col.is_in([3, 4, 5])).then(1)  \n        .when(month_col.is_in([6, 7, 8])).then(2)  \n        .otherwise(3) \n        .alias(f\"time_{c}_season\"),\n\n        (weekday_col.is_between(1, 5) & hour_col.is_between(8, 18)\n        ).cast(pl.Int8).alias(f\"time_{c}_is_business_hours\"),\n\n        (hour_col * (2 * np.pi / 24)).sin().alias(f\"time_{c}_hour_sin\"),\n        (hour_col * (2 * np.pi / 24)).cos().alias(f\"time_{c}_hour_cos\"),\n        (weekday_col * (2 * np.pi / 7)).sin().alias(f\"time_{c}_weekday_sin\"),\n        (weekday_col * (2 * np.pi / 7)).cos().alias(f\"time_{c}_weekday_cos\"),\n        (month_col * (2 * np.pi / 12)).sin().alias(f\"time_{c}_month_sin\"),\n        (month_col * (2 * np.pi / 12)).cos().alias(f\"time_{c}_month_cos\"),\n    ]\n  df = df.with_columns(exprs)\n  time_cols = [col for col in time_cols if col in df.columns]\n  df = df.drop(time_cols)\n\n  return df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T02:47:20.561607Z","iopub.execute_input":"2025-08-25T02:47:20.561856Z","iopub.status.idle":"2025-08-25T02:47:20.636009Z","shell.execute_reply.started":"2025-08-25T02:47:20.561834Z","shell.execute_reply":"2025-08-25T02:47:20.631484Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@time_feats()\ndef add_cabin_feats(df: pl.DataFrame) -> pl.DataFrame:\n  exprs = []\n  for leg in (0, 1):\n    cabin_cols = [\n        f\"legs{leg}_segments{s}_cabinClass\" for s in range(4)\n        if f\"legs{leg}_segments{s}_cabinClass\" in df.columns\n    ]\n\n    assert cabin_cols, f\"leg{leg} cabin_cols is empty\"\n    exprs.append(\n        pl.mean_horizontal(\n            pl.col(c) for c in cabin_cols).alias(f\"avg_cabin_legs{leg}\"))\n\n  df = df.with_columns(exprs)\n\n  df = df.with_columns([\n      pl.mean_horizontal(\n          pl.col(c) for c in [\"avg_cabin_legs0\", \"avg_cabin_legs1\"]).alias(\n              \"avg_cabin_legs_all\"),\n  ])\n\n  return df\n\n@time_feats()\ndef add_baggage_feats(df: pl.DataFrame) -> pl.DataFrame:\n  drop_cols = []\n  exprs = []\n  for leg in (0, 1):\n    for s in range(4):\n      exprs.extend([\n          # baggage_count: only keep quantity if type is 'piece'\n          pl.when(\n              pl.col(\n                  f\"legs{leg}_segments{s}_baggageAllowance_weightMeasurementType\"\n              ) == 0\n          ).then(\n              pl.col(f\"legs{leg}_segments{s}_baggageAllowance_quantity\").cast(\n                  pl.Int8)\n          ).otherwise(None).alias(f\"legs{leg}_segments{s}_baggage_count\"),\n          # baggage_weight: only keep quantity if type is 'weight'\n          pl.when(\n              pl.col(\n                  f\"legs{leg}_segments{s}_baggageAllowance_weightMeasurementType\"\n              ) == 1\n          ).then(\n              pl.col(f\"legs{leg}_segments{s}_baggageAllowance_quantity\").cast(\n                  pl.Float32)\n          ).otherwise(None).alias(f\"legs{leg}_segments{s}_baggage_weight\")\n      ])\n      drop_cols.append(f\"legs{leg}_segments{s}_baggageAllowance_quantity\")\n  df = df.with_columns(exprs)\n  df = df.drop(drop_cols)\n\n  exprs = []\n  for leg in (0, 1):\n    baggage_cols = [\n        f\"legs{leg}_segments{s}_baggage_count\" for s in range(4)\n        if f\"legs{leg}_segments{s}_baggage_count\" in df.columns\n    ]\n\n    assert baggage_cols, f\"leg{leg} baggage_cols is empty\"\n    exprs.append(\n        pl.mean_horizontal([pl.col(c) for c in baggage_cols\n                           ]).alias(f\"avg_baggage_count_legs{leg}\"))\n\n  for leg in (0, 1):\n    baggage_cols = [\n        f\"legs{leg}_segments{s}_baggage_weight\" for s in range(4)\n        if f\"legs{leg}_segments{s}_baggage_weight\" in df.columns\n    ]\n\n    assert baggage_cols, f\"leg{leg} baggage_cols is empty\"\n    exprs.append(\n        pl.mean_horizontal([pl.col(c) for c in baggage_cols\n                           ]).alias(f\"avg_baggage_weight_legs{leg}\"))\n\n  df = df.with_columns(exprs)\n\n  df = df.with_columns([\n      pl.mean_horizontal(\n          pl.col(c)\n          for c in [\"avg_baggage_count_legs0\", \"avg_baggage_count_legs1\"\n                   ]).alias(\"avg_baggage_count_legs_all\"),\n      pl.mean_horizontal(\n          pl.col(c)\n          for c in [\"avg_baggage_weight_legs0\", \"avg_baggage_weight_legs1\"\n                   ]).alias(\"avg_baggage_weight_legs_all\"),\n  ])\n\n  return df\n\n@time_feats()\ndef add_seats_feats(df: pl.DataFrame) -> pl.DataFrame:\n  exprs = []\n  for leg in (0, 1):\n    seats_cols = [\n        f\"legs{leg}_segments{s}_seatsAvailable\" for s in range(4)\n        if f\"legs{leg}_segments{s}_seatsAvailable\" in df.columns\n    ]\n\n    assert seats_cols, f\"leg{leg} seats_cols is empty\"\n    exprs.append(\n        pl.mean_horizontal(\n            pl.col(c) for c in seats_cols).alias(f\"avg_seats_count_legs{leg}\"))\n\n  df = df.with_columns(exprs)\n\n  df = df.with_columns([\n      pl.mean_horizontal(\n          pl.col(c) for c in [\"avg_seats_count_legs0\", \"avg_seats_count_legs1\"]\n      ).alias(\"avg_seats_count_legs_all\"),\n  ])\n  return df\n\n@time_feats()\ndef add_carrier_feats(df: pl.DataFrame) -> pl.DataFrame:\n  mc_cols = []\n  for leg in (0, 1):\n    mc_cols.extend([\n        f\"legs{leg}_segments{s}_marketingCarrier_code\" for s in range(4)\n        if f\"legs{leg}_segments{s}_marketingCarrier_code\" in df.columns\n    ])\n\n    assert mc_cols, f\"leg{leg} mc_cols is empty\"\n\n  df = df.with_columns(\n    pl.struct(mc_cols)\n      .map_elements(lambda s: len(set(v for v in s.values() if v is not None)), return_dtype=pl.UInt8)\n      .alias(\"num_unique_carriers\")\n  )\n\n  df = df.with_columns([\n      (pl.col('num_unique_carriers') / pl.max_horizontal(\n          pl.col('seg_legs_all_count'), 1)).alias('carrier_diversity_ratio'),\n  ])\n\n  return df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T02:47:20.638026Z","iopub.execute_input":"2025-08-25T02:47:20.638277Z","iopub.status.idle":"2025-08-25T02:47:20.713357Z","shell.execute_reply.started":"2025-08-25T02:47:20.638255Z","shell.execute_reply":"2025-08-25T02:47:20.708909Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@time_feats()\ndef add_ranking_feats(df, group_col, suffix=''):\n  exprs = []\n  for col, order in rank_order.items():\n    if col in df.columns:\n      exprs.append(\n          pl.col(col).rank(method='average', descending=(\n              order == 'desc')).over(group_col).alias(f'rank_{col}{suffix}'))\n    else:\n      logger.warning(f\"Column {col} not found in DataFrame, skipping ranking.\")\n  df = df.with_columns(exprs)\n\n  return df\n\n@time_feats()\ndef add_flighthash_feats(df):\n  df = (df.with_columns([\n      pl.len().over([\"ranker_id\", \"flight_hash\"]).alias(\"flight_hash_count\"),\n  ]).with_columns([\n      (pl.col(\"flight_hash_count\") /\n       pl.col(\"group_size\")).alias(\"flight_hash_ratio\"),\n      pl.col(\"flight_hash_count\").rank(\n          \"dense\",\n          descending=True).over(\"ranker_id\").alias(\"rank_flight_hash_count\"),\n  ]))\n  return df\n\n#Notice not consider label/selected and df is (train and test) combined\n@time_feats()\ndef add_stats_feats(df, group_col='profileId'):\n  exprs = []\n  for col in rank_order.keys():\n    if col in df.columns:\n      exprs.extend([\n        pl.col(col).mean().over(group_col).alias(f\"avg_{col}_{group_col}_stats\"),\n        pl.col(col).min().over(group_col).alias(f\"min_{col}_{group_col}_stats\"),\n        pl.col(col).max().over(group_col).alias(f\"max_{col}_{group_col}_stats\"),\n        pl.col(col).std().over(group_col).alias(f\"std_{col}_{group_col}_stats\"),\n        pl.col(col).median().over(group_col).alias(f\"median_{col}_{group_col}_stats\"),\n      ])\n  df = df.with_columns(exprs)\n  exprs = []\n  for col in rank_order.keys():\n    if col in df.columns:\n      exprs.extend([\n        ((pl.col(col) - pl.col(f\"avg_{col}_{group_col}_stats\")) / (pl.col(f\"std_{col}_{group_col}_stats\") + 1e-5)).alias(f\"{col}_zscore_{group_col}_stats\"),\n        ((pl.col(col) - pl.col(f\"min_{col}_{group_col}_stats\")) / (pl.col(f\"max_{col}_{group_col}_stats\") - pl.col(f\"min_{col}_{group_col}_stats\") + 1e-5)).alias(f\"{col}_minmax_{group_col}_stats\"),\n        (pl.col(col) / (pl.col(f\"avg_{col}_{group_col}_stats\") + 1e-5)).alias(f\"{col}_{group_col}_stats_ratio\"),\n      ])\n\n  df = df.with_columns(exprs)\n  return df\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T02:47:20.715625Z","iopub.execute_input":"2025-08-25T02:47:20.715914Z","iopub.status.idle":"2025-08-25T02:47:20.792370Z","shell.execute_reply.started":"2025-08-25T02:47:20.715891Z","shell.execute_reply":"2025-08-25T02:47:20.787515Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Make history avg feat \nhttps://www.kaggle.com/code/mikhailgolubchik/sm-xgboost-single\nNotice this is added after contest ends, it could boost online LB/PB +0.008","metadata":{}},{"cell_type":"code","source":"@time_feats()\ndef make_history_avg(df, source_cols, group_col, suffix):\n  ori_cols = [col for col in df.columns]\n  selected_df = df.filter(\n      pl.col(\"selected\") == 1).select([\"ranker_id\", \"requestDate\", group_col] +\n                                      source_cols)\n\n  ranker_to_profile = dict(\n      zip(selected_df[\"ranker_id\"].to_list(), selected_df[group_col].to_list()))\n\n  ranker_to_timestamp = dict(\n      zip(selected_df[\"ranker_id\"].to_list(),\n          selected_df[\"requestDate\"].to_list()))\n\n  history_df = selected_df.select([\"ranker_id\", \"requestDate\", group_col] +\n                                  source_cols)\n\n  all_stats_dict = {}  # ranker_id -> {col_mean: val, col_std: val}\n\n  unique_ranker_ids = df[\"ranker_id\"].unique().to_list()\n\n  for current_ranker_id in tqdm(unique_ranker_ids,\n                                desc=\"Обработка ranker_id\",\n                                mininterval=10.0):\n    current_profile_id = ranker_to_profile.get(current_ranker_id)\n\n    current_timestamp = ranker_to_timestamp.get(current_ranker_id)\n\n    profile_history = history_df.filter(\n        (pl.col(group_col) == current_profile_id) &\n        (pl.col(\"ranker_id\") != current_ranker_id) &\n        (pl.col(\"requestDate\") < current_timestamp))\n\n    agg_result = profile_history.select([\n        *[\n            pl.col(col).mean().alias(f\"{col}{suffix}_mean\")\n            for col in source_cols\n        ],\n        *[pl.col(col).std().alias(f\"{col}{suffix}_std\") for col in source_cols],\n        *[\n            pl.col(col).count().alias(f\"{col}{suffix}_count\")\n            for col in source_cols\n        ],  \n        *[\n            pl.col(col).median().alias(f\"{col}{suffix}_median\")\n            for col in source_cols\n        ],\n        *[\n            pl.col(col).quantile(0.25).alias(f\"{col}{suffix}_q25\")\n            for col in source_cols\n        ],\n        *[\n            pl.col(col).quantile(0.75).alias(f\"{col}{suffix}_q75\")\n            for col in source_cols\n        ]\n    ])\n\n    if agg_result.height > 0:\n      row = agg_result.row(0)\n      n_cols = len(source_cols)\n      all_stats_dict[current_ranker_id] = {\n          **{\n              f\"{col}{suffix}_mean\": row[i] for i, col in enumerate(source_cols)\n          },\n          **{\n              f\"{col}{suffix}_std\": row[i + n_cols] for i, col in enumerate(source_cols)\n          },\n          **{\n              f\"{col}{suffix}_count\": row[i + 2 * n_cols] for i, col in enumerate(source_cols)\n          },\n          **{\n              f\"{col}{suffix}_median\": row[i + 3 * n_cols] for i, col in enumerate(source_cols)\n          },\n          **{\n              f\"{col}{suffix}_q25\": row[i + 4 * n_cols] for i, col in enumerate(source_cols)\n          },\n          **{\n              f\"{col}{suffix}_q75\": row[i + 5 * n_cols] for i, col in enumerate(source_cols)\n          },\n      }\n\n  update_data = []\n  for ranker_id, stats in all_stats_dict.items():\n    row = {\"ranker_id\": ranker_id, **stats}\n    update_data.append(row)\n\n  schema = {\"ranker_id\": pl.Utf8} \n  for col in source_cols:\n    schema[f\"{col}{suffix}_mean\"] = pl.Float32\n    schema[f\"{col}{suffix}_std\"] = pl.Float32\n    schema[f\"{col}{suffix}_count\"] = pl.Int32\n    schema[f\"{col}{suffix}_median\"] = pl.Float32\n    schema[f\"{col}{suffix}_q25\"] = pl.Float32\n    schema[f\"{col}{suffix}_q75\"] = pl.Float32\n\n  update_df = pl.DataFrame(update_data, schema=schema)\n\n  df = df.join(update_df, on=\"ranker_id\", how=\"left\")\n\n  agg_exprs = []\n  for col in source_cols:\n    agg_exprs.extend([\n      pl.col(col).mean().alias(f\"{col}{suffix}_mean\"),\n      pl.col(col).std().alias(f\"{col}{suffix}_std\"),\n      pl.col(col).count().alias(f\"{col}{suffix}_count\"),\n      pl.col(col).median().alias(f\"{col}{suffix}_median\"),\n      pl.col(col).quantile(0.25).alias(f\"{col}{suffix}_q25\"),\n      pl.col(col).quantile(0.75).alias(f\"{col}{suffix}_q75\")\n    ])\n\n  df_stats = (df.filter(\n      pl.col(\"selected\") == 1).group_by(group_col).agg(agg_exprs))\n\n  stats_cols = [c for c in df.columns if c.endswith((\n    \"_mean\", \"_std\", \"_median\", \"_q25\", \"_q75\")) and c not in ori_cols]\n  count_cols = [c for c in df.columns if c.endswith(\"_count\") and c not in ori_cols]\n\n  df = df.with_columns([\n    pl.col(stats_cols).cast(pl.Float32),\n    pl.col(count_cols).cast(pl.Int32)\n  ])\n  df_stats = df_stats.with_columns([\n    pl.col(stats_cols).cast(pl.Float32),\n    pl.col(count_cols).cast(pl.Int32)\n  ])\n\n  return df, df_stats","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T02:47:20.795079Z","iopub.execute_input":"2025-08-25T02:47:20.795341Z","iopub.status.idle":"2025-08-25T02:47:20.821612Z","shell.execute_reply.started":"2025-08-25T02:47:20.795317Z","shell.execute_reply":"2025-08-25T02:47:20.816849Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@time_feats()\ndef gen_feats(df):\n  df = add_group_feats(df)\n  df = add_user_feats(df)\n  df = add_searchRoute_feats(df)\n  df = add_segment_feats(df)\n  df = add_flight_duration_feats(df)\n  df = add_segment_geography_time_feats(df)\n  df = add_travel_duration_feats(df)\n  df = add_time_feats(df)\n  \n  temp_geo_cols = []\n  for leg in [0, 1]:\n    temp_geo_cols.extend([\n        f'legs{leg}_dep_lat', f'legs{leg}_dep_lon', f'legs{leg}_arr_lat',\n        f'legs{leg}_arr_lon'\n    ])\n    for seg in range(4):\n      temp_geo_cols.extend([\n          f\"legs{leg}_seg{seg}_dep_lat\",\n          f\"legs{leg}_seg{seg}_dep_lon\",\n          f\"legs{leg}_seg{seg}_arr_lat\",\n          f\"legs{leg}_seg{seg}_arr_lon\",\n\n      ])\n\n  drop_cols = [col for col in temp_geo_cols if col in df.columns]\n  ic(drop_cols)\n  df = df.drop(drop_cols)\n  \n  df = add_cabin_feats(df)\n  df = add_baggage_feats(df)\n  df = add_seats_feats(df)\n  df = add_carrier_feats(df)\n  \n  df = add_flighthash_feats(df)\n  df = add_ranking_feats(df, 'ranker_id')\n  df = add_ranking_feats(df, ['ranker_id', 'flight_hash'], '_in_hash_group')\n  \n  df = add_stats_feats(df, 'uid')\n  df = add_stats_feats(df, 'companyID')\n  \n  drop_cols = [\n      col for col in df.columns if any(['_utc' in col, '_local' in col])\n  ]\n  ic(drop_cols)\n  df = df.drop(drop_cols)\n  \n  if FLAGS.history_avg:\n    test = get_test(df)\n    \n    df = get_nontest(df)\n    train = get_train(df)\n    \n    # if not training using all train data, valid data need to merge stats similar as test\n    if not FLAGS.online:\n      valid = get_valid(df)\n      test = pl.concat([valid, test], how='vertical')\n\n    train, df_stats_pr = make_history_avg(train,\n                                       source_cols=source_cols,\n                                       group_col=\"uid\",\n                                       suffix='_uid')\n    test = test.join(df_stats_pr, on='uid', how='left')\n    \n    train, df_stats_co = make_history_avg(train,\n                                       source_cols=source_cols,\n                                       group_col=\"companyID\",\n                                       suffix='_company')\n    test = test.join(df_stats_co, on='companyID', how='left')\n\n    # test = test.select(train.columns)\n    \n    df = align_and_concat([train, test])\n  return df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T02:47:20.824060Z","iopub.execute_input":"2025-08-25T02:47:20.824332Z","iopub.status.idle":"2025-08-25T02:47:20.870921Z","shell.execute_reply.started":"2025-08-25T02:47:20.824307Z","shell.execute_reply":"2025-08-25T02:47:20.866409Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def preprocess(add_feats=True):\n  df = load_df(use_ext=FLAGS.use_ext)\n  df = set_fold(df)\n  if FLAGS.fast:\n    df_train = get_nontest(df)\n    df_test = get_test(df)\n    df_train = filter(df_train, 0.01, FLAGS.seed)\n    df = pl.concat([df_train, df_test], how='vertical')\n       \n  if add_feats:\n    df = gen_feats(df)\n  \n  numer_cols = get_numer_cols(df)\n  cat_cols = get_cat_cols(df)\n  feat_cols = numer_cols + cat_cols\n  \n  df = smart_fillnull(df, numer_cols, cat_cols)\n  \n  ignore_cols = [col for col in IGNORE_COLS]\n  cat_cols = [col for col in cat_cols if col not in ignore_cols]\n  numer_cols = [col for col in numer_cols if col not in ignore_cols]\n  feat_cols = [col for col in feat_cols if col not in ignore_cols]\n  \n  train = get_train(df) if not FLAGS.stats_all else df\n  cats = encode_cat_unified(train, cat_cols, method=FLAGS.cat_method, num_workers=1)\n  ic(cats.keys())\n  df = df.with_columns([\n      pl.col(col).replace_strict(cats[get_unified_cat(col)], default=-1) for col in cat_cols\n  ])\n  \n  if FLAGS.remove_cats:  \n    if not FLAGS.reserve_cats:\n      numer_cols += cat_cols\n      cat_cols = []\n    else:\n      reserve_cats = ['profileId', 'companyID', 'corporateTariffCode', 'nationality', 'companyCode']\n      numer_cols += [col for col in cat_cols if col not in reserve_cats]\n      cat_cols = [col for col in cat_cols if col in reserve_cats] \n   \n  icl(numer_cols, 10) \n  icl(cat_cols, 10) \n  cols_dict = {\n        'numer': numer_cols,  \n        'cat': cat_cols,    \n        'feat': feat_cols,           \n  }  \n  return df, cols_dict","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T02:47:20.873167Z","iopub.execute_input":"2025-08-25T02:47:20.873431Z","iopub.status.idle":"2025-08-25T02:47:21.101972Z","shell.execute_reply.started":"2025-08-25T02:47:20.873408Z","shell.execute_reply":"2025-08-25T02:47:21.096707Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df, cols_dict = preprocess(add_feats=FLAGS.add_feats)\ndf","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T02:47:21.103937Z","iopub.execute_input":"2025-08-25T02:47:21.105247Z","iopub.status.idle":"2025-08-25T02:59:13.176563Z","shell.execute_reply.started":"2025-08-25T02:47:21.105207Z","shell.execute_reply":"2025-08-25T02:59:13.171156Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def preprocess_cat_cols(df, cat_cols):\n  if not cat_cols:\n    return df\n  \n  for col in cat_cols:\n    df[col] = df[col].astype('category')\n  return df\n\ndef get_num_boost_round(params):\n  if FLAGS.fast:\n    return 100\n  if FLAGS.trees:\n    return FLAGS.trees\n  if 'iterations' in params:\n    return params['iterations']\n  if 'num_iterations' in params:\n    return params['num_iterations']\n  if 'n_estimators' in params:\n    return params['n_estimators']","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T02:59:13.178848Z","iopub.execute_input":"2025-08-25T02:59:13.179064Z","iopub.status.idle":"2025-08-25T02:59:13.188575Z","shell.execute_reply.started":"2025-08-25T02:59:13.179044Z","shell.execute_reply":"2025-08-25T02:59:13.184513Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cat_cols = cols_dict['cat']\nfeat_cols = cols_dict['feat']\ndf_train, df_valid = get_train_valid(df)  \ndf_test = get_test(df)\n\nif FLAGS.mode == 'train':\n  X_train = df_train.to_pandas()\n  X_train = X_train[feat_cols]\n  X_train = preprocess_cat_cols(X_train, cat_cols)\n  \n  y_train = df_train['selected'].to_pandas()\n  group = df_train.select('ranker_id').group_by('ranker_id', maintain_order=True).agg(pl.len())['len'].to_numpy()\n  ic(X_train.shape, group.shape)\n\n  X_valid = df_valid.to_pandas()\n  X_valid = X_valid[feat_cols]\n  X_valid = preprocess_cat_cols(X_valid, cat_cols)\n\nX_test = df_test.to_pandas()\nX_test = X_test[feat_cols]\nX_test = preprocess_cat_cols(X_test, cat_cols)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T02:59:13.191017Z","iopub.execute_input":"2025-08-25T02:59:13.191306Z","iopub.status.idle":"2025-08-25T02:59:28.164342Z","shell.execute_reply.started":"2025-08-25T02:59:13.191283Z","shell.execute_reply":"2025-08-25T02:59:28.160155Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import xgboost as xgb\nif FLAGS.mode == 'train':\n  dtrain = xgb.DMatrix(\n    X_train,\n    y_train,\n    group=group,\n    enable_categorical=True)\n# dvalid = xgb.DMatrix(X_valid, enable_categorical=True)\n# dtest = xgb.DMatrix(X_test, enable_categorical=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T03:01:35.105066Z","iopub.execute_input":"2025-08-25T03:01:35.105452Z","iopub.status.idle":"2025-08-25T03:01:35.116091Z","shell.execute_reply.started":"2025-08-25T03:01:35.105421Z","shell.execute_reply":"2025-08-25T03:01:35.111274Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def show_feat_importance(model):\n  imp = model.get_score(importance_type=\"gain\")\n\n  imp_df = (\n    pd.DataFrame({\n      \"feat\": list(imp.keys()), \n      \"importance\": list(imp.values())\n      })\n    .sort_values(\"importance\", ascending=False)\n    .reset_index(drop=True)\n  )\n  return imp_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T03:01:36.038156Z","iopub.execute_input":"2025-08-25T03:01:36.038452Z","iopub.status.idle":"2025-08-25T03:01:36.049586Z","shell.execute_reply.started":"2025-08-25T03:01:36.038428Z","shell.execute_reply":"2025-08-25T03:01:36.043999Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class XGBTQDMCallback(xgb.callback.TrainingCallback):\n\n  def __init__(self, total_iterations, desc='', eval_name='valid'):\n    self.pbar = tqdm(total=total_iterations, desc=desc)\n    self.eval_name = eval_name\n\n  def after_iteration(self, model, epoch, evals_log):\n    self.pbar.update(1)\n    for eval_name, metrics in evals_log.items():\n      if eval_name == self.eval_name:\n        m = {}\n        for metric_name, values in metrics.items():\n          m.update({metric_name: values[-1]})\n        self.pbar.set_postfix(m)\n\n    if epoch + 1 == self.pbar.total:\n      self.pbar.close()\n\n    return False","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T03:01:36.937080Z","iopub.execute_input":"2025-08-25T03:01:36.937426Z","iopub.status.idle":"2025-08-25T03:01:36.949327Z","shell.execute_reply.started":"2025-08-25T03:01:36.937398Z","shell.execute_reply":"2025-08-25T03:01:36.945100Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def batch_predict(model, X, batch_size=50_000, proba=False, inplace=False, progress=True):\n  results = []\n  n = len(X)\n  iterator = range(0, n, batch_size)\n  if progress:\n    iterator = tqdm(iterator, desc=\"Batch Prediction\")\n\n  for start in iterator:\n    end = min(start + batch_size, n)\n    batch = X[start:end]\n\n    if inplace and hasattr(model, \"inplace_predict\"):\n      batch_pred = model.inplace_predict(batch)\n    elif proba and hasattr(model, \"predict_proba\"):\n      batch_pred = model.predict_proba(batch)\n    else:\n      batch_pred = model.predict(batch)\n\n    results.append(batch_pred)\n\n  first = results[0]\n  if isinstance(first, np.ndarray):\n    if first.ndim == 1:\n      return np.concatenate(results)\n    else:\n      return np.vstack(results)\n  else:\n    return [item for sublist in results for item in sublist]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T03:01:40.905541Z","iopub.execute_input":"2025-08-25T03:01:40.905876Z","iopub.status.idle":"2025-08-25T03:01:40.919577Z","shell.execute_reply.started":"2025-08-25T03:01:40.905850Z","shell.execute_reply":"2025-08-25T03:01:40.914655Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"models = []\npreds = []","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T03:01:41.310641Z","iopub.execute_input":"2025-08-25T03:01:41.310935Z","iopub.status.idle":"2025-08-25T03:01:41.322010Z","shell.execute_reply.started":"2025-08-25T03:01:41.310911Z","shell.execute_reply":"2025-08-25T03:01:41.316071Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"params = params_xgb.copy()\nobjectives = [\n  'rank:ndcg',\n  'rank:map',\n  'rank:pairwise',\n  'binary:logistic'\n]\nif FLAGS.n_models > 0:\n  objectives = objectives[:FLAGS.n_models]\nparams['device'] = FLAGS.device\n# params['predictor'] = 'cpu_predictor'\nic(params)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T03:01:47.111077Z","iopub.execute_input":"2025-08-25T03:01:47.111382Z","iopub.status.idle":"2025-08-25T03:01:47.141324Z","shell.execute_reply.started":"2025-08-25T03:01:47.111356Z","shell.execute_reply":"2025-08-25T03:01:47.136215Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# XGB training","metadata":{}},{"cell_type":"code","source":"if FLAGS.mode == 'train':\n  for objective in tqdm(objectives, desc='objectives'):\n    params['objective'] = objective\n    model = xgb.train(\n              params=params,\n              dtrain=dtrain,\n              num_boost_round=get_num_boost_round(params),\n              evals=None,\n              callbacks=[XGBTQDMCallback(get_num_boost_round(params), f'xgb_train_{objective}')],\n              verbose_eval=100,\n          )\n    model.set_param({\"device\": \"cpu\"}) \n    # pred = model.predict(dvalid)\n    # pred = model.inplace_predict(X_valid)\n    pred = batch_predict(model, X_valid, inplace=True)\n    ic(objective, pred)\n    preds.append(pred)\n    ic(objective)\n    display(show_feat_importance(model))\n    models.append(model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T03:01:50.932016Z","iopub.execute_input":"2025-08-25T03:01:50.932313Z","iopub.status.idle":"2025-08-25T03:01:50.946043Z","shell.execute_reply.started":"2025-08-25T03:01:50.932288Z","shell.execute_reply":"2025-08-25T03:01:50.940226Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def hitrate_at_3(y_true, y_pred, groups):\n  df = pl.DataFrame({'group': groups, 'pred': y_pred, 'true': y_true})\n\n  return (df.filter(pl.col(\"group\").count().over(\"group\") > 10).sort(\n      [\"group\", \"pred\"], descending=[False, True]).group_by(\n          \"group\", maintain_order=True).head(3).group_by(\"group\").agg(\n              pl.col(\"true\").max()).select(pl.col(\"true\").mean()).item())\n  \ndef eval_df(df):\n  score = hitrate_at_3(\n      df['selected'].to_numpy(),\n      df['pred'].to_numpy(),\n      df['ranker_id'].to_numpy()\n  )\n  return score","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T02:59:30.199735Z","iopub.status.idle":"2025-08-25T02:59:30.200398Z","shell.execute_reply.started":"2025-08-25T02:59:30.199872Z","shell.execute_reply":"2025-08-25T02:59:30.199885Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def rerank(df: pl.DataFrame, penalty_factor=0.12):\n  df = df.with_columns(\n      pl.max(\"pred\").over([\"ranker_id\", \"flight_hash\"]).alias(\"max_score_same_flight\"))\n\n  df = df.with_columns((pl.col(\"pred\") - penalty_factor * (pl.col(\"max_score_same_flight\") - pl.col(\"pred\"))).alias(\"pred\"))\n\n  return df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T03:01:53.665221Z","iopub.execute_input":"2025-08-25T03:01:53.665512Z","iopub.status.idle":"2025-08-25T03:01:53.678856Z","shell.execute_reply.started":"2025-08-25T03:01:53.665488Z","shell.execute_reply":"2025-08-25T03:01:53.673039Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# XGB eval","metadata":{}},{"cell_type":"code","source":"for objective, pred in zip(objectives, preds):\n  df_valid = df_valid.with_columns(\n    pl.Series(\"pred\", pred)\n  )\n  score = eval_df(df_valid)\n  ic(objective, score)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T03:01:54.982986Z","iopub.execute_input":"2025-08-25T03:01:54.983265Z","iopub.status.idle":"2025-08-25T03:01:54.993815Z","shell.execute_reply.started":"2025-08-25T03:01:54.983227Z","shell.execute_reply":"2025-08-25T03:01:54.989522Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def ensemble(preds):\n  preds = np.array(preds)\n  i = 0\n  for pred in preds:\n    ic(pred.shape)\n    min_val = pred.min()\n    max_val = pred.max()\n    if not (min_val >= 0 and max_val <= 1):\n      pred = 1 / (1 + np.exp(-pred))\n    preds[i] = pred\n    i += 1\n  pred = preds.mean(0)\n  return pred","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T03:01:57.247404Z","iopub.execute_input":"2025-08-25T03:01:57.247714Z","iopub.status.idle":"2025-08-25T03:01:57.259803Z","shell.execute_reply.started":"2025-08-25T03:01:57.247688Z","shell.execute_reply":"2025-08-25T03:01:57.254898Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if FLAGS.mode == 'train':\n  for i, model in tqdm(enumerate(models), total=len(models)):\n    with open(f'{FLAGS.out_dir}/{i}.pkl', 'wb') as f:\n      pickle.dump(model, f)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T03:02:00.679791Z","iopub.execute_input":"2025-08-25T03:02:00.680075Z","iopub.status.idle":"2025-08-25T03:02:00.690310Z","shell.execute_reply.started":"2025-08-25T03:02:00.680049Z","shell.execute_reply":"2025-08-25T03:02:00.685922Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# XGB load pretrain model","metadata":{}},{"cell_type":"code","source":"if FLAGS.mode != 'train':\n  models = []\n  preds = []\n  model_dir = FLAGS.model_dir if os.path.exists(FLAGS.model_dir) else FLAGS.out_dir\n  for i, objective in tqdm(enumerate(objectives), total=len(objectives)):\n    with open(f'{model_dir}/{i}.pkl', 'rb') as f:\n      model = pickle.load(f)\n      models.append(model)\n      #preds.append(model.inplace_predict(X_valid))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T03:02:03.071278Z","iopub.execute_input":"2025-08-25T03:02:03.071555Z","iopub.status.idle":"2025-08-25T03:02:04.100073Z","shell.execute_reply.started":"2025-08-25T03:02:03.071532Z","shell.execute_reply":"2025-08-25T03:02:04.095352Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if preds:\n  pred = ensemble(preds)\n  ic(pred, pred.shape)\n  df_valid = df_valid.with_columns(\n    pl.Series(\"pred\", pred)\n  )\n  score = eval_df(df_valid)\n  ic('ori', score)\n  df_valid = df_valid.with_columns(\n    pl.Series(\"pred\", pred)\n  )\n  df_valid = rerank(df_valid)\n  score = eval_df(df_valid)\n  ic('rerank', score)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T03:02:08.544538Z","iopub.execute_input":"2025-08-25T03:02:08.544851Z","iopub.status.idle":"2025-08-25T03:02:08.556170Z","shell.execute_reply.started":"2025-08-25T03:02:08.544825Z","shell.execute_reply":"2025-08-25T03:02:08.551208Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dump result of 4 xgb models ensemble to submission.parquet","metadata":{}},{"cell_type":"code","source":"preds = []\nfor model in tqdm(models):\n  # pred = model.predict(dtest)\n  # pred = model.inplace_predict(X_test)\n  pred = batch_predict(model, X_test, inplace=True)\n  preds.append(pred)\n  \npred = ensemble(preds)\npred","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T03:02:10.832905Z","iopub.execute_input":"2025-08-25T03:02:10.833191Z","iopub.status.idle":"2025-08-25T03:02:35.050347Z","shell.execute_reply.started":"2025-08-25T03:02:10.833166Z","shell.execute_reply":"2025-08-25T03:02:35.044436Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def probs2rank(df):\n  df = df.with_columns(\n      pl.col('pred').rank(method='ordinal',\n                          descending=True).over('ranker_id').cast(\n                              pl.Int32).alias('selected')).select(\n                                  ['Id', 'ranker_id', 'selected'])\n  return df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T03:02:38.040981Z","iopub.execute_input":"2025-08-25T03:02:38.041378Z","iopub.status.idle":"2025-08-25T03:02:38.054471Z","shell.execute_reply.started":"2025-08-25T03:02:38.041347Z","shell.execute_reply":"2025-08-25T03:02:38.047516Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_test = df_test.with_columns(\n    pl.Series(\"pred\", pred)\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T03:02:38.871780Z","iopub.execute_input":"2025-08-25T03:02:38.872163Z","iopub.status.idle":"2025-08-25T03:02:38.884649Z","shell.execute_reply.started":"2025-08-25T03:02:38.872132Z","shell.execute_reply":"2025-08-25T03:02:38.879600Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sub = probs2rank(df_test)\nsub","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T03:02:42.533714Z","iopub.execute_input":"2025-08-25T03:02:42.534056Z","iopub.status.idle":"2025-08-25T03:02:42.744580Z","shell.execute_reply.started":"2025-08-25T03:02:42.534028Z","shell.execute_reply":"2025-08-25T03:02:42.739895Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"FLAGS.out_dir","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T03:02:46.456802Z","iopub.execute_input":"2025-08-25T03:02:46.457147Z","iopub.status.idle":"2025-08-25T03:02:46.471201Z","shell.execute_reply.started":"2025-08-25T03:02:46.457119Z","shell.execute_reply":"2025-08-25T03:02:46.466933Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dump result of best single model to single.parquet","metadata":{}},{"cell_type":"code","source":"df_test = df_test.with_columns(\n    pl.Series(\"pred\", preds[0])\n)\nsub_single = probs2rank(df_test)\nsub_single","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T03:02:48.126920Z","iopub.execute_input":"2025-08-25T03:02:48.127341Z","iopub.status.idle":"2025-08-25T03:02:48.340599Z","shell.execute_reply.started":"2025-08-25T03:02:48.127306Z","shell.execute_reply":"2025-08-25T03:02:48.336338Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sub_single.write_parquet(f'{FLAGS.out_dir}/best_single.parquet')\nsub.write_csv(f'submission.csv')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-25T03:02:50.606465Z","iopub.execute_input":"2025-08-25T03:02:50.606803Z","iopub.status.idle":"2025-08-25T03:02:52.079763Z","shell.execute_reply.started":"2025-08-25T03:02:50.606777Z","shell.execute_reply":"2025-08-25T03:02:52.075443Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}