{"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-12T10:46:46.538805Z","iopub.execute_input":"2022-08-12T10:46:46.540024Z","iopub.status.idle":"2022-08-12T10:46:46.564139Z","shell.execute_reply.started":"2022-08-12T10:46:46.539885Z","shell.execute_reply":"2022-08-12T10:46:46.563122Z"},"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-12T10:46:46.569719Z","iopub.execute_input":"2022-08-12T10:46:46.572034Z","iopub.status.idle":"2022-08-12T10:48:54.444331Z","shell.execute_reply.started":"2022-08-12T10:46:46.571997Z","shell.execute_reply":"2022-08-12T10:48:54.443149Z"},"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-12T10:48:54.447349Z","iopub.execute_input":"2022-08-12T10:48:54.447752Z","iopub.status.idle":"2022-08-12T10:49:00.486614Z","shell.execute_reply.started":"2022-08-12T10:48:54.447714Z","shell.execute_reply":"2022-08-12T10:49:00.485001Z"},"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-12T10:49:00.488339Z","iopub.execute_input":"2022-08-12T10:49:00.489084Z","iopub.status.idle":"2022-08-12T10:49:00.566148Z","shell.execute_reply.started":"2022-08-12T10:49:00.489048Z","shell.execute_reply":"2022-08-12T10:49:00.564959Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preprocessor.get_category_dictionary_sizes()","metadata":{"execution":{"iopub.status.busy":"2022-08-12T10:49:00.568779Z","iopub.execute_input":"2022-08-12T10:49:00.569123Z","iopub.status.idle":"2022-08-12T10:49:00.579364Z","shell.execute_reply.started":"2022-08-12T10:49:00.569085Z","shell.execute_reply":"2022-08-12T10:49:00.578213Z"},"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-12T10:49:00.580859Z","iopub.execute_input":"2022-08-12T10:49:00.581569Z","iopub.status.idle":"2022-08-12T10:49:00.717550Z","shell.execute_reply.started":"2022-08-12T10:49:00.581534Z","shell.execute_reply":"2022-08-12T10:49:00.716429Z"},"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-12T10:49:00.718901Z","iopub.execute_input":"2022-08-12T10:49:00.719665Z","iopub.status.idle":"2022-08-12T10:49:00.730979Z","shell.execute_reply.started":"2022-08-12T10:49:00.719614Z","shell.execute_reply":"2022-08-12T10:49:00.729513Z"},"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-12T10:49:00.734228Z","iopub.execute_input":"2022-08-12T10:49:00.734778Z","iopub.status.idle":"2022-08-12T10:49:00.743375Z","shell.execute_reply.started":"2022-08-12T10:49:00.734741Z","shell.execute_reply":"2022-08-12T10:49:00.742352Z"},"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-12T10:49:00.745319Z","iopub.execute_input":"2022-08-12T10:49:00.746005Z","iopub.status.idle":"2022-08-12T10:49:00.755786Z","shell.execute_reply.started":"2022-08-12T10:49:00.745972Z","shell.execute_reply":"2022-08-12T10:49:00.754697Z"},"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-12T10:49:00.758059Z","iopub.execute_input":"2022-08-12T10:49:00.758797Z","iopub.status.idle":"2022-08-12T10:49:00.769743Z","shell.execute_reply.started":"2022-08-12T10:49:00.758763Z","shell.execute_reply":"2022-08-12T10:49:00.768528Z"},"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-12T10:49:00.775379Z","iopub.execute_input":"2022-08-12T10:49:00.776179Z","iopub.status.idle":"2022-08-12T10:49:00.786780Z","shell.execute_reply.started":"2022-08-12T10:49:00.776134Z","shell.execute_reply":"2022-08-12T10:49:00.785586Z"},"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-12T10:49:00.788409Z","iopub.execute_input":"2022-08-12T10:49:00.788950Z","iopub.status.idle":"2022-08-12T10:49:00.807595Z","shell.execute_reply.started":"2022-08-12T10:49:00.788917Z","shell.execute_reply":"2022-08-12T10:49:00.806303Z"},"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-12T10:49:00.809095Z","iopub.execute_input":"2022-08-12T10:49:00.810319Z","iopub.status.idle":"2022-08-12T10:49:00.827676Z","shell.execute_reply.started":"2022-08-12T10:49:00.810271Z","shell.execute_reply":"2022-08-12T10:49:00.826404Z"},"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-12T10:49:00.829020Z","iopub.execute_input":"2022-08-12T10:49:00.830297Z","iopub.status.idle":"2022-08-12T10:49:00.843898Z","shell.execute_reply.started":"2022-08-12T10:49:00.830228Z","shell.execute_reply":"2022-08-12T10:49:00.842606Z"},"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-12T10:49:00.845293Z","iopub.execute_input":"2022-08-12T10:49:00.846325Z","iopub.status.idle":"2022-08-12T10:49:00.861307Z","shell.execute_reply.started":"2022-08-12T10:49:00.846290Z","shell.execute_reply":"2022-08-12T10:49:00.859857Z"},"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-12T10:49:00.864035Z","iopub.execute_input":"2022-08-12T10:49:00.864916Z","iopub.status.idle":"2022-08-12T10:49:00.877576Z","shell.execute_reply.started":"2022-08-12T10:49:00.864865Z","shell.execute_reply":"2022-08-12T10:49:00.876229Z"},"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-12T10:49:00.879428Z","iopub.execute_input":"2022-08-12T10:49:00.880180Z","iopub.status.idle":"2022-08-12T10:49:00.979456Z","shell.execute_reply.started":"2022-08-12T10:49:00.880129Z","shell.execute_reply":"2022-08-12T10:49:00.977956Z"},"trusted":true},"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-12T11:11:47.860366Z","iopub.execute_input":"2022-08-12T11:11:47.860976Z","iopub.status.idle":"2022-08-12T11:35:53.483193Z","shell.execute_reply.started":"2022-08-12T11:11:47.860942Z","shell.execute_reply":"2022-08-12T11:35:53.481525Z"},"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-12T11:35:53.485517Z","iopub.execute_input":"2022-08-12T11:35:53.486120Z","iopub.status.idle":"2022-08-12T11:35:53.495237Z","shell.execute_reply.started":"2022-08-12T11:35:53.486060Z","shell.execute_reply":"2022-08-12T11:35:53.493318Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Without error if batchsize = 1","metadata":{}},{"cell_type":"code","source":"train_dataloader = torch.utils.data.DataLoader(\n    all_train_dataset,\n    collate_fn=collate_feature_dict,\n    num_workers=4,\n    batch_size=1,\n)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_embeds = torch.vstack(trainer.predict(model, train_dataloader))","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_embeds","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# ERROR without torch.nn.Flatten(start_dim=0)","metadata":{}},{"cell_type":"code","source":"model.head = torch.nn.Identity()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataloader = torch.utils.data.DataLoader(\n    all_train_dataset,\n    collate_fn=collate_feature_dict,\n    num_workers=4,\n    batch_size=1024,\n)","metadata":{"execution":{"iopub.status.busy":"2022-08-12T11:49:07.565256Z","iopub.execute_input":"2022-08-12T11:49:07.566102Z","iopub.status.idle":"2022-08-12T11:49:07.572536Z","shell.execute_reply.started":"2022-08-12T11:49:07.566063Z","shell.execute_reply":"2022-08-12T11:49:07.571472Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_embeds = torch.vstack(trainer.predict(model, train_dataloader))","metadata":{"execution":{"iopub.status.busy":"2022-08-12T11:49:11.878695Z","iopub.execute_input":"2022-08-12T11:49:11.879129Z","iopub.status.idle":"2022-08-12T11:52:01.031655Z","shell.execute_reply.started":"2022-08-12T11:49:11.879092Z","shell.execute_reply":"2022-08-12T11:52:01.029508Z"},"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-12T11:35:53.497499Z","iopub.execute_input":"2022-08-12T11:35:53.498049Z","iopub.status.idle":"2022-08-12T11:41:10.618034Z","shell.execute_reply.started":"2022-08-12T11:35:53.498007Z","shell.execute_reply":"2022-08-12T11:41:10.616496Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_predict[0].head()","metadata":{"execution":{"iopub.status.busy":"2022-08-12T11:41:10.620266Z","iopub.execute_input":"2022-08-12T11:41:10.620853Z","iopub.status.idle":"2022-08-12T11:41:10.644936Z","shell.execute_reply.started":"2022-08-12T11:41:10.620797Z","shell.execute_reply":"2022-08-12T11:41:10.643683Z"},"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-12T11:41:10.646871Z","iopub.execute_input":"2022-08-12T11:41:10.647208Z","iopub.status.idle":"2022-08-12T11:41:10.913547Z","shell.execute_reply.started":"2022-08-12T11:41:10.647176Z","shell.execute_reply":"2022-08-12T11:41:10.912313Z"},"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-12T11:41:10.916137Z","iopub.execute_input":"2022-08-12T11:41:10.916502Z","iopub.status.idle":"2022-08-12T11:41:13.555767Z","shell.execute_reply.started":"2022-08-12T11:41:10.916469Z","shell.execute_reply":"2022-08-12T11:41:13.554674Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!head ptls_submission.csv","metadata":{"execution":{"iopub.status.busy":"2022-08-12T11:41:13.563197Z","iopub.execute_input":"2022-08-12T11:41:13.563620Z","iopub.status.idle":"2022-08-12T11:41:14.928367Z","shell.execute_reply.started":"2022-08-12T11:41:13.563586Z","shell.execute_reply":"2022-08-12T11:41:14.926602Z"},"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-12T11:41:14.930950Z","iopub.execute_input":"2022-08-12T11:41:14.931655Z","iopub.status.idle":"2022-08-12T11:41:16.315402Z","shell.execute_reply.started":"2022-08-12T11:41:14.931594Z","shell.execute_reply":"2022-08-12T11:41:16.313032Z"},"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":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-08-12T11:41:16.319556Z","iopub.execute_input":"2022-08-12T11:41:16.323536Z","iopub.status.idle":"2022-08-12T11:41:16.338191Z","shell.execute_reply.started":"2022-08-12T11:41:16.323481Z","shell.execute_reply":"2022-08-12T11:41:16.336529Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}