{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":38760,"databundleVersionId":4493939,"sourceType":"competition"},{"sourceId":4474043,"sourceType":"datasetVersion","datasetId":2601572},{"sourceId":4483558,"sourceType":"datasetVersion","datasetId":2623568}],"dockerImageVersionId":30588,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install google-cloud-bigquery-storage omegaconf","metadata":{"execution":{"iopub.status.busy":"2023-11-26T04:59:39.615081Z","iopub.execute_input":"2023-11-26T04:59:39.615476Z","iopub.status.idle":"2023-11-26T04:59:51.450073Z","shell.execute_reply.started":"2023-11-26T04:59:39.615444Z","shell.execute_reply":"2023-11-26T04:59:51.449127Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\nimport io\nimport os\nimport sys\nimport logging\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport polars as pl\nfrom glob import glob\nfrom jinja2 import Template\nfrom tqdm.auto import tqdm, trange\nfrom google.api_core.exceptions import Conflict, NotFound\nfrom google.cloud import bigquery\n\nimport matplotlib.pyplot as plt\nplt.style.use('seaborn-v0_8-whitegrid')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-11-26T04:59:51.452176Z","iopub.execute_input":"2023-11-26T04:59:51.452483Z","iopub.status.idle":"2023-11-26T04:59:51.459402Z","shell.execute_reply.started":"2023-11-26T04:59:51.452455Z","shell.execute_reply":"2023-11-26T04:59:51.458445Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\n","metadata":{"execution":{"iopub.status.busy":"2023-11-26T04:59:51.460721Z","iopub.execute_input":"2023-11-26T04:59:51.461088Z","iopub.status.idle":"2023-11-26T04:59:51.468917Z","shell.execute_reply.started":"2023-11-26T04:59:51.46105Z","shell.execute_reply":"2023-11-26T04:59:51.468063Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Add-onsからGCP ProjectをPROJECT_IDに登録してください\nPROJECT_ID = user_secrets.get_secret(\"PROJECT_ID\")\nDATASET_ID = \"otto\"","metadata":{"execution":{"iopub.status.busy":"2023-11-26T04:59:51.470348Z","iopub.execute_input":"2023-11-26T04:59:51.470652Z","iopub.status.idle":"2023-11-26T04:59:51.665947Z","shell.execute_reply.started":"2023-11-26T04:59:51.470628Z","shell.execute_reply":"2023-11-26T04:59:51.665234Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"logger = logging.getLogger()\nhandler = logging.StreamHandler(sys.stdout)\nhandler.setLevel(logging.INFO)\nlogger.addHandler(handler)\nlogger.setLevel(logging.INFO)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T04:59:51.668158Z","iopub.execute_input":"2023-11-26T04:59:51.66844Z","iopub.status.idle":"2023-11-26T04:59:51.673467Z","shell.execute_reply.started":"2023-11-26T04:59:51.668415Z","shell.execute_reply":"2023-11-26T04:59:51.672632Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Table of Contents\n今回のNotebookでは訓練部分までを行います。  \ntrain_test_datasetを用いれば、予測までを行えるはずなので、チャレンジしてみてください。\n* [BQデータセットの準備](#section1)\n* [候補作成](#section2)\n* [特徴量作成](#section3)\n* [訓練](#section4)","metadata":{"execution":{"iopub.status.busy":"2023-11-24T07:59:04.115014Z","iopub.execute_input":"2023-11-24T07:59:04.115491Z","iopub.status.idle":"2023-11-24T07:59:04.123352Z","shell.execute_reply.started":"2023-11-24T07:59:04.11546Z","shell.execute_reply":"2023-11-24T07:59:04.121552Z"}}},{"cell_type":"markdown","source":"## BQデータセットの準備  <a class=\"anchor\" id=\"section1\"></a>\nコンペデータセットをBQにアップロードする","metadata":{}},{"cell_type":"code","source":"class BigQueryClient:\n    def __init__(self, project_id: str) -> None:\n        self.project_id = project_id\n        self.client = bigquery.Client(project=project_id)\n\n    def create_dataset(self, dataset_id: str) -> None:\n        try:\n            dataset = self.client.dataset(dataset_id)\n            dataset.location = \"US\"\n            dataset = self.client.create_dataset(dataset)\n        except Conflict:\n            logger.info(f\"{dataset_id} is already created\")\n\n    def upload_dataset(self, df: pl.DataFrame, dataset_id: str, table_name: str,) -> None:\n        with io.BytesIO() as stream:\n            df.write_parquet(stream)\n            stream.seek(0)\n            job = self.client.load_table_from_file(\n                stream,\n                destination=f\"{self.project_id}.{dataset_id}.{table_name}\",\n                project=PROJECT_ID,\n                job_config=bigquery.LoadJobConfig(\n                    write_disposition=bigquery.WriteDisposition.WRITE_TRUNCATE,\n                    source_format=bigquery.SourceFormat.PARQUET,\n                    autodetect=True,\n                ),\n            )\n            job.result()\n            logger.info(f\"Finished uploading {table_name}.\")\n\n    def exist_table(self, dataset_id:str, table_name:str) -> bool:\n        ref = bigquery.DatasetReference(self.project_id, dataset_id).table(table_name)\n        try:\n            self.client.get_table(ref)\n            logger.info(f\"table: {dataset_id}.{table_name} exists.\")\n            return True\n        except NotFound:\n            logger.info(f\"table: {dataset_id}.{table_name} not found.\")\n            return False\n    \n    def execute_query(self, query: str) -> None:\n        job_config = bigquery.job.QueryJobConfig(\n            default_dataset=None,\n            allow_large_results=True,\n            dry_run=True,\n        )\n        query_job = self.client.query(query, job_config=job_config)\n        logger.info(f\"This query will process {query_job.total_bytes_processed / 1e9} GB.\")\n        job_config.dry_run = False\n        query_job = self.client.query(query, job_config=job_config)\n        query_job.result()\n        logger.info(\"Executed query.\")\n        \n    def read_gbq(self, query: str, progress_bar_type: str = \"tqdm\", use_pandas: bool = False) -> pl.DataFrame | pl.DataFrame:\n        if use_pandas:\n            return self.client.query(query).to_arrow(progress_bar_type=progress_bar_type, create_bqstorage_client=True).to_pandas()\n        else:\n            return pl.from_arrow(self.client.query(query).to_arrow(progress_bar_type=progress_bar_type, create_bqstorage_client=True))","metadata":{"execution":{"iopub.status.busy":"2023-11-26T04:59:51.674652Z","iopub.execute_input":"2023-11-26T04:59:51.674976Z","iopub.status.idle":"2023-11-26T04:59:51.690226Z","shell.execute_reply.started":"2023-11-26T04:59:51.674945Z","shell.execute_reply":"2023-11-26T04:59:51.689347Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"bq = BigQueryClient(project_id=PROJECT_ID)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T04:59:51.691216Z","iopub.execute_input":"2023-11-26T04:59:51.691493Z","iopub.status.idle":"2023-11-26T04:59:51.707305Z","shell.execute_reply.started":"2023-11-26T04:59:51.691465Z","shell.execute_reply":"2023-11-26T04:59:51.706417Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create Dataset\nbq.create_dataset(DATASET_ID)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T04:59:51.70837Z","iopub.execute_input":"2023-11-26T04:59:51.708666Z","iopub.status.idle":"2023-11-26T04:59:52.300843Z","shell.execute_reply.started":"2023-11-26T04:59:51.708642Z","shell.execute_reply":"2023-11-26T04:59:52.299998Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 訓練データセット  <a class=\"anchor\" id=\"subsection1\"></a>\nデータセットの分け方\n- validは正解が含まれるもののみに絞っている","metadata":{}},{"cell_type":"code","source":"# local to BQ\nif not bq.exist_table(dataset_id=DATASET_ID, table_name='train_valid'):\n    paths = glob(\"/kaggle/input/otto-validation/train_parquet/*.parquet\")\n    train_df = pl.concat([pl.read_parquet(path) for path in tqdm(paths)])\n\n    print('train size', train_df.shape)\n\n    paths = glob(\"/kaggle/input/otto-validation/test_parquet/*.parquet\")\n    valid_df = pl.concat([pl.read_parquet(path) for path in tqdm(paths)])\n\n    train_df = train_df.with_columns(pl.lit(\"train\").alias(\"split\"))\n    valid_df = valid_df.with_columns(pl.lit(\"valid\").alias(\"split\"))\n    df = pl.concat([train_df, valid_df])\n    df = df.with_columns(\n        pl.from_epoch(pl.col('ts'), time_unit='ms'),\n    )\n    bq.upload_dataset(df, dataset_id=DATASET_ID, table_name='train_valid')\n    df = df.with_columns(pl.col('ts').dt.date().alias('date'))\n    # データセットの分け方可視化\n    ax = df.group_by(['split', 'date']).agg(pl.count()).to_pandas().set_index(['date', 'split'])['count'].unstack(-1).plot.bar(figsize=(20, 5))\n    ax.set(ylabel='count')","metadata":{"execution":{"iopub.status.busy":"2023-11-26T04:59:52.30202Z","iopub.execute_input":"2023-11-26T04:59:52.302284Z","iopub.status.idle":"2023-11-26T04:59:52.490051Z","shell.execute_reply.started":"2023-11-26T04:59:52.302261Z","shell.execute_reply":"2023-11-26T04:59:52.489277Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"markdown","source":"### 正解ラベル　 <a class=\"anchor\" id=\"subsection2\"></a>","metadata":{}},{"cell_type":"code","source":"if not bq.exist_table(dataset_id=DATASET_ID, table_name='ground_truth'):\n    # train_parquetに対する正解ラベル\n    gt_df = pl.read_parquet(\"/kaggle/input/otto-validation/test_labels.parquet\")\n    gt_df = gt_df.pivot(index='session', columns='type', values='ground_truth').select(\n        'session', pl.exclude('session').name.suffix('_label')\n    )\n    # clickは直後のアクションのみ評価対象\n    gt_df = gt_df.with_columns(pl.col('clicks_label').list[0].alias('click_label'))\n    # local to BQ\n    bq.upload_dataset(gt_df, dataset_id=DATASET_ID, table_name='ground_truth')","metadata":{"execution":{"iopub.status.busy":"2023-11-26T04:59:52.491068Z","iopub.execute_input":"2023-11-26T04:59:52.491316Z","iopub.status.idle":"2023-11-26T04:59:52.672451Z","shell.execute_reply.started":"2023-11-26T04:59:52.491293Z","shell.execute_reply":"2023-11-26T04:59:52.671532Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### テストデータ　 <a class=\"anchor\" id=\"subsection3\"></a>\n訓練データとテストデータを組み合わせたデータを作成する","metadata":{}},{"cell_type":"code","source":"if not bq.exist_table(dataset_id=DATASET_ID, table_name='train_test'):\n    train_df = pl.read_parquet(\"/kaggle/input/otto-full-optimized-memory-footprint/train.parquet\")\n    test_df = pl.read_parquet(\"/kaggle/input/otto-full-optimized-memory-footprint/test.parquet\")\n    train_df = train_df.with_columns(pl.lit(\"train\").alias(\"split\"))\n    test_df = test_df.with_columns(pl.lit(\"test\").alias(\"split\"))\n    df = pl.concat([train_df, test_df])\n    df = df.with_columns(\n        pl.from_epoch(pl.col('ts'), time_unit='s'),\n        pl.col('type').map_dict({0: \"clicks\", 1: \"carts\", 2: \"orders\"})\n    )\n    bq.upload_dataset(df, dataset_id=DATASET_ID, table_name='train_test')","metadata":{"execution":{"iopub.status.busy":"2023-11-26T04:59:52.673535Z","iopub.execute_input":"2023-11-26T04:59:52.67382Z","iopub.status.idle":"2023-11-26T04:59:52.870287Z","shell.execute_reply.started":"2023-11-26T04:59:52.673795Z","shell.execute_reply":"2023-11-26T04:59:52.869433Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 候補作成  <a class=\"anchor\" id=\"section2\"></a>","metadata":{}},{"cell_type":"markdown","source":"- 今回は、以下の3つの候補を作成します\n  - そのsessionで訪問したaid (visited_items.sql)\n  - session内で前後1日以内にactionしたaid同士から共起行列を作って、それをsession毎の最新aidに対して適用して候補を作成 (covisit_items_1days.sql)\n  - session内で次にactionしたaid同士から共起行列を作って、それをsession毎の最新aidに対して適用して候補を作成 (covisit_items_bigram.sql)","metadata":{}},{"cell_type":"code","source":"def read_sql(sql_path: str, params: dict[str, str | int] = {}) -> str:\n    sql_name = os.path.basename(sql_path).split('.')[0]\n    params.update({\n        'sql_name': sql_name,\n        'project_id': PROJECT_ID,\n        'dataset_id': DATASET_ID,\n    })\n    with open(sql_path, 'r') as f:\n        query = Template(f.read())\n    query = query.render(params)\n    return query","metadata":{"execution":{"iopub.status.busy":"2023-11-26T04:59:52.871439Z","iopub.execute_input":"2023-11-26T04:59:52.871765Z","iopub.status.idle":"2023-11-26T04:59:52.877511Z","shell.execute_reply.started":"2023-11-26T04:59:52.871739Z","shell.execute_reply":"2023-11-26T04:59:52.876711Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile visited_items.sql\n\n{% for dataset_name in ['train_valid', 'train_test'] %}\n\nCREATE OR REPLACE TABLE `{{ project_id }}.{{ dataset_id }}.{{ sql_name }}_{{ dataset_name }}_feature`\n\nAS\n\nWITH base AS (\n  SELECT\n    session,\n    aid,\n    ts,\n  FROM `{{ project_id }}.{{ dataset_id }}.{{ dataset_name }}`\n  WHERE split != 'train'\n), session_ts AS (\n  SELECT\n    session,\n    MAX(ts) as max_session_ts,\n  FROM base AS a\n  GROUP BY session\n)\n\nSELECT\n  session,\n  aid,\n  TIMESTAMP_DIFF(max_session_ts, MAX(ts), SECOND) as seconds_to_latest,\n  ROW_NUMBER() OVER(PARTITION BY session ORDER BY TIMESTAMP_DIFF(max_session_ts, MAX(ts), SECOND)) as rank_seconds_to_latest,\nFROM base\nLEFT JOIN session_ts USING(session)\nGROUP BY session, aid, max_session_ts\n;\n    \nCREATE OR REPLACE TABLE `{{ project_id }}.{{ dataset_id }}.{{ sql_name }}_{{ dataset_name }}_candidates`\n\n  AS\n\n  SELECT * FROM `{{ project_id }}.{{ dataset_id }}.{{ sql_name }}_{{ dataset_name }}_feature`\n  WHERE\n    rank_seconds_to_latest <= 30\n;\n{% endfor %}\n","metadata":{"execution":{"iopub.status.busy":"2023-11-26T04:59:52.878841Z","iopub.execute_input":"2023-11-26T04:59:52.879125Z","iopub.status.idle":"2023-11-26T04:59:52.887899Z","shell.execute_reply.started":"2023-11-26T04:59:52.8791Z","shell.execute_reply":"2023-11-26T04:59:52.887017Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile covisit_items_1days.sql\n\n{% for dataset_name in ['train_valid', 'train_test'] %}\n\nCREATE OR REPLACE TABLE `{{ project_id }}.{{ dataset_id }}.{{ sql_name }}_{{ dataset_name }}_feature`\n\nAS\n\nWITH base AS (\n  SELECT\n    session,\n    aid,\n    ts,\n    type,\n    CASE \n      WHEN type = 'clicks' THEN 1\n      WHEN type = 'carts' THEN 6\n      WHEN type = 'orders' THEN 4\n    END AS weight,\n    split,\n  FROM `{{ project_id }}.{{ dataset_id }}.{{ dataset_name }}`\n), pairs AS (\n    SELECT\n      a.aid AS from_aid,\n      b.aid AS to_aid,\n      SUM(b.weight) AS sum_weight,\n      COUNT(1) AS transition_count,\n    FROM base AS a\n    LEFT JOIN base AS b ON a.session = b.session\n    WHERE\n      a.ts != b.ts\n      AND\n      -- 前後24時間のアクションを集める\n      abs(timestamp_diff(a.ts, b.ts, HOUR)) <= 24\n    GROUP BY 1, 2\n    -- 2回以上登場するpairのみに絞る\n    HAVING transition_count > 1\n), latest_aids AS (\n  -- 最後にアクションしたaidを取ってくる\n  SELECT\n    session,\n    aid,\n  FROM `{{ project_id }}.{{ dataset_id }}.{{ dataset_name }}`\n  -- matrix作成には全期間使うが訓練にはvalidの期間のみ用いる\n  WHERE split != \"train\"\n  QUALIFY ROW_NUMBER() OVER(PARTITION BY session ORDER BY ts DESC) = 1\n)\n\nSELECT\n  session,\n  to_aid as aid,\n  {% for fn in ['sum', 'avg', 'stddev'] %}\n    {{ fn }}(sum_weight) AS {{ fn }}_{{ sql_name }}_sum_weight,\n    {{ fn }}(transition_count) AS {{ fn }}_{{ sql_name }}_transition_count,\n  {% endfor %}\n  COUNT(*) AS {{ sql_name }}_count,\n  ROW_NUMBER() OVER(PARTITION BY session ORDER BY SUM(sum_weight) DESC) as rn_{{ sql_name }}_sum_weight,\n  ROW_NUMBER() OVER(PARTITION BY session ORDER BY SUM(transition_count) DESC) as rn_{{ sql_name }}_transition_count,\n  ROW_NUMBER() OVER(PARTITION BY session ORDER BY COUNT(1) DESC) as rn_{{ sql_name }}_aid_count,\nFROM latest_aids\nLEFT JOIN pairs\nON latest_aids.aid = pairs.from_aid\nWHERE to_aid is not NULL\nGROUP BY 1, 2\nQUALIFY\n  -- 特徴量のために多めに残しておく\n  rn_{{ sql_name }}_sum_weight <= 500\n  OR rn_{{ sql_name }}_transition_count <= 500\n  OR rn_{{ sql_name }}_aid_count <= 500\n\n;\n    \nCREATE OR REPLACE TABLE `{{ project_id }}.{{ dataset_id }}.{{ sql_name }}_{{ dataset_name }}_candidates`\n\nAS\n\nSELECT DISTINCT\n  session,\n  aid,\nFROM `{{ project_id }}.{{ dataset_id }}.{{ sql_name }}_{{ dataset_name }}_feature`\nWHERE\n  rn_{{ sql_name }}_sum_weight <= {{ top_k }}\n  OR rn_{{ sql_name }}_transition_count <= {{ top_k }}\n  OR rn_{{ sql_name }}_aid_count <= {{ top_k }}\n;\n{% endfor %}\n","metadata":{"execution":{"iopub.status.busy":"2023-11-26T04:59:52.892665Z","iopub.execute_input":"2023-11-26T04:59:52.892921Z","iopub.status.idle":"2023-11-26T04:59:52.900839Z","shell.execute_reply.started":"2023-11-26T04:59:52.892898Z","shell.execute_reply":"2023-11-26T04:59:52.90003Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile covisit_items_bigram.sql\n\n{% for dataset_name in ['train_valid', 'train_test'] %}\n\nCREATE OR REPLACE TABLE `{{ project_id }}.{{ dataset_id }}.{{ sql_name }}_{{ dataset_name }}_feature`\n\nAS\n\nWITH base AS (\n  SELECT\n    session,\n    aid,\n    type,\n    ts,\n    LEAD(aid) OVER(PARTITION BY session ORDER BY ts) AS lead_aid,\n    LEAD(type) OVER(PARTITION BY session ORDER BY ts) AS lead_type,\n    ROW_NUMBER() OVER(PARTITION BY session ORDER BY ts) AS rn\n  FROM `{{ project_id }}.{{ dataset_id }}.{{ dataset_name }}`\n  QUALIFY lead_aid IS NOT NULL\n),\n\npairs AS (\n  SELECT\n    aid AS from_aid,\n    lead_aid AS to_aid,\n    count(*) AS transition_count,\n    count(DISTINCT session) AS transition_user_count\n  FROM base\n  GROUP BY 1, 2\n  HAVING transition_count > 1\n),\n\naid_count AS (\n  SELECT\n    from_aid AS aid,\n    sum(transition_count) AS sum_transition_count,\n    sum(transition_user_count) AS sum_transition_user_count\n  FROM pairs\n  GROUP BY 1\n),\n\ncvr_count AS (\n  SELECT\n    from_aid,\n    to_aid,\n    transition_count / sqrt(coalesce(r1.sum_transition_count, 1) * coalesce(r2.sum_transition_count, 1)) AS norm_transition_count,\n    transition_user_count / sqrt(coalesce(r1.sum_transition_user_count, 1) * coalesce(r2.sum_transition_user_count, 1)) AS norm_transition_user_count\n  FROM pairs\n  LEFT JOIN aid_count AS r1\n    ON pairs.from_aid = r1.aid\n  LEFT JOIN aid_count AS r2\n    ON pairs.to_aid = r2.aid\n),\n\ncovisit AS (\n  SELECT *\n  FROM (\n    SELECT DISTINCT *\n    FROM (\n      SELECT\n        from_aid,\n        to_aid\n      FROM cvr_count\n      QUALIFY\n        --  上位500に絞る\n        ROW_NUMBER() OVER(PARTITION BY from_aid ORDER BY norm_transition_count DESC) <= 500\n        OR\n        ROW_NUMBER() OVER(PARTITION BY from_aid ORDER BY norm_transition_user_count DESC) <= 500\n    )\n  )\n  LEFT JOIN cvr_count USING (from_aid, to_aid)\n),\n\nlatest_aids AS (\n  -- 最後にアクションしたaidを取ってくる\n  SELECT\n    session,\n    aid,\n  FROM `{{ project_id }}.{{ dataset_id }}.{{ dataset_name }}`\n  -- matrix作成には全期間使うが訓練にはvalidの期間のみ用いる\n  WHERE split != \"train\"\n  QUALIFY ROW_NUMBER() OVER(PARTITION BY session ORDER BY ts DESC) = 1\n)\n\nSELECT\n  session,\n  to_aid as aid,\n  -- feature\n  {% for fn in ['sum', 'avg', 'stddev'] %}\n    {{ fn }}(norm_transition_count) AS {{ fn }}_{{ sql_name }}_norm_transition_count,\n    {{ fn }}(norm_transition_user_count) AS {{ fn }}_{{ sql_name }}_norm_transition_user_count,\n  {% endfor %}\n  COUNT(*) AS {{ sql_name }}_count,\n  ROW_NUMBER() OVER(PARTITION BY session ORDER BY sum(norm_transition_count) DESC) AS rn_{{ sql_name }}_norm_transition_count,\n  ROW_NUMBER() OVER(PARTITION BY session ORDER BY sum(norm_transition_user_count) DESC) AS rn_{{ sql_name }}_norm_transition_user_count,\n  ROW_NUMBER() OVER(PARTITION BY session ORDER BY count(*) DESC) AS rn_{{ sql_name }}_aid_count\nFROM latest_aids AS l\nLEFT JOIN covisit AS r\n  ON l.aid = r.from_aid\nWHERE to_aid IS NOT NULL\nGROUP BY 1, 2\nQUALIFY\n  -- 特徴量のために多めに残しておく\n  rn_{{ sql_name }}_norm_transition_count <= 500\n  OR rn_{{ sql_name }}_norm_transition_user_count <= 500\n  OR rn_{{ sql_name }}_aid_count <= 500\n\n;\n    \nCREATE OR REPLACE TABLE `{{ project_id }}.{{ dataset_id }}.{{ sql_name }}_{{ dataset_name }}_candidates`\n\nAS\n\nSELECT DISTINCT\n  session,\n  aid,\nFROM `{{ project_id }}.{{ dataset_id }}.{{ sql_name }}_{{ dataset_name }}_feature`\nWHERE\n  rn_{{ sql_name }}_norm_transition_count <= {{ top_k }}\n  OR rn_{{ sql_name }}_norm_transition_user_count <= {{ top_k }}\n  OR rn_{{ sql_name }}_aid_count <= {{ top_k }}\n;\n{% endfor %}\n","metadata":{"execution":{"iopub.status.busy":"2023-11-26T04:59:52.902014Z","iopub.execute_input":"2023-11-26T04:59:52.902261Z","iopub.status.idle":"2023-11-26T04:59:52.917283Z","shell.execute_reply.started":"2023-11-26T04:59:52.902237Z","shell.execute_reply":"2023-11-26T04:59:52.916484Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for sql_path in ['visited_items.sql', 'covisit_items_1days.sql', 'covisit_items_bigram.sql']:\n    sql_name = sql_path.split(\".\")[0]\n    if not bq.exist_table(dataset_id=DATASET_ID, table_name=f'{sql_name}_train_valid_candidates'):\n        query = read_sql(sql_path, params={'top_k': 30})\n        bq.execute_query(query)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T04:59:52.91854Z","iopub.execute_input":"2023-11-26T04:59:52.918905Z","iopub.status.idle":"2023-11-26T04:59:53.468102Z","shell.execute_reply.started":"2023-11-26T04:59:52.918874Z","shell.execute_reply":"2023-11-26T04:59:53.46728Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"candidates = [                          \n    \"visited_items\",\n    \"covisit_items_1days\",\n    \"covisit_items_bigram\",\n]","metadata":{"execution":{"iopub.status.busy":"2023-11-26T04:59:53.469209Z","iopub.execute_input":"2023-11-26T04:59:53.469472Z","iopub.status.idle":"2023-11-26T04:59:53.473887Z","shell.execute_reply.started":"2023-11-26T04:59:53.469448Z","shell.execute_reply":"2023-11-26T04:59:53.472922Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### candidatesをまとめる","metadata":{}},{"cell_type":"code","source":"%%writefile user_items.sql\n\n{% for dataset_name in ['train_valid', 'train_test'] %}      \n\n  CREATE OR REPLACE TABLE `{{ project_id }}.{{ dataset_id }}.{{ dataset_name }}_{{ sql_name }}`\n  CLUSTER BY session\n\n  AS\n\n  WITH\n  sampled_session AS (\n    -- 最後の1weekだけで学習\n    SELECT DISTINCT session FROM `{{ project_id }}.{{ dataset_id }}.{{ dataset_name }}`\n    WHERE split != \"train\"\n  )\n\n  SELECT DISTINCT\n    session,\n    aid,\n  FROM (\n      {% for candidate in candidates %}\n        SELECT\n          session,\n          aid\n        FROM `{{ project_id }}.{{ dataset_id }}.{{ candidate }}_{{ dataset_name }}_candidates`\n      {% if not loop.last %}\n        UNION ALL\n      {% endif %}\n      {% endfor %}\n    )\n  WHERE session IN (SELECT session FROM sampled_session);\n{% endfor %}\n","metadata":{"execution":{"iopub.status.busy":"2023-11-26T04:59:53.474987Z","iopub.execute_input":"2023-11-26T04:59:53.47529Z","iopub.status.idle":"2023-11-26T04:59:53.487433Z","shell.execute_reply.started":"2023-11-26T04:59:53.475265Z","shell.execute_reply":"2023-11-26T04:59:53.486604Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not bq.exist_table(dataset_id=DATASET_ID, table_name='train_valid_user_items'):\n    query = read_sql('user_items.sql', params={'candidates': candidates})\n    bq.execute_query(query)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T04:59:53.488361Z","iopub.execute_input":"2023-11-26T04:59:53.488647Z","iopub.status.idle":"2023-11-26T04:59:53.708985Z","shell.execute_reply.started":"2023-11-26T04:59:53.488624Z","shell.execute_reply":"2023-11-26T04:59:53.708193Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 特徴量作成  <a class=\"anchor\" id=\"section3\"></a>\n今回は以下の二つのみを作成します。\nsession x aidの特徴など更に色々作れると思います\n- sessionの特徴\n- aidの特徴","metadata":{}},{"cell_type":"code","source":"%%writefile session_feature.sql\n\n{% for dataset_name in ['train_valid', 'train_test'] %}\n\n  CREATE OR REPLACE TABLE `{{ project_id }}.{{ dataset_id }}.{{ sql_name }}_{{ dataset_name }}`\n\n  AS\n\n  WITH base AS (\n    SELECT\n      session,\n      aid,\n      type,\n      ts\n    FROM `{{ project_id }}.{{ dataset_id }}.{{ dataset_name }}`\n  )\n\n  SELECT\n    session,\n    timestamp_diff(max(ts), min(ts), SECOND) AS search_secnonds,\n    count(*) AS session_size,\n    count(DISTINCT if(type = \"clicks\", aid, null)) AS click_aid_session_count,\n    count(DISTINCT if(type = \"carts\", aid, null)) AS cart_aid_session_count,\n    count(DISTINCT if(type = \"orders\", aid, null)) AS order_aid_session_count\n  FROM base\n  GROUP BY session;\n{% endfor %}","metadata":{"execution":{"iopub.status.busy":"2023-11-26T04:59:53.71007Z","iopub.execute_input":"2023-11-26T04:59:53.710366Z","iopub.status.idle":"2023-11-26T04:59:53.716276Z","shell.execute_reply.started":"2023-11-26T04:59:53.710335Z","shell.execute_reply":"2023-11-26T04:59:53.715325Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile aid_feature.sql\n\n{% for dataset_name in ['train_valid', 'train_test'] %}\n\n  CREATE OR REPLACE TABLE `{{ project_id }}.{{ dataset_id }}.{{ sql_name }}_{{ dataset_name }}`\n\n  AS\n\n  WITH base AS (\n    SELECT\n      session,\n      aid,\n      type,\n      ts\n    FROM `{{ project_id }}.{{ dataset_id }}.{{ dataset_name }}`\n  ),\n\n  n_session_feature AS (\n    -- 期間中全てのアクションから計算\n    SELECT\n      aid,\n      n_session_clicks,\n      n_session_carts,\n      n_session_orders\n    FROM (\n      SELECT\n        aid,\n        type,\n        count(DISTINCT session) AS n_session\n      FROM base\n      GROUP BY 1, 2\n    )\n    PIVOT (sum(n_session) AS n_session FOR type IN (\"clicks\", \"carts\", \"orders\"))\n  ),\n\n  count_feature AS (\n    -- 期間中全てのアクションから計算\n    SELECT\n      aid,\n      action_count_clicks,\n      action_count_carts,\n      action_count_orders\n    FROM (\n      SELECT\n        aid,\n        type,\n        count(*) AS action_count\n      FROM base\n      GROUP BY 1, 2\n    )\n    PIVOT (sum(action_count) AS action_count FOR type IN (\"clicks\", \"carts\", \"orders\"))\n  ),\n\n  time_feature AS (\n    -- dataset作成時に使う\n    -- 最後のアクションからどれくらい時間が経ってるか\n    -- 最初のアクションからどれくらい時間が経ってるか\n    SELECT\n      aid,\n      min(ts) AS action_min_ts,\n      max(ts) AS action_max_ts\n    FROM base\n    GROUP BY 1\n  ),\n\n  cvr_feature AS (\n    SELECT\n      aid,\n      (n_session_carts + n_session_orders) / (n_session_clicks + avg(n_session_clicks) OVER()) AS n_session_cvr_click_to_cart_order,\n      (action_count_carts + action_count_orders) / (action_count_clicks + avg(action_count_clicks) OVER()) AS action_count_cvr_click_to_cart_order,\n      (n_session_carts) / (n_session_clicks + avg(n_session_clicks) OVER()) AS n_session_cvr_click_to_cart,\n      (action_count_carts) / (action_count_clicks + avg(action_count_clicks) OVER()) AS action_count_cvr_click_to_cart,\n      (n_session_orders) / (n_session_clicks + avg(n_session_clicks) OVER()) AS n_session_cvr_click_to_order,\n      (action_count_orders) / (action_count_clicks + avg(action_count_clicks) OVER()) AS action_count_cvr_click_to_order,\n      (n_session_orders) / (n_session_carts + avg(n_session_carts) OVER()) AS n_session_cvr_cart_to_order,\n      (action_count_orders) / (action_count_carts + avg(action_count_carts) OVER()) AS action_count_cvr_cart_to_order\n    FROM n_session_feature\n    LEFT JOIN count_feature USING (aid)\n  )\n\n  SELECT *\n  FROM n_session_feature\n  LEFT JOIN count_feature USING (aid)\n  LEFT JOIN time_feature USING (aid)\n  LEFT JOIN cvr_feature USING (aid);\n{% endfor %}\n","metadata":{"execution":{"iopub.status.busy":"2023-11-26T04:59:53.717395Z","iopub.execute_input":"2023-11-26T04:59:53.717681Z","iopub.status.idle":"2023-11-26T04:59:53.730118Z","shell.execute_reply.started":"2023-11-26T04:59:53.717653Z","shell.execute_reply":"2023-11-26T04:59:53.728991Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for sql_path in ['session_feature.sql', 'aid_feature.sql']:\n    sql_name = sql_path.split(\".\")[0]\n    if not bq.exist_table(dataset_id=DATASET_ID, table_name=f'{sql_name}_train_valid'):\n        query = read_sql(sql_path)\n        bq.execute_query(query)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T04:59:53.73129Z","iopub.execute_input":"2023-11-26T04:59:53.731608Z","iopub.status.idle":"2023-11-26T04:59:54.100278Z","shell.execute_reply.started":"2023-11-26T04:59:53.731578Z","shell.execute_reply":"2023-11-26T04:59:54.099475Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## データセット作成","metadata":{}},{"cell_type":"code","source":"%%writefile dataset.sql\n\n{% for dataset_name in ['train_valid', 'train_test'] %}    \n\n  CREATE OR REPLACE TABLE `{{ project_id }}.{{ dataset_id }}.{{ dataset_name }}_{{ sql_name }}`\n  CLUSTER BY session\n\n  AS\n\n  WITH session_ts AS (\n    SELECT\n      session,\n      timestamp_trunc(timestamp_add(max(ts), INTERVAL 1 SECOND), MINUTE) AS ts_minute,\n    FROM `{{ project_id }}.{{ dataset_id }}.{{ dataset_name }}`\n    GROUP BY 1\n  ), joined AS (\n    SELECT\n      * EXCEPT(\n        click_label, clicks_label, carts_label, orders_label, ts_minute, action_min_ts, action_max_ts\n      ),\n      cast(aid IN (select l.item from unnest(carts_label.list) as l) AS int64) AS cart_label,\n      cast(aid IN (select l.item from unnest(orders_label.list) as l) AS int64) AS order_label,\n      timestamp_diff(ts_minute, action_max_ts, SECOND) AS seconds_from_last_aid_action,\n      timestamp_diff(ts_minute, action_min_ts, SECOND) AS seconds_from_first_aid_action,\n      coalesce(cast(aid = click_label AS int64), 0) AS click_label\n    FROM (\n      SELECT DISTINCT\n        session,\n        aid,\n      FROM `{{ project_id }}.{{ dataset_id }}.{{ dataset_name }}_user_items`\n    )\n    LEFT JOIN `{{ project_id }}.{{ dataset_id }}.ground_truth` USING (session)\n    LEFT JOIN session_ts USING (session)\n    LEFT JOIN `{{ project_id }}.{{ dataset_id }}.session_feature_{{ dataset_name }}` USING (session)\n    LEFT JOIN `{{ project_id }}.{{ dataset_id }}.aid_feature_{{ dataset_name }}` USING (aid)\n    {% for candidate in candidates %}    \n      LEFT JOIN `{{ project_id }}.{{ dataset_id }}.{{ candidate }}_{{ dataset_name }}_feature` USING (session, aid)\n    {% endfor %}\n  )\n\n  SELECT * FROM joined;\n\n    {% endfor %}\n","metadata":{"execution":{"iopub.status.busy":"2023-11-26T04:59:54.10129Z","iopub.execute_input":"2023-11-26T04:59:54.101529Z","iopub.status.idle":"2023-11-26T04:59:54.107936Z","shell.execute_reply.started":"2023-11-26T04:59:54.101507Z","shell.execute_reply":"2023-11-26T04:59:54.107017Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not bq.exist_table(dataset_id=DATASET_ID, table_name='train_valid_dataset'):\n    query = read_sql('dataset.sql', params={'candidates': candidates})\n    bq.execute_query(query)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T04:59:54.109051Z","iopub.execute_input":"2023-11-26T04:59:54.109383Z","iopub.status.idle":"2023-11-26T04:59:54.303922Z","shell.execute_reply.started":"2023-11-26T04:59:54.10935Z","shell.execute_reply":"2023-11-26T04:59:54.303108Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### ネガティブサンプリング","metadata":{}},{"cell_type":"code","source":"%%writefile sampled_dataset.sql\n\n{% for dataset_name in ['train_valid'] %}\n\n  CREATE OR REPLACE TABLE `{{ project_id }}.{{ dataset_id }}.{{ dataset_name }}_{{ sql_name }}`\n  CLUSTER BY session, fold\n\n  AS\n\n  WITH base AS (\n    SELECT * FROM `{{ project_id }}.{{ dataset_id }}.{{ dataset_name }}_dataset`\n  ),\n\n  folds AS (\n    SELECT\n      session,\n      MOD(ROW_NUMBER() OVER(ORDER BY FARM_FINGERPRINT(CONCAT({{ seed }}, CAST(session AS STRING)))), 5) AS fold\n    FROM (SELECT DISTINCT session FROM base)\n  ),\n\n  positive AS (\n    SELECT *\n    FROM base\n    WHERE click_label = 1 OR cart_label = 1 OR order_label = 1\n  ),\n\n  positive_count AS (\n    SELECT count(*) AS positive_size\n    FROM positive\n  ),\n\n  negative AS (\n    SELECT base.*\n    FROM base, positive_count\n    WHERE\n      click_label != 1 AND cart_label != 1 AND order_label != 1\n    -- 正例の10倍の負例をsampling\n    QUALIFY ROW_NUMBER() OVER(ORDER BY FARM_FINGERPRINT(CONCAT({{ seed }}, session, aid))) <= 10 * positive_size\n  )\n\n  SELECT\n    *,\n    cast(cart_label + order_label >= 1 AS INT64) AS cart_order_label\n  FROM (\n      SELECT * FROM positive\n      UNION ALL\n      SELECT * FROM negative\n    )\n  LEFT JOIN folds USING (session);\n{% endfor %}\n","metadata":{"execution":{"iopub.status.busy":"2023-11-26T04:59:54.3051Z","iopub.execute_input":"2023-11-26T04:59:54.305377Z","iopub.status.idle":"2023-11-26T04:59:54.313856Z","shell.execute_reply.started":"2023-11-26T04:59:54.305354Z","shell.execute_reply":"2023-11-26T04:59:54.313048Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not bq.exist_table(dataset_id=DATASET_ID, table_name='train_valid_sampled_dataset'):\n    query = read_sql('sampled_dataset.sql', params={'seed': '1024'})\n    bq.execute_query(query)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T04:59:54.315021Z","iopub.execute_input":"2023-11-26T04:59:54.315552Z","iopub.status.idle":"2023-11-26T04:59:54.498602Z","shell.execute_reply.started":"2023-11-26T04:59:54.315527Z","shell.execute_reply":"2023-11-26T04:59:54.497801Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 訓練  <a class=\"anchor\" id=\"section4\"></a>","metadata":{}},{"cell_type":"code","source":"import pickle\nimport catboost as cat\nfrom omegaconf import DictConfig, OmegaConf","metadata":{"execution":{"iopub.status.busy":"2023-11-26T05:22:54.879974Z","iopub.execute_input":"2023-11-26T05:22:54.88084Z","iopub.status.idle":"2023-11-26T05:22:54.88522Z","shell.execute_reply.started":"2023-11-26T05:22:54.880806Z","shell.execute_reply":"2023-11-26T05:22:54.884216Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CatModel(object):\n    \"\"\"\n    label_col毎にcatboost modelを作成するためのクラス\n    \"\"\"\n\n    def __init__(\n        self,\n        config: DictConfig,\n    ):\n        self.config = config\n        self.model_dicts: Dict[int, cat.Booster] = {}\n\n    def store_model(self, bst: cat.CatBoost, n_fold: int) -> None:\n        self.model_dicts[n_fold] = bst\n\n    def store_importance(self, importance_df: pd.DataFrame) -> None:\n        self.importance_df = importance_df\n\n    def cv(\n        self,\n        df: pd.DataFrame,\n    ) -> pl.DataFrame:\n        importances = []\n        preds = np.zeros(len(df))\n        df = df.with_columns(pl.arange(0, len(df)).alias(\"index\"))\n        for n_fold in range(self.config.n_fold):\n            train_df = df.filter(df[\"fold\"] != n_fold)\n            valid_df = df.filter(df[\"fold\"] == n_fold)\n            logger.info(\n                f\"{self.config.label_col}[fold {n_fold}] train shape: {train_df.shape}, valid shape: {valid_df.shape}\"\n            )\n            bst, importance = self.fit(train_df, valid_df)\n            valid_pool = cat.Pool(\n                valid_df[self.config.feature_cols].to_numpy(),\n                cat_features=list(self.config.categorical_features_indices),\n            )\n            preds[valid_df[\"index\"].to_numpy()] = bst.predict(valid_pool)\n            self.store_model(bst, n_fold)\n            importances.append(importance)\n        df = df.with_columns(pl.Series(self.config.pred_col, preds))\n        importances_mean = np.mean(importances, axis=0)\n        importances_std = np.std(importances, axis=0)\n        importance_df = pd.DataFrame(\n            {\"mean\": importances_mean, \"std\": importances_std},\n            index=self.config.feature_cols,\n        ).sort_values(by=\"mean\", ascending=False)\n        self.store_importance(importance_df)\n        return df\n\n    def fit(\n        self,\n        train_df: pl.DataFrame,\n        valid_df: pl.DataFrame,\n    ) -> cat.CatBoost:\n\n        X_train = train_df.select(self.config.feature_cols)\n        y_train = train_df.select(self.config.label_col)\n\n        X_valid = valid_df.select(self.config.feature_cols)\n        y_valid = valid_df.select(self.config.label_col)\n        dtrain = cat.Pool(\n            X_train.to_numpy(),\n            label=y_train.to_numpy(),\n            feature_names=self.config.feature_cols,\n            cat_features=list(self.config.categorical_features_indices),\n        )\n        dvalid = cat.Pool(\n            X_valid.to_numpy(),\n            label=y_valid.to_numpy(),\n            feature_names=self.config.feature_cols,\n            cat_features=list(self.config.categorical_features_indices),\n        )\n        if self.config.params.loss_function in [\n            \"YetiRank\",\n            \"PairLogit\",\n            \"PairLogitPairwise\",\n        ]:\n            dtrain.set_group_id(train_df[\"session\"].to_numpy())\n            dvalid.set_group_id(valid_df[\"session\"].to_numpy())\n\n        bst = cat.train(\n            pool=dtrain,\n            params=dict(self.config.params),\n            evals=dvalid,\n            early_stopping_rounds=100,\n            verbose_eval=100,\n        )\n        importance = bst.get_feature_importance(dtrain, type=\"FeatureImportance\")\n        return bst, importance\n\n    def save_model(self, model_dir: str, suffix: str = \"\") -> None:\n        with open(\n            f\"{model_dir}/booster_{self.config.label_col + suffix}.pkl\", \"wb\"\n        ) as f:\n            pickle.dump(self.model_dicts, f)\n\n    def save_importance(\n        self,\n        result_path: str,\n        suffix: str = \"\",\n    ) -> None:\n        self.importance_df.sort_values(\"mean\").iloc[-50:].plot.barh(\n            xerr=\"std\", figsize=(10, 20)\n        )\n        plt.tight_layout()\n        plt.savefig(\n            os.path.join(\n                result_path,\n                f\"importance_{self.config.label_col + suffix}.png\",\n            )\n        )\n        self.importance_df.name = \"feature_name\"\n        self.importance_df = self.importance_df.reset_index().sort_values(\n            by=\"mean\", ascending=False\n        )\n        self.importance_df.to_csv(\n            os.path.join(\n                result_path,\n                f\"importance_{self.config.label_col + suffix}.csv\",\n            ),\n            index=False,\n        )\n","metadata":{"execution":{"iopub.status.busy":"2023-11-26T05:22:27.711458Z","iopub.execute_input":"2023-11-26T05:22:27.711887Z","iopub.status.idle":"2023-11-26T05:22:27.734488Z","shell.execute_reply.started":"2023-11-26T05:22:27.711854Z","shell.execute_reply":"2023-11-26T05:22:27.733524Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"config = '''\nseed: 777\ncatboost:\n  n_fold: 5\n  feature_cols:\n  cat_cols: []\n  label_col:\n  pred_col:\n  early_stopping_rounds: 200\n  verbose_eval: 200\n  categorical_features_indices: []\n  params:\n    task_type: GPU\n    iterations: 100000\n    loss_function: PairLogit\n    eval_metric: PairLogit\n    custom_metric: PairLogit\n    max_depth: 8\n    learning_rate: 0.5\n    max_bin: 32\n    verbose: 100\n    devices: 0:1\n    use_best_model: True\n    od_type: Iter\n    od_wait: 100\n    random_seed: ${seed}\n    gpu_ram_part: 0.95\n'''","metadata":{"execution":{"iopub.status.busy":"2023-11-26T05:22:28.788732Z","iopub.execute_input":"2023-11-26T05:22:28.789095Z","iopub.status.idle":"2023-11-26T05:22:28.793804Z","shell.execute_reply.started":"2023-11-26T05:22:28.789065Z","shell.execute_reply":"2023-11-26T05:22:28.792833Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"config = OmegaConf.create(config)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T04:59:55.513064Z","iopub.execute_input":"2023-11-26T04:59:55.51343Z","iopub.status.idle":"2023-11-26T04:59:55.531278Z","shell.execute_reply.started":"2023-11-26T04:59:55.513404Z","shell.execute_reply":"2023-11-26T04:59:55.530385Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = bq.read_gbq(\n    f'SELECT * FROM `{PROJECT_ID}.{DATASET_ID}.train_valid_sampled_dataset`', \n    progress_bar_type='tqdm_notebook'\n)","metadata":{"execution":{"iopub.status.busy":"2023-11-26T04:59:55.532201Z","iopub.execute_input":"2023-11-26T04:59:55.532441Z","iopub.status.idle":"2023-11-26T05:00:42.202193Z","shell.execute_reply.started":"2023-11-26T04:59:55.532419Z","shell.execute_reply":"2023-11-26T05:00:42.201363Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ランキングロスを使う場合は、sessionごとに並び替える\ndf = df.sort(['session', 'aid'])","metadata":{"execution":{"iopub.status.busy":"2023-11-26T05:00:42.203569Z","iopub.execute_input":"2023-11-26T05:00:42.203856Z","iopub.status.idle":"2023-11-26T05:00:51.79761Z","shell.execute_reply.started":"2023-11-26T05:00:42.203825Z","shell.execute_reply":"2023-11-26T05:00:51.796606Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"config.catboost.feature_cols = [\n    col\n    for col in df.columns\n    if col not in [\"session\", \"aid\", \"click_label\", \"cart_label\", \"order_label\", \"cart_order_label\", \"fold\", \"action_max_ts\", \"action_min_ts\"]\n]","metadata":{"execution":{"iopub.status.busy":"2023-11-26T05:00:51.799029Z","iopub.execute_input":"2023-11-26T05:00:51.799417Z","iopub.status.idle":"2023-11-26T05:00:51.81112Z","shell.execute_reply.started":"2023-11-26T05:00:51.799375Z","shell.execute_reply":"2023-11-26T05:00:51.81032Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for label_col in ['click_label', 'cart_order_label']:\n    config.catboost.label_col = label_col\n    config.catboost.pred_col = f\"{label_col}_pred\"\n    model = CatModel(config.catboost)\n    df = model.cv(df)\n    model.save_model('./', suffix=f'_{label_col}')\n    model.save_importance('./', suffix=f'_{label_col}')\n","metadata":{"execution":{"iopub.status.busy":"2023-11-26T05:00:51.812306Z","iopub.execute_input":"2023-11-26T05:00:51.812654Z","iopub.status.idle":"2023-11-26T05:14:03.847079Z","shell.execute_reply.started":"2023-11-26T05:00:51.812623Z","shell.execute_reply":"2023-11-26T05:14:03.845795Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.write_parquet('pred_results.parquet')","metadata":{"execution":{"iopub.status.busy":"2023-11-26T05:14:03.856336Z","iopub.status.idle":"2023-11-26T05:14:03.856729Z","shell.execute_reply.started":"2023-11-26T05:14:03.85654Z","shell.execute_reply":"2023-11-26T05:14:03.856569Z"},"trusted":true},"execution_count":null,"outputs":[]}]}