{"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":"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\n\n# Input data files are available in the read-only \"../input/\" directory\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-08-11T07:46:00.107637Z","iopub.execute_input":"2022-08-11T07:46:00.108107Z","iopub.status.idle":"2022-08-11T07:46:00.130051Z","shell.execute_reply.started":"2022-08-11T07:46:00.108017Z","shell.execute_reply":"2022-08-11T07:46:00.129070Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# pytorch-lifestream (ptls)\n\nThis is library for sequential data: https://github.com/dllllb/pytorch-lifestream\n\nData preprocessed with https://www.kaggle.com/code/ivkireev/amex-ptls-data-preprocessing\n","metadata":{}},{"cell_type":"code","source":"!pip uninstall tensorflow-transform tfx-bsl -y\n!pip install pyspark\n!pip install git+https://github.com/dllllb/pytorch-lifestream.git@main\n# !pip install pyarrow==5.0.0  # uncomment it for gpu kernel","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-08-11T07:46:01.754566Z","iopub.execute_input":"2022-08-11T07:46:01.754985Z","iopub.status.idle":"2022-08-11T07:47:53.216566Z","shell.execute_reply.started":"2022-08-11T07:46:01.754953Z","shell.execute_reply":"2022-08-11T07:47:53.213979Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Imports","metadata":{}},{"cell_type":"code","source":"import pickle\nfrom copy import deepcopy\nfrom functools import partial\n\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport pandas as pd\nimport pytorch_lightning as pl\n\nimport torch\nimport torchmetrics\n\nfrom sklearn.model_selection import train_test_split\n\nfrom ptls.data_load.datasets import ParquetDataset, parquet_file_scan\nfrom ptls.data_load.iterable_processing import IterableShuffle\nfrom ptls.data_load.padded_batch import PaddedBatch\nfrom ptls.data_load.utils import collate_feature_dict\nfrom ptls.frames import PtlsDataModule\nfrom ptls.frames.inference_module import InferenceModule\nfrom ptls.frames.supervised import SeqToTargetIterableDataset, SequenceToTarget\nfrom ptls.nn import TrxEncoder, RnnEncoder, TransformerEncoder, PBLinear, PBL2Norm, PBLayerNorm, PBReLU\n","metadata":{"execution":{"iopub.status.busy":"2022-08-11T07:57:30.199257Z","iopub.execute_input":"2022-08-11T07:57:30.199813Z","iopub.status.idle":"2022-08-11T07:57:30.210481Z","shell.execute_reply.started":"2022-08-11T07:57:30.199775Z","shell.execute_reply":"2022-08-11T07:57:30.209214Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Read the data\n\nData was prepared in `ptls` format. This is parquet files where each row represent one data sample - one client. Features are the sequential fields with arrays of length `sequence_length`\n\nDemo of data usage is in https://www.kaggle.com/code/ivkireev/ptls-data-usage","metadata":{}},{"cell_type":"code","source":"# Read preprocessor which keep column names and dictionary sizes for categorical fields\nwith open('../input/amex-ptls-data/preprocessor.pickle', 'rb') as f:\n    preprocessor = pickle.load(f)","metadata":{"execution":{"iopub.status.busy":"2022-08-11T07:50:11.299002Z","iopub.execute_input":"2022-08-11T07:50:11.299852Z","iopub.status.idle":"2022-08-11T07:50:11.385358Z","shell.execute_reply.started":"2022-08-11T07:50:11.299810Z","shell.execute_reply":"2022-08-11T07:50:11.384147Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preprocessor.get_category_dictionary_sizes()","metadata":{"execution":{"iopub.status.busy":"2022-08-11T07:50:12.619696Z","iopub.execute_input":"2022-08-11T07:50:12.620132Z","iopub.status.idle":"2022-08-11T07:50:12.630304Z","shell.execute_reply.started":"2022-08-11T07:50:12.620095Z","shell.execute_reply":"2022-08-11T07:50:12.629130Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"all_train_files = parquet_file_scan('../input/amex-ptls-data/data_train.parquet')\ntest_files = parquet_file_scan('../input/amex-ptls-data/data_test.parquet')\n\nlen(all_train_files), len(test_files)","metadata":{"execution":{"iopub.status.busy":"2022-08-11T07:50:15.246666Z","iopub.execute_input":"2022-08-11T07:50:15.247466Z","iopub.status.idle":"2022-08-11T07:50:15.598303Z","shell.execute_reply.started":"2022-08-11T07:50:15.247425Z","shell.execute_reply":"2022-08-11T07:50:15.597399Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_files, valid_files = train_test_split(all_train_files, test_size=0.2, random_state=242)\nlen(train_files), len(valid_files)","metadata":{"execution":{"iopub.status.busy":"2022-08-11T07:50:15.600183Z","iopub.execute_input":"2022-08-11T07:50:15.601053Z","iopub.status.idle":"2022-08-11T07:50:15.612844Z","shell.execute_reply.started":"2022-08-11T07:50:15.601015Z","shell.execute_reply":"2022-08-11T07:50:15.611055Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class NumConcat:\n    \"\"\"Concatenete numerical columns which a collected to single field\n    Work with one tensor is faster than iterating over 150+ numerical columns in amex dataset.\n    \"\"\"\n    def __call__(self, data):\n        self.data = data\n        return iter(self)\n        \n    def __iter__(self):\n        for rec in self.data:\n            rec['num_cols'] = np.stack(rec['num_cols'].tolist(), axis=1)\n            rec['num_cols'] = torch.from_numpy(rec['num_cols'])\n            yield rec\n","metadata":{"execution":{"iopub.status.busy":"2022-08-11T07:50:19.599273Z","iopub.execute_input":"2022-08-11T07:50:19.600510Z","iopub.status.idle":"2022-08-11T07:50:19.608015Z","shell.execute_reply.started":"2022-08-11T07:50:19.600458Z","shell.execute_reply":"2022-08-11T07:50:19.607115Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# We will use ParquetDataset in iterable mode\nall_train_dataset = ParquetDataset(all_train_files, shuffle_files=True, i_filters=[\n    NumConcat(),\n    IterableShuffle(buffer_size=1000),  # this shuffle rows with buffer \n])\ntrain_dataset = ParquetDataset(train_files, shuffle_files=True, i_filters=[\n    NumConcat(),\n    IterableShuffle(buffer_size=1000),\n])\nvalid_dataset = ParquetDataset(valid_files, i_filters=[NumConcat()])\ntest_dataset = ParquetDataset(test_files, i_filters=[NumConcat()])","metadata":{"execution":{"iopub.status.busy":"2022-08-11T07:50:21.686388Z","iopub.execute_input":"2022-08-11T07:50:21.687548Z","iopub.status.idle":"2022-08-11T07:50:21.693489Z","shell.execute_reply.started":"2022-08-11T07:50:21.687501Z","shell.execute_reply":"2022-08-11T07:50:21.692488Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Encoder network","metadata":{}},{"cell_type":"code","source":"class NaImputer(torch.nn.Module):\n    \"\"\"Replace NaN with 0 value\n    \"\"\"\n    def forward(self, x: PaddedBatch):\n        new_x = {k: self.impute(k, v) for k, v in x.payload.items()}\n        return PaddedBatch(new_x, x.seq_lens)\n    \n    def impute(self, k, v):\n        if not PaddedBatch.is_seq_feature(k, v):\n            return v\n        return v.masked_fill(torch.isnan(v), 0.0)","metadata":{"execution":{"iopub.status.busy":"2022-08-11T07:50:52.238984Z","iopub.execute_input":"2022-08-11T07:50:52.239412Z","iopub.status.idle":"2022-08-11T07:50:52.248212Z","shell.execute_reply.started":"2022-08-11T07:50:52.239380Z","shell.execute_reply":"2022-08-11T07:50:52.246783Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def sign_log(x):\n    return torch.log1p(torch.abs(x)) * torch.sign(x)\n    \nclass TrxEncoderExt(torch.nn.Module):\n    \"\"\"Supports numerical features stacked with `NumConcat`.\n    Implements the same behavior as ptls.nn.TrxEncoder but faster.\n    \"\"\"\n    def __init__(self, trx_encoder, numeric_count=177):\n        super().__init__()\n        \n        self.trx_encoder = trx_encoder\n        self.numeric_count = numeric_count\n        self.bn = torch.nn.BatchNorm1d(numeric_count)\n        \n    def forward(self, x):\n        x_categorical = self.trx_encoder(x)\n        \n        num_cols = x.payload['num_cols']\n        B, T, H = num_cols.size()\n        \n        num_log = (self.bn(sign_log(num_cols).view(B * T, H)).view(B, T, H) *\n            x.seq_len_mask.float().unsqueeze(2)).clamp(-3, 3)\n        return PaddedBatch(\n            torch.cat([\n                x_categorical.payload, \n                num_log,\n            ], dim=2),\n            x.seq_lens,\n        )\n    \n    @property\n    def output_size(self):\n        return self.trx_encoder.output_size + self.numeric_count\n","metadata":{"execution":{"iopub.status.busy":"2022-08-11T07:54:52.208220Z","iopub.execute_input":"2022-08-11T07:54:52.208700Z","iopub.status.idle":"2022-08-11T07:54:52.221235Z","shell.execute_reply.started":"2022-08-11T07:54:52.208665Z","shell.execute_reply":"2022-08-11T07:54:52.220296Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Metrics\n\nTaken from https://www.kaggle.com/code/inversion/amex-competition-metric-python\nand covered to `torchmetrics.Metric` for NN usage","metadata":{}},{"cell_type":"code","source":"def top_four_percent_captured(y_true: pd.DataFrame, y_pred: pd.DataFrame) -> float:\n    df = (pd.concat([y_true, y_pred], axis='columns')\n          .sort_values('prediction', ascending=False))\n    df['weight'] = df['target'].apply(lambda x: 20 if x==0 else 1)\n    four_pct_cutoff = int(0.04 * df['weight'].sum())\n    df['weight_cumsum'] = df['weight'].cumsum()\n    df_cutoff = df.loc[df['weight_cumsum'] <= four_pct_cutoff]\n    return (df_cutoff['target'] == 1).sum() / (df['target'] == 1).sum()\n\ndef weighted_gini(y_true: pd.DataFrame, y_pred: pd.DataFrame) -> float:\n    df = (pd.concat([y_true, y_pred], axis='columns')\n          .sort_values('prediction', ascending=False))\n    df['weight'] = df['target'].apply(lambda x: 20 if x==0 else 1)\n    df['random'] = (df['weight'] / df['weight'].sum()).cumsum()\n    total_pos = (df['target'] * df['weight']).sum()\n    df['cum_pos_found'] = (df['target'] * df['weight']).cumsum()\n    df['lorentz'] = df['cum_pos_found'] / total_pos\n    df['gini'] = (df['lorentz'] - df['random']) * df['weight']\n    return df['gini'].sum()\n\ndef normalized_weighted_gini(y_true: pd.DataFrame, y_pred: pd.DataFrame) -> float:\n    y_true_pred = y_true.rename(columns={'target': 'prediction'})\n    return weighted_gini(y_true, y_pred) / weighted_gini(y_true, y_true_pred)\n\ndef amex_metric(y_true: pd.DataFrame, y_pred: pd.DataFrame) -> float:\n    g = normalized_weighted_gini(y_true, y_pred)\n    d = top_four_percent_captured(y_true, y_pred)\n\n    return 0.5 * (g + d), g, d","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-08-11T07:58:46.419929Z","iopub.execute_input":"2022-08-11T07:58:46.420430Z","iopub.status.idle":"2022-08-11T07:58:46.436385Z","shell.execute_reply.started":"2022-08-11T07:58:46.420395Z","shell.execute_reply":"2022-08-11T07:58:46.435444Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TopFourPercentCaptured(torchmetrics.Metric):\n    full_state_update = True\n    compute_on_cpu = True\n    \n    def __init__(self):\n        super().__init__()\n    \n        self.add_state('y_pred', []),\n        self.add_state('y_true', []),\n\n    def update(self, y_pred, y_true):\n        self.y_pred.append(y_pred.detach().cpu())\n        self.y_true.append(y_true.detach().cpu())\n        \n    def compute(self):\n        y_pred = pd.DataFrame({'prediction': torch.cat(self.y_pred).numpy()})\n        y_true = pd.DataFrame({'target': torch.cat(self.y_true).numpy()})\n        return top_four_percent_captured(y_true, y_pred)\n    \n    \nclass NormalizedWeightedGini(torchmetrics.Metric):\n    full_state_update = True\n    compute_on_cpu = True\n    \n    def __init__(self):\n        super().__init__()\n    \n        self.add_state('y_pred', []),\n        self.add_state('y_true', []),\n\n    def update(self, y_pred, y_true):\n        self.y_pred.append(y_pred.detach().cpu())\n        self.y_true.append(y_true.detach().cpu())\n        \n    def compute(self):\n        y_pred = pd.DataFrame({'prediction': torch.cat(self.y_pred).numpy()})\n        y_true = pd.DataFrame({'target': torch.cat(self.y_true).numpy()})\n        return normalized_weighted_gini(y_true, y_pred)\n    \n    \nclass AmexMetric(torchmetrics.Metric):\n    full_state_update = True\n    compute_on_cpu = True\n    \n    def __init__(self):\n        super().__init__()\n    \n        self.add_state('y_pred', []),\n        self.add_state('y_true', []),\n\n    def update(self, y_pred, y_true):\n        self.y_pred.append(y_pred.detach().cpu())\n        self.y_true.append(y_true.detach().cpu())\n        \n    def compute(self):\n        y_pred = pd.DataFrame({'prediction': torch.cat(self.y_pred).numpy()})\n        y_true = pd.DataFrame({'target': torch.cat(self.y_true).numpy()})\n        return amex_metric(y_true, y_pred)[0]","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-08-11T07:58:47.458723Z","iopub.execute_input":"2022-08-11T07:58:47.459122Z","iopub.status.idle":"2022-08-11T07:58:47.474777Z","shell.execute_reply.started":"2022-08-11T07:58:47.459092Z","shell.execute_reply":"2022-08-11T07:58:47.473690Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loss","metadata":{}},{"cell_type":"code","source":"class PairwiseMarginRankingLoss(torch.nn.Module):\n    \"\"\"Ranking loss for amex task\n    \"\"\"\n    def __init__(self, margin):\n        super().__init__()\n\n        self.margin = margin\n\n    def forward(self, pred, true):\n        pred_0 = pred[true == 0]\n        pred_1 = pred[true == 1]\n        loss = pred_0.view(1, -1) - pred_1.view(-1, 1) + self.margin\n        loss = torch.nn.functional.relu(loss)\n        return loss.mean()","metadata":{"execution":{"iopub.status.busy":"2022-08-11T07:58:48.266120Z","iopub.execute_input":"2022-08-11T07:58:48.266560Z","iopub.status.idle":"2022-08-11T07:58:48.277534Z","shell.execute_reply.started":"2022-08-11T07:58:48.266525Z","shell.execute_reply":"2022-08-11T07:58:48.276149Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Lightning Module\n\nhttps://pytorch-lightning.readthedocs.io/en/1.6.4/common/lightning_module.html","metadata":{}},{"cell_type":"code","source":"class SequenceToTargetExt(SequenceToTarget):\n    \"\"\"LightningModule for supervised task or for inference\n    \"\"\"\n    def configure_optimizers(self):\n        parameters = [\n            {\n                'params': [p for n, p in self.named_parameters() \n                           if n.startswith('seq_encoder.1.')],\n                 'weight_decay': 1e-4,\n            },\n            {\n                'params': [p for n, p in self.named_parameters()\n                           if not n.startswith('seq_encoder.1.')]\n            },\n        ]\n        optimizer = self.optimizer_partial(parameters)\n        scheduler = self.lr_scheduler_partial(optimizer)\n        return [optimizer], [scheduler]\n","metadata":{"execution":{"iopub.status.busy":"2022-08-11T08:00:02.379160Z","iopub.execute_input":"2022-08-11T08:00:02.379656Z","iopub.status.idle":"2022-08-11T08:00:02.388361Z","shell.execute_reply.started":"2022-08-11T08:00:02.379615Z","shell.execute_reply":"2022-08-11T08:00:02.387179Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train\n\nPrepare `datamodule`, model as `pl.LightningModile` and `pl.Trainer`\n\nRun train with validation","metadata":{}},{"cell_type":"code","source":"def get_model():\n    trx_encoder = TrxEncoderExt(TrxEncoder(\n        embeddings={col: {'in': d, 'out': 4} for col, d in preprocessor.get_category_dictionary_sizes().items()},\n        numeric_values={},  # processed by TrxEncoderExt\n    #     numeric_values={col: 'identity' for col in num_fields},\n    ))\n\n    seq_encoder = torch.nn.Sequential(\n        NaImputer(),\n        trx_encoder,\n        RnnEncoder(\n            input_size=trx_encoder.output_size, \n            hidden_size=128,\n            is_reduce_sequence=True,\n        ),\n    )\n    \n    model = SequenceToTargetExt(\n        seq_encoder=seq_encoder,\n        head=torch.nn.Sequential(\n            torch.nn.Linear(seq_encoder[-1].embedding_size, 1),\n            torch.nn.Sigmoid(),\n            torch.nn.Flatten(start_dim=0),\n        ),\n        loss=PairwiseMarginRankingLoss(0.1),\n        metric_list={\n    #         'auc': torchmetrics.AUROC(),\n    #         'top4': TopFourPercentCaptured(),\n    #         'nwg': NormalizedWeightedGini(),\n            'amex': AmexMetric(),\n        },\n        optimizer_partial=partial(\n            torch.optim.Adam,\n            lr=0.001,\n            weight_decay=0,\n        ),\n        lr_scheduler_partial=partial(\n            torch.optim.lr_scheduler.StepLR,\n            step_size=1,\n            gamma=0.1 ** (1/6),\n        ),\n    )\n    for n, p in model.named_parameters():\n        print(n, p.shape)\n\n    return model","metadata":{"execution":{"iopub.status.busy":"2022-08-11T08:06:16.219733Z","iopub.execute_input":"2022-08-11T08:06:16.220247Z","iopub.status.idle":"2022-08-11T08:06:16.231693Z","shell.execute_reply.started":"2022-08-11T08:06:16.220201Z","shell.execute_reply":"2022-08-11T08:06:16.230365Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dm = PtlsDataModule(\n    train_data=SeqToTargetIterableDataset(train_dataset, target_col_name='target'),\n    valid_data=SeqToTargetIterableDataset(valid_dataset, target_col_name='target'),\n    train_batch_size=512,\n    valid_batch_size=1024,\n    train_num_workers=8,\n    valid_num_workers=8,\n)\n\nmodel = get_model()\n\ntrainer = pl.Trainer(\n    gpus=0,\n    max_epochs=6,\n    enable_checkpointing=False,\n    enable_progress_bar=False,\n    callbacks=[\n        pl.callbacks.LearningRateMonitor(),\n    ],\n)","metadata":{"execution":{"iopub.status.busy":"2022-08-05T14:24:48.497998Z","iopub.execute_input":"2022-08-05T14:24:48.499137Z","iopub.status.idle":"2022-08-05T14:24:48.584590Z","shell.execute_reply.started":"2022-08-05T14:24:48.499020Z","shell.execute_reply":"2022-08-05T14:24:48.583528Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('\\n'.join([\n    'Running...',\n    f'version = {trainer.logger.version}',\n]))\n\ntrainer.fit(model, dm)\n\nprint(trainer.logged_metrics)","metadata":{"execution":{"iopub.status.busy":"2022-08-05T14:25:08.810222Z","iopub.execute_input":"2022-08-05T14:25:08.810652Z","iopub.status.idle":"2022-08-05T14:26:48.712998Z","shell.execute_reply.started":"2022-08-05T14:25:08.810619Z","shell.execute_reply":"2022-08-05T14:26:48.706540Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Results","metadata":{}},{"cell_type":"code","source":"import ptls.tb_interface as tb","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_tb = tb.get_scalars('lightning_logs/')","metadata":{"execution":{"iopub.status.busy":"2022-08-03T12:02:35.850003Z","iopub.execute_input":"2022-08-03T12:02:35.851003Z","iopub.status.idle":"2022-08-03T12:02:35.871219Z","shell.execute_reply.started":"2022-08-03T12:02:35.850957Z","shell.execute_reply":"2022-08-03T12:02:35.869790Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_tb['tag'].value_counts()","metadata":{"execution":{"iopub.status.busy":"2022-08-03T12:02:36.505676Z","iopub.execute_input":"2022-08-03T12:02:36.507217Z","iopub.status.idle":"2022-08-03T12:02:36.518451Z","shell.execute_reply.started":"2022-08-03T12:02:36.507150Z","shell.execute_reply":"2022-08-03T12:02:36.516951Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for tag in ['val_amex', 'train_amex']:\n    df = df_tb[lambda x: x['tag'] == tag].set_index(['step', 'version'])['value'].unstack()\n    df.plot(figsize=(8, 4), title=tag)\n    plt.show()\n    print(df.max())","metadata":{"execution":{"iopub.status.busy":"2022-08-03T12:03:06.782936Z","iopub.execute_input":"2022-08-03T12:03:06.784142Z","iopub.status.idle":"2022-08-03T12:03:07.333996Z","shell.execute_reply.started":"2022-08-03T12:03:06.784074Z","shell.execute_reply":"2022-08-03T12:03:07.332653Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Final Train\n\nTrain on full frain data","metadata":{}},{"cell_type":"code","source":"dm = PtlsDataModule(\n    train_data=SeqToTargetIterableDataset(all_train_dataset, target_col_name='target'),\n    train_batch_size=512,\n    train_num_workers=4,\n)\n\nmodel = get_model()\n\ntrainer = pl.Trainer(\n    gpus=0,\n    max_epochs=6,\n    enable_checkpointing=False,\n    enable_progress_bar=False,\n    callbacks=[\n        pl.callbacks.LearningRateMonitor(),\n    ],\n)\n\nprint('\\n'.join([\n    'Running train on full data...',\n    f'version = {trainer.logger.version}',\n]))\n\n\ntrainer.fit(model, dm)\n\nprint(trainer.logged_metrics)","metadata":{"execution":{"iopub.status.busy":"2022-08-03T12:05:58.610535Z","iopub.execute_input":"2022-08-03T12:05:58.610964Z","iopub.status.idle":"2022-08-03T12:26:17.269323Z","shell.execute_reply.started":"2022-08-03T12:05:58.610929Z","shell.execute_reply":"2022-08-03T12:26:17.267653Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference on test data","metadata":{}},{"cell_type":"code","source":"test_dataloader = torch.utils.data.DataLoader(\n    test_dataset,\n    collate_fn=collate_feature_dict,\n    num_workers=4,\n    batch_size=1024,\n    \n)","metadata":{"execution":{"iopub.status.busy":"2022-08-03T12:28:06.756938Z","iopub.execute_input":"2022-08-03T12:28:06.757457Z","iopub.status.idle":"2022-08-03T12:28:06.764338Z","shell.execute_reply.started":"2022-08-03T12:28:06.757413Z","shell.execute_reply":"2022-08-03T12:28:06.763143Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_predict = trainer.predict(InferenceModule(model, model_out_name='prediction'), test_dataloader)","metadata":{"execution":{"iopub.status.busy":"2022-08-03T12:34:14.651917Z","iopub.execute_input":"2022-08-03T12:34:14.652404Z","iopub.status.idle":"2022-08-03T12:38:33.096526Z","shell.execute_reply.started":"2022-08-03T12:34:14.652367Z","shell.execute_reply":"2022-08-03T12:38:33.095470Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_predict[0].head()","metadata":{"execution":{"iopub.status.busy":"2022-08-03T12:39:18.610974Z","iopub.execute_input":"2022-08-03T12:39:18.611476Z","iopub.status.idle":"2022-08-03T12:39:18.623216Z","shell.execute_reply.started":"2022-08-03T12:39:18.611433Z","shell.execute_reply":"2022-08-03T12:39:18.622245Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_predict = pd.concat(df_predict, axis=0)","metadata":{"execution":{"iopub.status.busy":"2022-08-03T12:39:22.750879Z","iopub.execute_input":"2022-08-03T12:39:22.751337Z","iopub.status.idle":"2022-08-03T12:39:22.983039Z","shell.execute_reply.started":"2022-08-03T12:39:22.751299Z","shell.execute_reply":"2022-08-03T12:39:22.981691Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_predict.to_csv('ptls_submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2022-08-03T12:39:52.450897Z","iopub.execute_input":"2022-08-03T12:39:52.451373Z","iopub.status.idle":"2022-08-03T12:39:55.225840Z","shell.execute_reply.started":"2022-08-03T12:39:52.451337Z","shell.execute_reply":"2022-08-03T12:39:55.224473Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!head ptls_submission.csv","metadata":{"execution":{"iopub.status.busy":"2022-08-03T12:40:29.379294Z","iopub.execute_input":"2022-08-03T12:40:29.379748Z","iopub.status.idle":"2022-08-03T12:40:30.527671Z","shell.execute_reply.started":"2022-08-03T12:40:29.379709Z","shell.execute_reply":"2022-08-03T12:40:30.526038Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!head ../input/amex-default-prediction/sample_submission.csv","metadata":{"execution":{"iopub.status.busy":"2022-08-03T12:40:41.828431Z","iopub.execute_input":"2022-08-03T12:40:41.829907Z","iopub.status.idle":"2022-08-03T12:40:42.941883Z","shell.execute_reply.started":"2022-08-03T12:40:41.829847Z","shell.execute_reply":"2022-08-03T12:40:42.940330Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# More info about ptls\n\nLinks:\n- Repo: https://github.com/dllllb/pytorch-lifestream\n- Demos: https://github.com/dllllb/pytorch-lifestream/tree/main/demo\n- Docs: https://dllllb.github.io/pytorch-lifestream/","metadata":{}},{"cell_type":"code","source":"from IPython.display import HTML\n\nHTML('''\n<!-- Place this tag in your head or just before your close body tag. -->\n<script async defer src=\"https://buttons.github.io/buttons.js\"></script>\n\n<p>Star <code>pytorch-lifestream</code> on github if you like this notebook</p>\n<!-- Place this tag where you want the button to render. -->\n<a class=\"github-button\" href=\"https://github.com/dllllb/pytorch-lifestream\" data-icon=\"octicon-star\" data-size=\"large\" data-show-count=\"true\" aria-label=\"Star dllllb/pytorch-lifestream on GitHub\">Star</a>\n''')","metadata":{"execution":{"iopub.status.busy":"2022-08-04T04:18:22.473196Z","iopub.execute_input":"2022-08-04T04:18:22.474436Z","iopub.status.idle":"2022-08-04T04:18:22.483128Z","shell.execute_reply.started":"2022-08-04T04:18:22.474351Z","shell.execute_reply":"2022-08-04T04:18:22.482339Z"},"_kg_hide-input":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}