{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# H&M Recsys challenge\n\nThis notebook is a working demo of the new Feature Aggregator from skrub!\n\nThe dataset we use is a list of transactions, customers x articles, with metadata for both. The goal is to predict for each customer the next 12 articles to recommend.\n\nWe benefit a lot from the lazy mode of polars!\n\n**Our plan**\n\n- Perform a double aggregation on the transactions table (by customer and by article) in one-shot using the Feature Aggregator\n- Create the embeddings of the articles based on their description and clothing type, using `skrub.MinHashEncoder`\n- Create the embeddings of the customer by averaging the embeddings of their purchased articles\n- Join both embeddings to the transactions table to enrich it","metadata":{}},{"cell_type":"markdown","source":"## Loading data","metadata":{}},{"cell_type":"code","source":"from pathlib import Path\n\ninput_dir = Path(\"../input/h-and-m-personalized-fashion-recommendations\")\nlist(input_dir.iterdir())","metadata":{"execution":{"iopub.status.busy":"2023-06-15T11:06:09.307238Z","iopub.execute_input":"2023-06-15T11:06:09.307628Z","iopub.status.idle":"2023-06-15T11:06:09.318341Z","shell.execute_reply.started":"2023-06-15T11:06:09.307595Z","shell.execute_reply":"2023-06-15T11:06:09.317180Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pip install polars --upgrade -q","metadata":{"execution":{"iopub.status.busy":"2023-06-15T11:05:23.941795Z","iopub.execute_input":"2023-06-15T11:05:23.942945Z","iopub.status.idle":"2023-06-15T11:05:41.218650Z","shell.execute_reply.started":"2023-06-15T11:05:23.942900Z","shell.execute_reply":"2023-06-15T11:05:41.217488Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import polars as pl\n\npl.show_versions()","metadata":{"execution":{"iopub.status.busy":"2023-06-15T11:05:41.220971Z","iopub.execute_input":"2023-06-15T11:05:41.221304Z","iopub.status.idle":"2023-06-15T11:05:41.385857Z","shell.execute_reply.started":"2023-06-15T11:05:41.221273Z","shell.execute_reply":"2023-06-15T11:05:41.384898Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"customers = pl.read_csv(Path(input_dir) / \"customers.csv\")\nprint(customers.shape)\ncustomers.head()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-06-15T11:06:11.705837Z","iopub.execute_input":"2023-06-15T11:06:11.706539Z","iopub.status.idle":"2023-06-15T11:06:12.952387Z","shell.execute_reply.started":"2023-06-15T11:06:11.706499Z","shell.execute_reply":"2023-06-15T11:06:12.951628Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transactions = pl.read_csv(Path(input_dir) / \"transactions_train.csv\")\nprint(transactions.shape)\ntransactions.head()","metadata":{"execution":{"iopub.status.busy":"2023-06-15T11:06:12.954118Z","iopub.execute_input":"2023-06-15T11:06:12.954602Z","iopub.status.idle":"2023-06-15T11:06:31.096390Z","shell.execute_reply.started":"2023-06-15T11:06:12.954574Z","shell.execute_reply":"2023-06-15T11:06:31.095289Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"articles = pl.read_csv(Path(input_dir) / \"articles.csv\")\nprint(articles.shape)\narticles.head()","metadata":{"execution":{"iopub.status.busy":"2023-06-15T11:06:31.097954Z","iopub.execute_input":"2023-06-15T11:06:31.098263Z","iopub.status.idle":"2023-06-15T11:06:31.315474Z","shell.execute_reply.started":"2023-06-15T11:06:31.098238Z","shell.execute_reply":"2023-06-15T11:06:31.314688Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_submission = pl.read_csv(Path(input_dir) / \"sample_submission.csv\")\nprint(sample_submission.shape)\nsample_submission.head()","metadata":{"execution":{"iopub.status.busy":"2023-06-15T11:06:31.317642Z","iopub.execute_input":"2023-06-15T11:06:31.318178Z","iopub.status.idle":"2023-06-15T11:06:32.628477Z","shell.execute_reply.started":"2023-06-15T11:06:31.318149Z","shell.execute_reply":"2023-06-15T11:06:32.627313Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Split train/test by datetime","metadata":{}},{"cell_type":"code","source":"transactions = transactions.with_columns(\n    pl.col(\"t_dat\").str.to_datetime()\n).with_columns(\n    pl.col(\"t_dat\").dt.weekday().alias(\"weekday\")\n)\ntransactions.head()","metadata":{"execution":{"iopub.status.busy":"2023-06-15T11:06:32.630012Z","iopub.execute_input":"2023-06-15T11:06:32.631145Z","iopub.status.idle":"2023-06-15T11:07:00.108616Z","shell.execute_reply.started":"2023-06-15T11:06:32.631107Z","shell.execute_reply":"2023-06-15T11:07:00.106969Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"t_min, t_max = transactions[\"t_dat\"].min(), transactions[\"t_dat\"].max()\nduration = t_max - t_min\nduration","metadata":{"execution":{"iopub.status.busy":"2023-06-15T11:07:00.112107Z","iopub.execute_input":"2023-06-15T11:07:00.113095Z","iopub.status.idle":"2023-06-15T11:07:00.194790Z","shell.execute_reply.started":"2023-06-15T11:07:00.113055Z","shell.execute_reply":"2023-06-15T11:07:00.193649Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = transactions.filter(\n    pl.col(\"t_dat\") <= t_max - duration/3\n)\ntest = transactions.filter(\n    pl.col(\"t_dat\") > t_max - duration/3\n)\ntrain.shape, test.shape","metadata":{"execution":{"iopub.status.busy":"2023-06-15T11:07:00.196716Z","iopub.execute_input":"2023-06-15T11:07:00.197327Z","iopub.status.idle":"2023-06-15T11:07:07.714455Z","shell.execute_reply.started":"2023-06-15T11:07:00.197293Z","shell.execute_reply":"2023-06-15T11:07:07.713646Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Only make predictions on last month of data","metadata":{}},{"cell_type":"code","source":"last_month = t_max - duration/3 - pl.duration(weeks=4)\ntrain_last_month = train.filter(\n    pl.col(\"t_dat\") > last_month\n)\ntrain_last_month.shape","metadata":{"execution":{"iopub.status.busy":"2023-06-15T11:07:07.716093Z","iopub.execute_input":"2023-06-15T11:07:07.716492Z","iopub.status.idle":"2023-06-15T11:07:07.832864Z","shell.execute_reply.started":"2023-06-15T11:07:07.716449Z","shell.execute_reply":"2023-06-15T11:07:07.831536Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"last_month = t_max - pl.duration(weeks=4)\ntest_last_month = test.filter(\n    pl.col(\"t_dat\") > last_month\n)\ntest_last_month.shape","metadata":{"execution":{"iopub.status.busy":"2023-06-15T11:07:07.837065Z","iopub.execute_input":"2023-06-15T11:07:07.837402Z","iopub.status.idle":"2023-06-15T11:07:07.934496Z","shell.execute_reply.started":"2023-06-15T11:07:07.837374Z","shell.execute_reply":"2023-06-15T11:07:07.933694Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Quick preprocessing","metadata":{}},{"cell_type":"code","source":"def dummies_and_cast(df):\n    return pl.concat([\n        df[[\"t_dat\", \"customer_id\", \"article_id\", \"price\"]],\n        df[\"weekday\"].to_dummies(),\n        df[\"sales_channel_id\"].cast(pl.Utf8).to_frame(),\n    ], how=\"horizontal\")","metadata":{"execution":{"iopub.status.busy":"2023-06-15T11:07:07.935768Z","iopub.execute_input":"2023-06-15T11:07:07.936086Z","iopub.status.idle":"2023-06-15T11:07:07.942001Z","shell.execute_reply.started":"2023-06-15T11:07:07.936060Z","shell.execute_reply":"2023-06-15T11:07:07.940833Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_last_month = dummies_and_cast(train_last_month)\ntrain_last_month.head()","metadata":{"execution":{"iopub.status.busy":"2023-06-15T11:07:07.943481Z","iopub.execute_input":"2023-06-15T11:07:07.944642Z","iopub.status.idle":"2023-06-15T11:07:08.100025Z","shell.execute_reply.started":"2023-06-15T11:07:07.944602Z","shell.execute_reply":"2023-06-15T11:07:08.098937Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Loading skrub! (previously dirty_cat)","metadata":{}},{"cell_type":"code","source":"!pip install dirty_cat -q","metadata":{"execution":{"iopub.status.busy":"2023-06-15T11:07:08.101790Z","iopub.execute_input":"2023-06-15T11:07:08.102776Z","iopub.status.idle":"2023-06-15T11:07:23.095887Z","shell.execute_reply.started":"2023-06-15T11:07:08.102732Z","shell.execute_reply":"2023-06-15T11:07:23.094494Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from abc import abstractmethod\nfrom itertools import product\n\nimport numpy as np\nimport pandas as pd\ntry:\n    import polars as pl\n    POLARS_SETUP = True\nexcept ImportError:\n    POLARS_SETUP = False\n\nfrom sklearn.base import BaseEstimator, TransformerMixin\nfrom sklearn.utils.validation import check_is_fitted\n\nNUM_OPS = [\"sum\", \"mean\", \"std\", \"max\", \"min\"]\nCATEG_OPS = [\"mode\"]\nALL_OPS = NUM_OPS + CATEG_OPS\n\nPANDAS_OPS_MAPPING = {\n    \"mode\": pd.Series.mode\n}\n\n\ndef split_num_categ_ops(agg_ops):\n    \"\"\"Separate aggregagor operators input \n    by their type.\n\n    Parameters\n    ----------\n    agg_ops : list of str,\n        The input operators names.\n\n    Returns\n    -------\n    num_ops, categ_ops : Tuple of List of str\n        List of operator names\n    \"\"\"\n    num_ops, categ_ops = [], []\n    for op_name in agg_ops:\n        if op_name in NUM_OPS:\n            num_ops.append(op_name)\n        elif op_name in CATEG_OPS:\n            categ_ops.append(op_name)\n        else:\n            ValueError(\n                f\"'ops' options are {ALL_OPS}, got: {op_name}.\"\n            )\n    return num_ops, categ_ops\n\n\ndef dispatch_assembling_engine(tables):\n    \"\"\"Returns the AssemblingEngine implementation for given tables\n    module.\n\n    Parameters\n    ----------\n    tables : List of Tuple of (table, cols_to_join, cols_to_agg)\n        We only use table to detect the module used.\n    \n    Returns\n    -------\n    assembling_engine : AssemblingEngine\n        The suited AssemblingEngine implementation.\n    \"\"\"\n\n    use_pandas = all(\n        [isinstance(table, pd.DataFrame) for table, _, _ in tables]\n    )\n    use_polars = False\n    if POLARS_SETUP:\n        # we don't mix DataFrame and LazyFrame\n        use_polars = (\n            all(\n                [isinstance(table, pl.DataFrame) for table, _, _ in tables]\n            )\n            or all(\n                [isinstance(table, pl.LazyFrame) for table, _, _ in tables]\n            )\n        )\n    if use_pandas:\n        return PandasAssemblingEngine()\n    elif use_polars:\n        return PolarsAssemblingEngine()\n    else:\n        raise NotImplementedError(\n            \"Only Pandas or Polars DataFrame are currently supported.\"\n        )\n\n\nclass AssemblingEngine:\n    \"\"\"Helper class to perform the join and aggregate operations.\n\n    This is an abstract base class that is specialized depending\n    on the module of the dataframe used (Pandas or Polars).\n    \"\"\"\n    @abstractmethod\n    def agg(self, table, cols_to_join, cols_to_agg, agg_ops, suffix):\n        pass\n\n    @abstractmethod\n    def join(self, left, right, left_on, right_on):\n        pass\n\n\nclass PandasAssemblingEngine(AssemblingEngine):\n\n    def agg(self, table, cols_to_join, cols_to_agg, agg_ops, suffix):\n\n        def get_agg_ops(cols, agg_ops):\n            stats = {}\n            for col, op_name in product(cols, agg_ops):\n                op = PANDAS_OPS_MAPPING.get(op_name, op_name)\n                stats[f\"{col}_{op_name}\"] = pd.NamedAgg(col, op)\n            return stats\n\n        def split_num_categ_cols(table):\n            num_cols = table.select_dtypes(\"number\").columns\n            categ_cols = table.select_dtypes(\n                [\"object\", \"string\", \"category\"]\n            ).columns\n            return num_cols, categ_cols\n            \n        num_cols, categ_cols = split_num_categ_cols(\n            table[cols_to_agg]\n        )\n        num_ops, categ_ops = split_num_categ_ops(agg_ops)\n\n        num_stats = get_agg_ops(num_cols, num_ops)\n        categ_stats = get_agg_ops(categ_cols, categ_ops)\n                    \n        table = table.groupby(cols_to_join).agg(**num_stats, **categ_stats)\n        table.columns = [\n            f\"{col}{suffix}\" if col not in cols_to_join else col for col in table.columns\n        ]\n        \n        return table\n\n    def join(self, left, right, left_on, right_on):\n        \n        return left.merge(\n            right,\n            how=\"left\",\n            left_on=left_on,\n            right_on=right_on,\n        )\n\n\nclass PolarsAssemblingEngine(AssemblingEngine):\n\n    def agg(self, table, cols_to_join, cols_to_agg, agg_ops, suffix):\n\n        def split_num_categ_cols(table):\n            \n            num_cols = table.select(\n                pl.col(pl.NUMERIC_DTYPES)\n            ).columns\n            \n            categ_cols = table.select(\n                pl.col(pl.Utf8)\n            ).columns\n            \n            return num_cols, categ_cols\n\n        def get_agg_ops(cols, agg_ops):\n            stats, mode_cols = [], []\n            for col, op_name in product(cols, agg_ops):\n                op_dict = {\n                    \"mean\": pl.col(col).mean().alias(f\"{col}_{op_name}\"),\n                    \"std\": pl.col(col).std().alias(f\"{col}_{op_name}\"),\n                    \"sum\": pl.col(col).sum().alias(f\"{col}_{op_name}\"),\n                    \"min\": pl.col(col).min().alias(f\"{col}_{op_name}\"),\n                    \"max\": pl.col(col).max().alias(f\"{col}_{op_name}\"),\n                    \"mode\": pl.col(col).mode().alias(f\"{col}_{op_name}\"),\n                }\n                op = op_dict.get(op_name, None)\n                if op is None:\n                    raise ValueError(\n                        f\"Polars operation '{op}' is not supported. \"\n                        f\"Available: {list(op_dict)}\"\n                    )\n                stats.append(op)\n\n                # mode() output needs a flattening post-processing\n                if op_name == \"mode\":\n                    mode_cols.append(f\"{col}_mode\")\n                    \n            return stats, mode_cols\n\n        num_cols, categ_cols = split_num_categ_cols(\n            table.select(cols_to_agg)\n        )\n        \n        num_ops, categ_ops = split_num_categ_ops(agg_ops)\n\n        num_ops, num_mode_cols = get_agg_ops(num_cols, num_ops)\n        categ_ops, categ_mode_cols = get_agg_ops(categ_cols, categ_ops)\n\n        all_ops = [*num_ops, *categ_ops]\n        table = table.groupby(cols_to_join).agg(all_ops)\n\n        # flattening post-processing of mode() cols\n        flatten_ops = []\n        for col in [*num_mode_cols, *categ_mode_cols]:\n            flatten_ops.append(\n                # pl.col(col).arr.get(0).alias(col)\n                pl.col(col).list[0].alias(col)\n            )\n        table = table.with_columns(flatten_ops)\n        \n        cols_renaming = {\n            col: f\"{col}{suffix}\"\n            for col in table.columns if col not in cols_to_join\n        }\n        table = table.rename(cols_renaming)\n        \n        return table\n\n    def join(self, left, right, left_on, right_on):\n        \n        return left.join(\n            right,\n            how=\"left\",\n            left_on=left_on,\n            right_on=right_on,\n        )\n    \n\nclass JoinAggregator(BaseEstimator, TransformerMixin):\n    \"\"\"Perform aggregation on auxilliary dataframes before joining\n    on the base dataframe.\n\n    Apply numerical (mean, std, min, max) and categorical (mode) aggregation \n    operations on the columns to agg, selected by dtypes.\n    \n    The grouping columns used during the aggregation are the columns used \n    as keys for joining.\n\n    These operations can run lazily by inputing polars LazyFrames. \n    Pandas and polars dataframes can't be mixed together, so the user has \n    to switch between format if needed.\n    \n    Parameters\n    ----------\n    tables : list of tuples\n        List of (dataframe, columns_to_join, columns_to_agg) tuple\n        specifying the auxilliary dataframes and their columns for joining \n        and aggregation operations.\n    \n        dataframe : {pandas.DataFrame, polars.DataFrame, polars.LazyFrame}\n            The auxilliary data to aggregate and join.\n        \n        columns_to_join : str or array-like\n            Select the columns from the dataframe to use as keys during the join operation.\n        \n        columns_to_agg : str or array-like\n            Select the columns from the dataframe to use as values during \n            the aggregation operations.\n\n    main_key : str or array-like\n        Select the columns from the base table to use as keys during \n        the join operation.\n\n    agg_ops : str or list of str, default=None\n        Aggregation operations to perform on the auxilliary table.\n        Options: {'mean', 'std', 'min', 'max', 'mode'}. If set to None, \n        ['mean', 'mode'] will be used.\n    \"\"\"\n\n    def __init__(self, tables, main_key, agg_ops=None, suffixes=None):\n        self.tables = tables\n        self.main_key = main_key\n        self.agg_ops = agg_ops\n        self.suffixes = suffixes\n\n    def fit(self, X, y=None):\n        \"\"\"Fit the instance to the auxiliary tables by aggregating them\n        and storing the outputs.\n\n        Parameters\n        ----------\n        X : {pandas.Dataframe, polars.DataFrame, polars.LazyFrame}\n            Input data, based table on which to left join the \n            auxilliary tables..\n        \n        y : array-like of shape (n_samples), default=None\n            Used to compute the correlation between the generated covariates\n            and the target for screening purposes.\n\n        Returns\n        -------\n        :obj:`JoinAggregator`\n            Fitted :class:`JoinAggregator` instance (self).\n        \"\"\"\n        self.check_cols(X)\n\n        if self.agg_ops is None:\n            agg_ops = [\"mean\", \"mode\"]\n        else:\n            agg_ops = np.atleast_1d(self.agg_ops).tolist()\n\n        self.assembly_engine = dispatch_assembling_engine(self.tables)\n        \n        self.agg_tables_ = []\n        for (table, cols_to_join, cols_to_agg), suffix in zip(self.tables, self.suffixes_):\n\n            agg_table = self.assembly_engine.agg(\n                table,\n                cols_to_join,\n                cols_to_agg,\n                agg_ops,\n                suffix,\n            )\n            agg_table = self._screen(agg_table, y)\n\n            self.agg_tables_.append((agg_table, cols_to_join))\n            \n        return self\n\n    def transform(self, X):\n        \"\"\"Transform `X` by left joining the pre-aggregated \n        auxiliary tables to it.\n\n        Parameters\n        ----------\n        X : {pandas.DataFrame, polars.DataFrame, polars.LazyFrame}\n            The input data to transform.\n        \"\"\"\n\n        check_is_fitted(self, \"agg_tables_\")\n\n        for main_key, (aux, aux_key) in zip(self.main_keys_, self.agg_tables_):\n            X = self.assembly_engine.join(\n                left=X,\n                right=aux,\n                left_on=main_key,\n                right_on=aux_key,\n            )\n\n        return X\n\n    def _screen(self, agg_table, y):\n        # TODO: Add logic\n        return agg_table\n    \n    def check_cols(self, X):\n        \"\"\"Check that all columns to join and columns to aggregate \n        belong to their respective dataframes.\n        \"\"\"\n        # Check main_keys\n        main_keys = np.atleast_1d(self.main_key).tolist()\n        missing_cols = set(main_keys) - set(X.columns)\n        if len(missing_cols) > 0:\n            raise ValueError(\n                f\"Got main_key={self.main_key!r}, but column not in {list(X.columns)}.\"\n            )\n            \n        n_main_keys, n_tables = len(main_keys), len(self.tables)\n        if (n_main_keys != 1) and (n_main_keys != n_tables):\n            raise ValueError(\n                \"The number of main keys must be either 1 or \"\n                \"match the number of tables\"\n            )\n        \n        # Ensure n_main_keys == n_tables\n        if n_main_keys == 1:\n            main_keys = main_keys * n_tables\n        \n        self.main_keys_ = main_keys\n        \n        # Check agg and join columns\n        for idx, (table, cols_to_join, cols_to_agg) in enumerate(self.tables, start=1):\n\n            cols_to_join = np.atleast_1d(cols_to_join).tolist()\n            cols_to_agg = np.atleast_1d(cols_to_agg).tolist()\n\n            table_cols = set(table.columns)\n            input_cols = set([*cols_to_join, *cols_to_agg])\n\n            missing_cols = input_cols - table_cols\n            if len(missing_cols) > 0:\n                raise ValueError(f\"{missing_cols} are missing in table {idx}\")\n\n        # Check suffixes\n        if self.suffixes is None:\n            if n_tables == 1:\n                suffixes = [\"\"]\n            else:\n                suffixes = [f\"_{idx}\" for idx in range(len(self.tables))]\n        elif hasattr(self.suffixes, \"__len__\"):\n            suffixes = np.atleast_1d(self.suffixes).tolist()\n            if len(suffixes) != n_tables:\n                raise ValueError(\"Suffixes must be None or match the number of tables.\")\n        else:\n            raise ValueError(\n                \"Suffixes must be a list of string matching the number of tables.\"\n            )\n        \n        self.suffixes_ = suffixes\n        \n        return","metadata":{"execution":{"iopub.status.busy":"2023-06-15T11:08:07.358427Z","iopub.execute_input":"2023-06-15T11:08:07.358842Z","iopub.status.idle":"2023-06-15T11:08:08.309922Z","shell.execute_reply.started":"2023-06-15T11:08:07.358809Z","shell.execute_reply":"2023-06-15T11:08:08.308817Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Join aggregate on `customer_id` and `article_id` from the same table!","metadata":{}},{"cell_type":"code","source":"cols_weekday = [f\"weekday_{idx}\" for idx in range(1, 8)]\ncols_to_agg = [\"price\", \"sales_channel_id\", *cols_weekday]\n\njoin_agg = JoinAggregator(\n    tables=[\n        (train_last_month.lazy(), [\"customer_id\"], cols_to_agg),\n        (train_last_month.lazy(), [\"article_id\"], cols_to_agg)\n    ],\n    suffixes=[\"_customer\", \"_article\"],\n    main_key=[\"customer_id\", \"article_id\"],\n    agg_ops=[\"sum\", \"mode\"],\n)\ntrain_last_month = join_agg.fit_transform(train_last_month.lazy())","metadata":{"execution":{"iopub.status.busy":"2023-06-15T11:08:15.074175Z","iopub.execute_input":"2023-06-15T11:08:15.074564Z","iopub.status.idle":"2023-06-15T11:08:15.095112Z","shell.execute_reply.started":"2023-06-15T11:08:15.074533Z","shell.execute_reply":"2023-06-15T11:08:15.094194Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_last_month = train_last_month.select(\n    pl.exclude(cols_weekday)\n).collect()\n\nprint(\n    train_last_month.select(pl.col(\"customer_id\").count()).item(),\n    len(train_last_month.columns)\n)\ntrain_last_month.head()","metadata":{"execution":{"iopub.status.busy":"2023-06-15T11:08:17.740224Z","iopub.execute_input":"2023-06-15T11:08:17.740623Z","iopub.status.idle":"2023-06-15T11:08:19.289315Z","shell.execute_reply.started":"2023-06-15T11:08:17.740588Z","shell.execute_reply":"2023-06-15T11:08:19.288394Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Get embeddings from articles","metadata":{}},{"cell_type":"code","source":"from dirty_cat import TableVectorizer, MinHashEncoder\n\nn_components = 10\n\ntv = TableVectorizer(\n    high_card_cat_transformer=MinHashEncoder(\n        n_components=n_components\n    ),\n)\n\ncols = [\"product_type_name\", \"detail_desc\"]\narticles_t = tv.fit_transform(articles[cols])\n\narticles_t = pl.DataFrame(articles_t)\ncols_embedding = [f\"x{idx}\" for idx in range(len(cols) * n_components)]\narticles_t.columns = cols_embedding\narticles_t = pl.concat(\n    [articles[[\"article_id\"]], articles_t], how=\"horizontal\"\n)\narticles_t.head()","metadata":{"execution":{"iopub.status.busy":"2023-06-15T11:23:48.993428Z","iopub.execute_input":"2023-06-15T11:23:48.993854Z","iopub.status.idle":"2023-06-15T11:24:01.034551Z","shell.execute_reply.started":"2023-06-15T11:23:48.993822Z","shell.execute_reply":"2023-06-15T11:24:01.033464Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Let's join aggregate again, on the embeddings of our customers and articles!","metadata":{}},{"cell_type":"markdown","source":"First, join the article embeddings to a `customer_id | article_id` table.","metadata":{}},{"cell_type":"code","source":"customers_t = (\n    train_last_month.select([\"customer_id\", \"article_id\"]).lazy()\n    .join(\n        articles_t.lazy(),\n        left_on=\"article_id\",\n        right_on=\"article_id\",\n        how=\"left\",\n    )\n    .select(pl.exclude(\"article_id\"))\n).collect()\ncustomers_t.head()","metadata":{"execution":{"iopub.status.busy":"2023-06-15T11:35:28.330580Z","iopub.execute_input":"2023-06-15T11:35:28.330994Z","iopub.status.idle":"2023-06-15T11:35:28.551730Z","shell.execute_reply.started":"2023-06-15T11:35:28.330953Z","shell.execute_reply":"2023-06-15T11:35:28.550940Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Now:\n- Compute the mean of embeddings over `customer_id`, then join on `customer_id`.\n- Compute the mean of embeddings over `article_id` (themselves, the mean is a no-op here), then join on `article_id`. ","metadata":{"execution":{"iopub.status.busy":"2023-06-15T11:15:26.365899Z","iopub.execute_input":"2023-06-15T11:15:26.366311Z","iopub.status.idle":"2023-06-15T11:15:26.374850Z","shell.execute_reply.started":"2023-06-15T11:15:26.366282Z","shell.execute_reply":"2023-06-15T11:15:26.373124Z"}}},{"cell_type":"code","source":"cols_to_agg = cols_embedding\n\njoin_agg = JoinAggregator(\n    tables=[\n        (customers_t.lazy(), [\"customer_id\"], cols_to_agg),\n        (articles_t.lazy(), [\"article_id\"], cols_to_agg)\n    ],\n    suffixes=[\"_customer\", \"_article\"],\n    main_key=[\"customer_id\", \"article_id\"],\n    agg_ops=[\"mean\"],\n)\ntrain_last_month = join_agg.fit_transform(train_last_month.lazy())","metadata":{"execution":{"iopub.status.busy":"2023-06-15T11:36:07.026570Z","iopub.execute_input":"2023-06-15T11:36:07.026971Z","iopub.status.idle":"2023-06-15T11:36:07.037218Z","shell.execute_reply.started":"2023-06-15T11:36:07.026942Z","shell.execute_reply":"2023-06-15T11:36:07.035769Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_last_month = train_last_month.collect()\n\nprint(\n    train_last_month.select(pl.col(\"customer_id\").count()).item(),\n    len(train_last_month.columns)\n)\ntrain_last_month.head()","metadata":{"execution":{"iopub.status.busy":"2023-06-15T11:36:09.211938Z","iopub.execute_input":"2023-06-15T11:36:09.212311Z","iopub.status.idle":"2023-06-15T11:36:10.002754Z","shell.execute_reply.started":"2023-06-15T11:36:09.212284Z","shell.execute_reply":"2023-06-15T11:36:10.001926Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# What's next?\n\n1. Negative Sampling 1:30 between customers and articles to create negative labels\n2. Fit using a `sklearn.ensemble.HistGradientBoostingClassifier`\n3. Predict: \n    - run join aggregator for customers and articles for the predict month\n    - for each customer, fetch the top 50 embeddings from articles\n    - join the customer embeddings and the article embeddings\n    - predict using the fitted classifier\n4. Score with MAP@K","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}