{"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":"import shutil\ntry:\n    src_dir = '/kaggle/input/einops'\n    dest_dir = '/kaggle/working/einops'\n\n    shutil.copytree(src_dir, dest_dir)\n\n    !pip install ../working/einops/einops-master\nexcept:\n    print('Did copy!')","metadata":{"execution":{"iopub.status.busy":"2023-06-27T11:27:51.438137Z","iopub.execute_input":"2023-06-27T11:27:51.438517Z","iopub.status.idle":"2023-06-27T11:28:25.878556Z","shell.execute_reply.started":"2023-06-27T11:27:51.438488Z","shell.execute_reply":"2023-06-27T11:28:25.877491Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nfrom torch import nn, einsum\nimport polars as pl\nfrom random import randrange\nimport numpy as np\nimport torch.nn.functional as F\nimport torch.optim as optim\nimport pandas as pd\nfrom einops import rearrange, repeat\nfrom einops.layers.torch import Rearrange, Reduce","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-06-27T11:28:25.880316Z","iopub.execute_input":"2023-06-27T11:28:25.880644Z","iopub.status.idle":"2023-06-27T11:28:29.208034Z","shell.execute_reply.started":"2023-06-27T11:28:25.880610Z","shell.execute_reply":"2023-06-27T11:28:29.207187Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model set up","metadata":{}},{"cell_type":"code","source":"# utils\ndef exists(val):\n    return val is not None\n\ndef pair(val):\n    return (val, val) if not isinstance(val, tuple) else val\n\ndef default(val, d):\n    return val if exists(val) else d","metadata":{"execution":{"iopub.status.busy":"2023-06-27T11:28:29.209110Z","iopub.execute_input":"2023-06-27T11:28:29.209784Z","iopub.status.idle":"2023-06-27T11:28:29.215355Z","shell.execute_reply.started":"2023-06-27T11:28:29.209753Z","shell.execute_reply":"2023-06-27T11:28:29.214351Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#g_mlp\ndef dropout_layers(layers, prob_survival):\n    if prob_survival == 1:\n        return layers\n\n    num_layers = len(layers)\n    to_drop = torch.zeros(num_layers).uniform_(0., 1.) > prob_survival\n\n    # make sure at least one layer makes it\n    if all(to_drop):\n        rand_index = randrange(num_layers)\n        to_drop[rand_index] = False\n\n    layers = [layer for (layer, drop) in zip(layers, to_drop) if not drop]\n    return layers\n\n\ndef shift(t, amount, mask = None):\n    if amount == 0:\n        return t\n    return F.pad(t, (0, 0, amount, -amount), value = 0.)\n\n\nclass PreShiftTokens(nn.Module):\n    def __init__(self, shifts, fn):\n        super().__init__()\n        self.fn = fn\n        self.shifts = tuple(shifts)\n\n    def forward(self, x, **kwargs):\n        if self.shifts == (0,):\n            return self.fn(x, **kwargs)\n\n        shifts = self.shifts\n        segments = len(shifts)\n        feats_per_shift = x.shape[-1] // segments\n        splitted = x.split(feats_per_shift, dim = -1)\n        segments_to_shift, rest = splitted[:segments], splitted[segments:]\n        segments_to_shift = list(map(lambda args: shift(*args), zip(segments_to_shift, shifts)))\n        x = torch.cat((*segments_to_shift, *rest), dim = -1)\n        return self.fn(x, **kwargs)\n\n\nclass Attention(nn.Module):\n    def __init__(self, dim_in, dim_out, dim_inner, causal = False):\n        super().__init__()\n        self.scale = dim_inner ** -0.5\n        self.causal = causal\n\n        self.to_qkv = nn.Linear(dim_in, dim_inner * 3, bias = False)\n        self.to_out = nn.Linear(dim_inner, dim_out)\n\n    def forward(self, x):\n        device = x.device\n        q, k, v = self.to_qkv(x).chunk(3, dim = -1)\n        sim = einsum('b i d, b j d -> b i j', q, k) * self.scale\n\n        if self.causal:\n            mask = torch.ones(sim.shape[-2:], device = device).triu(1).bool()\n            sim.masked_fill_(mask[None, ...], -torch.finfo(q.dtype).max)\n\n        attn = sim.softmax(dim = -1)\n        out = einsum('b i j, b j d -> b i d', attn, v)\n        return self.to_out(out)\n\n\nclass SpatialGatingUnit(nn.Module):\n    def __init__(\n        self,\n        dim,\n        dim_seq,\n        causal = False,\n        act = nn.Identity(),\n        heads = 1,\n        init_eps = 1e-3,\n        circulant_matrix = False\n    ):\n        super().__init__()\n        dim_out = dim // 2\n        self.heads = heads\n        self.causal = causal\n        self.norm = nn.LayerNorm(dim_out)\n\n        self.act = act\n\n        # parameters\n\n        if circulant_matrix:\n            self.circulant_pos_x = nn.Parameter(torch.ones(heads, dim_seq))\n            self.circulant_pos_y = nn.Parameter(torch.ones(heads, dim_seq))\n\n        self.circulant_matrix = circulant_matrix\n        shape = (heads, dim_seq,) if circulant_matrix else (heads, dim_seq, dim_seq)\n        weight = torch.zeros(shape)\n\n        self.weight = nn.Parameter(weight)\n        init_eps /= dim_seq\n        nn.init.uniform_(self.weight, -init_eps, init_eps)\n\n        self.bias = nn.Parameter(torch.ones(heads, dim_seq))\n\n    def forward(self, x, gate_res = None):\n        device, n, h = x.device, x.shape[1], self.heads\n\n        res, gate = x.chunk(2, dim = -1)\n        gate = self.norm(gate)\n\n        weight, bias = self.weight, self.bias\n\n        if self.circulant_matrix:\n            # build the circulant matrix\n\n            dim_seq = weight.shape[-1]\n            weight = F.pad(weight, (0, dim_seq), value = 0)\n            weight = repeat(weight, '... n -> ... (r n)', r = dim_seq)\n            weight = weight[:, :-dim_seq].reshape(h, dim_seq, 2 * dim_seq - 1)\n            weight = weight[:, :, (dim_seq - 1):]\n\n            # give circulant matrix absolute position awareness\n\n            pos_x, pos_y = self.circulant_pos_x, self.circulant_pos_y\n            weight = weight * rearrange(pos_x, 'h i -> h i ()') * rearrange(pos_y, 'h j -> h () j')\n\n        if self.causal:\n            weight, bias = weight[:, :n, :n], bias[:, :n]\n            mask = torch.ones(weight.shape[-2:], device = device).triu_(1).bool()\n            mask = rearrange(mask, 'i j -> () i j')\n            weight = weight.masked_fill(mask, 0.)\n\n        gate = rearrange(gate, 'b n (h d) -> b h n d', h = h)\n\n        gate = einsum('b h n d, h m n -> b h m d', gate, weight)\n        gate = gate + rearrange(bias, 'h n -> () h n ()')\n\n        gate = rearrange(gate, 'b h n d -> b n (h d)')\n\n        if exists(gate_res):\n            gate = gate + gate_res\n\n        return self.act(gate) * res\n\n\nclass gMLPBlock(nn.Module):\n    def __init__(\n        self,\n        *,\n        dim,\n        dim_ff,\n        seq_len,\n        heads = 1,\n        attn_dim = None,\n        causal = False,\n        act = nn.Identity(),\n        circulant_matrix = False\n    ):\n        super().__init__()\n        self.proj_in = nn.Sequential(\n            nn.Linear(dim, dim_ff),\n            nn.GELU()\n        )\n\n        self.attn = Attention(dim, dim_ff // 2, attn_dim, causal) if exists(attn_dim) else None\n\n        self.sgu = SpatialGatingUnit(dim_ff, seq_len, causal, act, heads, circulant_matrix = circulant_matrix)\n        self.proj_out = nn.Linear(dim_ff // 2, dim)\n\n    def forward(self, x):\n        gate_res = self.attn(x) if exists(self.attn) else None\n        x = self.proj_in(x)\n        x = self.sgu(x, gate_res = gate_res)\n        x = self.proj_out(x)\n        return x\n\n\nclass gMLP(nn.Module):\n    def __init__(\n        self,\n        *,\n        num_tokens = None,\n        dim,\n        depth,\n        seq_len,\n        heads = 1,\n        ff_mult = 4,\n        attn_dim = None,\n        prob_survival = 1.,\n        causal = False,\n        circulant_matrix = False,\n        shift_tokens = 0,\n        act = nn.Identity()\n    ):\n        super().__init__()\n        assert (dim % heads) == 0, 'dimension must be divisible by number of heads'\n\n        dim_ff = dim * ff_mult\n        self.seq_len = seq_len\n        self.prob_survival = prob_survival\n\n        self.to_embed = nn.Embedding(num_tokens, dim) if exists(num_tokens) else nn.Identity()\n\n        token_shifts = tuple(range(0 if causal else -shift_tokens, shift_tokens + 1))\n        self.layers = nn.ModuleList([Residual(PreNorm(dim, PreShiftTokens(token_shifts, gMLPBlock(dim = dim, heads = heads, dim_ff = dim_ff, seq_len = seq_len, attn_dim = attn_dim, causal = causal, act = act, circulant_matrix = circulant_matrix)))) for i in range(depth)])\n\n        self.to_logits = nn.Sequential(\n            nn.LayerNorm(dim),\n            Reduce('b n d -> b d', 'mean'),\n            nn.Linear(dim, 1)\n        )\n\n    def forward(self, x):\n        x = self.to_embed(x)\n        layers = self.layers if not self.training else dropout_layers(self.layers, self.prob_survival)\n        out = nn.Sequential(*layers)(x)\n        return self.to_logits(out)\n\n\nclass gMLPClassification(nn.Module):\n    def __init__(\n        self,\n        *,\n        patch_width,\n        seq_len,\n        num_classes,\n        dim,\n        depth,\n        heads = 1,\n        ff_mult = 4,\n        attn_dim = None,\n        prob_survival = 1.\n    ):\n        super().__init__()\n        assert (dim % heads) == 0, 'dimension must be divisible by number of heads'\n        num_patches = (seq_len // patch_width)\n\n        dim_ff = dim * ff_mult\n\n        self.to_patch_embed = nn.Sequential(\n            Rearrange('b (w p2) -> b (w) (p2)', p2 = patch_width),\n            nn.Linear(patch_width, dim)\n        )\n\n        self.prob_survival = prob_survival\n\n        self.layers = nn.ModuleList([Residual(PreNorm(dim, gMLPBlock(dim = dim, heads = heads, dim_ff = dim_ff, seq_len = num_patches, attn_dim = attn_dim))) for i in range(depth)])\n\n        self.to_logits = nn.Sequential(\n            nn.LayerNorm(dim),\n            Reduce('b n d -> b d', 'mean'),\n            nn.Linear(dim, num_classes)\n        )\n\n    def forward(self, x):\n        x = self.to_patch_embed(x)\n        layers = self.layers if not self.training else dropout_layers(self.layers, self.prob_survival)\n        x = nn.Sequential(*layers)(x)\n        return self.to_logits(x)\n","metadata":{"execution":{"iopub.status.busy":"2023-06-27T11:28:29.217769Z","iopub.execute_input":"2023-06-27T11:28:29.218206Z","iopub.status.idle":"2023-06-27T11:28:29.257539Z","shell.execute_reply.started":"2023-06-27T11:28:29.218172Z","shell.execute_reply":"2023-06-27T11:28:29.256015Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# shared classes\nclass Residual(nn.Module):\n    def __init__(self, fn):\n        super().__init__()\n        self.fn = fn\n\n    def forward(self, x, **kwargs):\n        return self.fn(x, **kwargs) + x\n    \nclass PreNorm(nn.Module):\n    def __init__(self, dim, fn):\n        super().__init__()\n        self.norm = nn.LayerNorm(dim)\n        self.fn = fn\n\n    def forward(self, x, **kwargs):\n        return self.fn(self.norm(x), **kwargs)","metadata":{"execution":{"iopub.status.busy":"2023-06-27T11:28:29.258902Z","iopub.execute_input":"2023-06-27T11:28:29.259740Z","iopub.status.idle":"2023-06-27T11:28:29.272983Z","shell.execute_reply.started":"2023-06-27T11:28:29.259705Z","shell.execute_reply":"2023-06-27T11:28:29.271844Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class GEGLU(nn.Module):\n    def forward(self, x):\n        x, gates = x.chunk(2, dim = -1)\n        return x * F.gelu(gates)\n\nclass MLP(nn.Module):\n    def __init__(self, dims, act = None):\n        super().__init__()\n        dims_pairs = list(zip(dims[:-1], dims[1:]))\n        layers = []\n        for ind, (dim_in, dim_out) in enumerate(dims_pairs):\n            is_last = ind >= (len(dims_pairs) - 1)\n            linear = nn.Linear(dim_in, dim_out)\n            layers.append(linear)\n\n            if is_last:\n                continue\n\n            act = default(act, nn.ReLU())\n            layers.append(act)\n\n        self.mlp = nn.Sequential(*layers)\n\n    def forward(self, x):\n        return self.mlp(x)\n\nclass HeadAttention(nn.Module):\n    def __init__(\n        self,\n        dim,\n        heads = 8,\n        dim_head = 16,\n        dropout = 0.\n    ):\n        super().__init__()\n        inner_dim = dim_head * heads\n        self.heads = heads\n        self.scale = dim_head ** -0.5\n\n        self.to_qkv = nn.Linear(dim, inner_dim * 3, bias = False)\n        self.to_out = nn.Linear(inner_dim, dim)\n\n        self.dropout = nn.Dropout(dropout)\n\n    def forward(self, x):\n        h = self.heads\n        q, k, v = self.to_qkv(x).chunk(3, dim = -1)\n        q, k, v = map(lambda t: rearrange(t, 'b n (h d) -> b h n d', h = h), (q, k, v))\n        sim = einsum('b h i d, b h j d -> b h i j', q, k) * self.scale\n\n        attn = sim.softmax(dim = -1)\n        attn = self.dropout(attn)\n\n        out = einsum('b h i j, b h j d -> b h i d', attn, v)\n        out = rearrange(out, 'b h n d -> b n (h d)', h = h)\n        return self.to_out(out)\n\nclass FeedForward(nn.Module):\n    def __init__(self, dim, mult = 4, dropout = 0.):\n        super().__init__()\n        self.net = nn.Sequential(\n            nn.Linear(dim, dim * mult * 2),\n            GEGLU(),\n            nn.Dropout(dropout),\n            nn.Linear(dim * mult, dim)\n        )\n\n    def forward(self, x, **kwargs):\n        return self.net(x)\n\nclass Transformer(nn.Module):\n    def __init__(self, num_tokens, dim, depth, heads, dim_head, attn_dropout, ff_dropout):\n        super().__init__()\n        self.embeds = nn.Embedding(num_tokens, dim)\n        self.layers = nn.ModuleList([])\n\n        for _ in range(depth):\n            self.layers.append(nn.ModuleList([\n                Residual(PreNorm(dim, HeadAttention(dim, heads = heads, dim_head = dim_head, dropout = attn_dropout))),\n                Residual(PreNorm(dim, FeedForward(dim, dropout = ff_dropout))),\n            ]))\n\n    def forward(self, x):\n        x = self.embeds(x)\n\n        for attn, ff in self.layers:\n            x = attn(x)\n            x = ff(x)\n\n        return x\n\n\nclass GatedTabTransformer(nn.Module):\n    def __init__(\n        self,\n        *,\n        categories,\n        num_continuous,\n        transformer_dim,\n        transformer_depth,\n        transformer_heads,\n        transformer_dim_head = 16,\n        dim_out = 1,\n        mlp_depth = 2,\n        mlp_act = None,\n        num_special_tokens = 2,\n        continuous_mean_std = None,\n        attn_dropout = 0.,\n        ff_dropout = 0.,\n        gmlp_enabled=False,\n        mlp_dimension=32,\n    ):\n        super().__init__()\n        assert all(map(lambda n: n > 0, categories)), 'number of each category must be positive'\n\n        # categories related calculations\n\n        self.num_categories = len(categories)\n        self.num_unique_categories = sum(categories)\n\n        # create category embeddings table\n\n        self.num_special_tokens = num_special_tokens\n        total_tokens = self.num_unique_categories + num_special_tokens\n\n        # for automatically offsetting unique category ids to the correct position in the categories embedding table\n\n        categories_offset = F.pad(torch.tensor(list(categories)), (1, 0), value = num_special_tokens)\n        categories_offset = categories_offset.cumsum(dim = -1)[:-1]\n        self.register_buffer('categories_offset', categories_offset)\n\n        # continuous\n\n        if exists(continuous_mean_std):\n            assert continuous_mean_std.shape == (num_continuous, 2), f'continuous_mean_std must have a shape of ({num_continuous}, 2) where the last dimension contains the mean and variance respectively'\n        self.register_buffer('continuous_mean_std', continuous_mean_std)\n\n        self.norm = nn.LayerNorm(num_continuous)\n        self.num_continuous = num_continuous\n\n        # transformer\n\n        self.transformer = Transformer(\n            num_tokens = total_tokens,\n            dim = transformer_dim,\n            depth = transformer_depth,\n            heads = transformer_heads,\n            dim_head = transformer_dim_head,\n            attn_dropout = attn_dropout,\n            ff_dropout = ff_dropout\n        )\n\n        # mlp to logits\n\n        input_size = (transformer_dim * self.num_categories) + num_continuous\n\n        if gmlp_enabled:\n            self.mlp = gMLPClassification(\n                patch_width=1,\n                seq_len=input_size,\n                num_classes = dim_out,\n                dim = mlp_dimension,\n                depth = mlp_depth\n            )\n        else:\n            hidden_dimensions = []\n\n            for i in range(mlp_depth):\n                if mlp_dimension == -1:\n                    hidden_dimensions.append((input_size // 8) * (2**(mlp_depth - i)))\n                else:\n                    hidden_dimensions.append(mlp_dimension)\n            \n            all_dimensions = [input_size, *hidden_dimensions, dim_out]\n            self.mlp = MLP(all_dimensions, act = mlp_act)\n\n    def forward(self, x_categ, x_cont=None):\n        assert x_categ.shape[-1] == self.num_categories, f'you must pass in {self.num_categories} values for your categories input'\n        x_categ += self.categories_offset\n\n        x = self.transformer(x_categ)\n\n        flat_categ = x.flatten(1)\n\n        if self.num_continuous != 0:\n            assert x_cont.shape[1] == self.num_continuous, f'you must pass in {self.num_continuous} values for your continuous input'\n\n            if exists(self.continuous_mean_std):\n                mean, std = self.continuous_mean_std.unbind(dim = -1)\n                x_cont = (x_cont - mean) / std\n\n            normed_cont = self.norm(x_cont)\n\n            x = torch.cat((flat_categ, normed_cont), dim = -1)\n        else:\n            x = flat_categ\n\n        return self.mlp(x)","metadata":{"execution":{"iopub.status.busy":"2023-06-27T11:28:29.274494Z","iopub.execute_input":"2023-06-27T11:28:29.274845Z","iopub.status.idle":"2023-06-27T11:28:29.303734Z","shell.execute_reply.started":"2023-06-27T11:28:29.274804Z","shell.execute_reply":"2023-06-27T11:28:29.302560Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def model_loader(n_cat, target_dim, shape, path):\n    model = GatedTabTransformer(\n        categories = n_cat,                          # tuple containing the number of unique values within each category\n        num_continuous = shape,               # number of continuous values\n        transformer_dim = config[\"transformer_dim\"],        # dimension, paper set at 32\n        dim_out = target_dim,                                       # binary prediction, but could be anything\n        transformer_depth = config[\"transformer_depth\"],    # depth, paper recommended 6\n        transformer_heads = config[\"transformer_heads\"],    # heads, paper recommends 8\n        attn_dropout = config[\"dropout\"],                   # post-attention dropout\n        ff_dropout = config[\"dropout\"],                     # feed forward dropout\n        mlp_act = nn.LeakyReLU(config[\"relu_slope\"]),       # activation for final mlp, defaults to relu, but could be anything else (selu, etc.)\n        mlp_depth=config[\"mlp_depth\"],                      # mlp hidden layers depth\n        mlp_dimension=config[\"mlp_dimension\"],              # dimension of mlp layers\n        gmlp_enabled=config[\"gmlp_enabled\"]                 # gmlp or standard mlp\n    )\n    \n    checkpoint = torch.load(path, map_location='cpu')\n    model.load_state_dict(checkpoint)\n    model.to('cpu')\n    print('Loaded model type:', path)\n    return model","metadata":{"execution":{"iopub.status.busy":"2023-06-27T11:28:29.305337Z","iopub.execute_input":"2023-06-27T11:28:29.305912Z","iopub.status.idle":"2023-06-27T11:28:29.315543Z","shell.execute_reply.started":"2023-06-27T11:28:29.305874Z","shell.execute_reply":"2023-06-27T11:28:29.314525Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"path04 = '/kaggle/input/gated-tf-psp/model_0to4.pt'\npath512 = '/kaggle/input/gated-tf-psp/model_5to12.pt'\npath1322 = '/kaggle/input/gated-tf-psp/model_13to22.pt'\n\nconfig = {\"batch_size\": 128,\n            \"patience\": 5,\n            \"initial_lr\": 1e-3,\n            \"scheduler_gamma\": 0.1,\n            \"scheduler_step\": 8,\n            \"relu_slope\": 0,\n            \"transformer_heads\": 8,\n            \"transformer_depth\": 6,\n            \"transformer_dim\": 8,\n            \"gmlp_enabled\": True,\n            \"mlp_depth\": 6,\n            \"mlp_dimension\": 64,\n            \"dropout\": 0.2\n        }","metadata":{"execution":{"iopub.status.busy":"2023-06-27T11:29:07.442348Z","iopub.execute_input":"2023-06-27T11:29:07.442763Z","iopub.status.idle":"2023-06-27T11:29:07.448380Z","shell.execute_reply.started":"2023-06-27T11:29:07.442732Z","shell.execute_reply":"2023-06-27T11:29:07.447552Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tf1322 = model_loader([4, 4, 38, 6, 24], 5, 1064, path1322)\ntf04 = model_loader([5, 4, 17, 3, 18], 3, 600, path04)\ntf512 = model_loader([4, 4, 23, 4, 24], 10, 916, path512)","metadata":{"execution":{"iopub.status.busy":"2023-06-27T11:29:10.155091Z","iopub.execute_input":"2023-06-27T11:29:10.155454Z","iopub.status.idle":"2023-06-27T11:29:11.342886Z","shell.execute_reply.started":"2023-06-27T11:29:10.155427Z","shell.execute_reply":"2023-06-27T11:29:11.342099Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Process data","metadata":{}},{"cell_type":"code","source":"with open('/kaggle/input/gated-tf-psp/bowl_1322_col_names.txt', 'r') as f:\n    col1322 = f.read().splitlines()\nwith open('/kaggle/input/gated-tf-psp/bowl_04_col_names.txt', 'r') as f:\n    col04 = f.read().splitlines()\nwith open('/kaggle/input/gated-tf-psp/bowl_512_col_names.txt', 'r') as f:\n    col512 = f.read().splitlines()","metadata":{"execution":{"iopub.status.busy":"2023-06-27T11:29:15.944243Z","iopub.execute_input":"2023-06-27T11:29:15.944623Z","iopub.status.idle":"2023-06-27T11:29:15.959603Z","shell.execute_reply.started":"2023-06-27T11:29:15.944594Z","shell.execute_reply":"2023-06-27T11:29:15.958684Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CATS = ['event_name', 'name', 'fqid', 'room_fqid', 'text_fqid']\nNUMS = ['page', 'room_coor_x', 'room_coor_y', 'screen_coor_x', 'screen_coor_y',\n        'hover_duration', 'elapsed_time_diff']\n\nname_feature = ['basic', 'undefined', 'close', 'open', 'prev', 'next']\nevent_name_feature = ['cutscene_click', 'person_click', 'navigate_click',\n       'observation_click', 'notification_click', 'object_click',\n       'object_hover', 'map_hover', 'map_click', 'checkpoint',\n       'notebook_click']\n\n# from https://www.kaggle.com/code/leehomhuang/catboost-baseline-with-lots-features-inference :\nfqid_lists = ['worker', 'archivist', 'gramps', 'wells', 'toentry', 'confrontation', 'crane_ranger', 'groupconvo', 'flag_girl', 'tomap', 'tostacks', 'tobasement', 'archivist_glasses', 'boss', 'journals', 'seescratches', 'groupconvo_flag', 'cs', 'teddy', 'expert', 'businesscards', 'ch3start', 'tunic.historicalsociety', 'tofrontdesk', 'savedteddy', 'plaque', 'glasses', 'tunic.drycleaner', 'reader_flag', 'tunic.library', 'tracks', 'tunic.capitol_2', 'trigger_scarf', 'reader', 'directory', 'tunic.capitol_1', 'journals.pic_0.next', 'unlockdoor', 'tunic', 'what_happened', 'tunic.kohlcenter', 'tunic.humanecology', 'colorbook', 'logbook', 'businesscards.card_0.next', 'journals.hub.topics', 'logbook.page.bingo', 'journals.pic_1.next', 'journals_flag', 'reader.paper0.next', 'tracks.hub.deer', 'reader_flag.paper0.next', 'trigger_coffee', 'wellsbadge', 'journals.pic_2.next', 'tomicrofiche', 'journals_flag.pic_0.bingo', 'plaque.face.date', 'notebook', 'tocloset_dirty', 'businesscards.card_bingo.bingo', 'businesscards.card_1.next', 'tunic.wildlife', 'tunic.hub.slip', 'tocage', 'journals.pic_2.bingo', 'tocollectionflag', 'tocollection', 'chap4_finale_c', 'chap2_finale_c', 'lockeddoor', 'journals_flag.hub.topics', 'tunic.capitol_0', 'reader_flag.paper2.bingo', 'photo', 'tunic.flaghouse', 'reader.paper1.next', 'directory.closeup.archivist', 'intro', 'businesscards.card_bingo.next', 'reader.paper2.bingo', 'retirement_letter', 'remove_cup', 'journals_flag.pic_0.next', 'magnify', 'coffee', 'key', 'togrampa', 'reader_flag.paper1.next', 'janitor', 'tohallway', 'chap1_finale', 'report', 'outtolunch', 'journals_flag.hub.topics_old', 'journals_flag.pic_1.next', 'reader.paper2.next', 'chap1_finale_c', 'reader_flag.paper2.next', 'door_block_talk', 'journals_flag.pic_1.bingo', 'journals_flag.pic_2.next', 'journals_flag.pic_2.bingo', 'block_magnify', 'reader.paper0.prev', 'block', 'reader_flag.paper0.prev', 'block_0', 'door_block_clean', 'reader.paper2.prev', 'reader.paper1.prev', 'doorblock', 'tocloset', 'reader_flag.paper2.prev', 'reader_flag.paper1.prev', 'block_tomap2', 'journals_flag.pic_0_old.next', 'journals_flag.pic_1_old.next', 'block_tocollection', 'block_nelson', 'journals_flag.pic_2_old.next', 'block_tomap1', 'block_badge', 'need_glasses', 'block_badge_2', 'fox', 'block_1']\ntext_lists = ['tunic.historicalsociety.cage.confrontation', 'tunic.wildlife.center.crane_ranger.crane', 'tunic.historicalsociety.frontdesk.archivist.newspaper', 'tunic.historicalsociety.entry.groupconvo', 'tunic.wildlife.center.wells.nodeer', 'tunic.historicalsociety.frontdesk.archivist.have_glass', 'tunic.drycleaner.frontdesk.worker.hub', 'tunic.historicalsociety.closet_dirty.gramps.news', 'tunic.humanecology.frontdesk.worker.intro', 'tunic.historicalsociety.frontdesk.archivist_glasses.confrontation', 'tunic.historicalsociety.basement.seescratches', 'tunic.historicalsociety.collection.cs', 'tunic.flaghouse.entry.flag_girl.hello', 'tunic.historicalsociety.collection.gramps.found', 'tunic.historicalsociety.basement.ch3start', 'tunic.historicalsociety.entry.groupconvo_flag', 'tunic.library.frontdesk.worker.hello', 'tunic.library.frontdesk.worker.wells', 'tunic.historicalsociety.collection_flag.gramps.flag', 'tunic.historicalsociety.basement.savedteddy', 'tunic.library.frontdesk.worker.nelson', 'tunic.wildlife.center.expert.removed_cup', 'tunic.library.frontdesk.worker.flag', 'tunic.historicalsociety.frontdesk.archivist.hello', 'tunic.historicalsociety.closet.gramps.intro_0_cs_0', 'tunic.historicalsociety.entry.boss.flag', 'tunic.flaghouse.entry.flag_girl.symbol', 'tunic.historicalsociety.closet_dirty.trigger_scarf', 'tunic.drycleaner.frontdesk.worker.done', 'tunic.historicalsociety.closet_dirty.what_happened', 'tunic.wildlife.center.wells.animals', 'tunic.historicalsociety.closet.teddy.intro_0_cs_0', 'tunic.historicalsociety.cage.glasses.afterteddy', 'tunic.historicalsociety.cage.teddy.trapped', 'tunic.historicalsociety.cage.unlockdoor', 'tunic.historicalsociety.stacks.journals.pic_2.bingo', 'tunic.historicalsociety.entry.wells.flag', 'tunic.humanecology.frontdesk.worker.badger', 'tunic.historicalsociety.stacks.journals_flag.pic_0.bingo', 'tunic.historicalsociety.closet.intro', 'tunic.historicalsociety.closet.retirement_letter.hub', 'tunic.historicalsociety.entry.directory.closeup.archivist', 'tunic.historicalsociety.collection.tunic.slip', 'tunic.kohlcenter.halloffame.plaque.face.date', 'tunic.historicalsociety.closet_dirty.trigger_coffee', 'tunic.drycleaner.frontdesk.logbook.page.bingo', 'tunic.library.microfiche.reader.paper2.bingo', 'tunic.kohlcenter.halloffame.togrampa', 'tunic.capitol_2.hall.boss.haveyougotit', 'tunic.wildlife.center.wells.nodeer_recap', 'tunic.historicalsociety.cage.glasses.beforeteddy', 'tunic.historicalsociety.closet_dirty.gramps.helpclean', 'tunic.wildlife.center.expert.recap', 'tunic.historicalsociety.frontdesk.archivist.have_glass_recap', 'tunic.historicalsociety.stacks.journals_flag.pic_1.bingo', 'tunic.historicalsociety.cage.lockeddoor', 'tunic.historicalsociety.stacks.journals_flag.pic_2.bingo', 'tunic.historicalsociety.collection.gramps.lost', 'tunic.historicalsociety.closet.notebook', 'tunic.historicalsociety.frontdesk.magnify', 'tunic.humanecology.frontdesk.businesscards.card_bingo.bingo', 'tunic.wildlife.center.remove_cup', 'tunic.library.frontdesk.wellsbadge.hub', 'tunic.wildlife.center.tracks.hub.deer', 'tunic.historicalsociety.frontdesk.key', 'tunic.library.microfiche.reader_flag.paper2.bingo', 'tunic.flaghouse.entry.colorbook', 'tunic.wildlife.center.coffee', 'tunic.capitol_1.hall.boss.haveyougotit', 'tunic.historicalsociety.basement.janitor', 'tunic.historicalsociety.collection_flag.gramps.recap', 'tunic.wildlife.center.wells.animals2', 'tunic.flaghouse.entry.flag_girl.symbol_recap', 'tunic.historicalsociety.closet_dirty.photo', 'tunic.historicalsociety.stacks.outtolunch', 'tunic.library.frontdesk.worker.wells_recap', 'tunic.historicalsociety.frontdesk.archivist_glasses.confrontation_recap', 'tunic.capitol_0.hall.boss.talktogramps', 'tunic.historicalsociety.closet.photo', 'tunic.historicalsociety.collection.tunic', 'tunic.historicalsociety.closet.teddy.intro_0_cs_5', 'tunic.historicalsociety.closet_dirty.gramps.archivist', 'tunic.historicalsociety.closet_dirty.door_block_talk', 'tunic.historicalsociety.entry.boss.flag_recap', 'tunic.historicalsociety.frontdesk.archivist.need_glass_0', 'tunic.historicalsociety.entry.wells.talktogramps', 'tunic.historicalsociety.frontdesk.block_magnify', 'tunic.historicalsociety.frontdesk.archivist.foundtheodora', 'tunic.historicalsociety.closet_dirty.gramps.nothing', 'tunic.historicalsociety.closet_dirty.door_block_clean', 'tunic.capitol_1.hall.boss.writeitup', 'tunic.library.frontdesk.worker.nelson_recap', 'tunic.library.frontdesk.worker.hello_short', 'tunic.historicalsociety.stacks.block', 'tunic.historicalsociety.frontdesk.archivist.need_glass_1', 'tunic.historicalsociety.entry.boss.talktogramps', 'tunic.historicalsociety.frontdesk.archivist.newspaper_recap', 'tunic.historicalsociety.entry.wells.flag_recap', 'tunic.drycleaner.frontdesk.worker.done2', 'tunic.library.frontdesk.worker.flag_recap', 'tunic.humanecology.frontdesk.block_0', 'tunic.library.frontdesk.worker.preflag', 'tunic.historicalsociety.basement.gramps.seeyalater', 'tunic.flaghouse.entry.flag_girl.hello_recap', 'tunic.historicalsociety.closet.doorblock', 'tunic.drycleaner.frontdesk.worker.takealook', 'tunic.historicalsociety.basement.gramps.whatdo', 'tunic.library.frontdesk.worker.droppedbadge', 'tunic.historicalsociety.entry.block_tomap2', 'tunic.library.frontdesk.block_nelson', 'tunic.library.microfiche.block_0', 'tunic.historicalsociety.entry.block_tocollection', 'tunic.historicalsociety.entry.block_tomap1', 'tunic.historicalsociety.collection.gramps.look_0', 'tunic.library.frontdesk.block_badge', 'tunic.historicalsociety.cage.need_glasses', 'tunic.library.frontdesk.block_badge_2', 'tunic.kohlcenter.halloffame.block_0', 'tunic.capitol_0.hall.chap1_finale_c', 'tunic.capitol_1.hall.chap2_finale_c', 'tunic.capitol_2.hall.chap4_finale_c', 'tunic.wildlife.center.fox.concern', 'tunic.drycleaner.frontdesk.block_0', 'tunic.historicalsociety.entry.gramps.hub', 'tunic.humanecology.frontdesk.block_1', 'tunic.drycleaner.frontdesk.block_1']\nroom_lists = ['tunic.historicalsociety.entry', 'tunic.wildlife.center', 'tunic.historicalsociety.cage', 'tunic.library.frontdesk', 'tunic.historicalsociety.frontdesk', 'tunic.historicalsociety.stacks', 'tunic.historicalsociety.closet_dirty', 'tunic.humanecology.frontdesk', 'tunic.historicalsociety.basement', 'tunic.kohlcenter.halloffame', 'tunic.library.microfiche', 'tunic.drycleaner.frontdesk', 'tunic.historicalsociety.collection', 'tunic.historicalsociety.closet', 'tunic.flaghouse.entry', 'tunic.historicalsociety.collection_flag', 'tunic.capitol_1.hall', 'tunic.capitol_0.hall', 'tunic.capitol_2.hall']","metadata":{"execution":{"iopub.status.busy":"2023-06-27T11:29:17.638091Z","iopub.execute_input":"2023-06-27T11:29:17.638476Z","iopub.status.idle":"2023-06-27T11:29:17.656197Z","shell.execute_reply.started":"2023-06-27T11:29:17.638447Z","shell.execute_reply":"2023-06-27T11:29:17.655106Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def feature_engineer_pl(x, grp, use_extra, feature_suffix):\n        \n    aggs = [\n        pl.col(\"index\").count().alias(f\"session_number_{feature_suffix}\"),\n      \n        *[pl.col(c).drop_nulls().n_unique().alias(f\"{c}_unique_{feature_suffix}\") for c in CATS],\n        *[pl.col(c).quantile(0.1, \"nearest\").alias(f\"{c}_quantile1_{feature_suffix}\") for c in NUMS],\n        *[pl.col(c).quantile(0.2, \"nearest\").alias(f\"{c}_quantile2_{feature_suffix}\") for c in NUMS],\n        *[pl.col(c).quantile(0.4, \"nearest\").alias(f\"{c}_quantile4_{feature_suffix}\") for c in NUMS],\n        *[pl.col(c).quantile(0.6, \"nearest\").alias(f\"{c}_quantile6_{feature_suffix}\") for c in NUMS],\n        *[pl.col(c).quantile(0.8, \"nearest\").alias(f\"{c}_quantile8_{feature_suffix}\") for c in NUMS],\n        *[pl.col(c).quantile(0.9, \"nearest\").alias(f\"{c}_quantile9_{feature_suffix}\") for c in NUMS],\n        \n        *[pl.col(c).mean().alias(f\"{c}_mean_{feature_suffix}\") for c in NUMS],\n        *[pl.col(c).std().alias(f\"{c}_std_{feature_suffix}\") for c in NUMS],\n        *[pl.col(c).min().alias(f\"{c}_min_{feature_suffix}\") for c in NUMS],\n        *[pl.col(c).max().alias(f\"{c}_max_{feature_suffix}\") for c in NUMS],\n        \n        *[pl.col(\"event_name\").filter(pl.col(\"event_name\") == c).count().alias(f\"{c}_event_name_counts{feature_suffix}\")for c in event_name_feature],\n        *[pl.col(\"elapsed_time_diff\").filter(pl.col(\"event_name\")==c).quantile(0.1, \"nearest\").alias(f\"{c}_ET_quantile1_{feature_suffix}\") for c in event_name_feature],\n        *[pl.col(\"elapsed_time_diff\").filter(pl.col(\"event_name\")==c).quantile(0.2, \"nearest\").alias(f\"{c}_ET_quantile2_{feature_suffix}\") for c in event_name_feature],\n        *[pl.col(\"elapsed_time_diff\").filter(pl.col(\"event_name\")==c).quantile(0.4, \"nearest\").alias(f\"{c}_ET_quantile4_{feature_suffix}\") for c in event_name_feature],\n        *[pl.col(\"elapsed_time_diff\").filter(pl.col(\"event_name\")==c).quantile(0.6, \"nearest\").alias(f\"{c}_ET_quantile6_{feature_suffix}\") for c in event_name_feature],\n        *[pl.col(\"elapsed_time_diff\").filter(pl.col(\"event_name\")==c).quantile(0.8, \"nearest\").alias(f\"{c}_ET_quantile8_{feature_suffix}\") for c in event_name_feature],\n        *[pl.col(\"elapsed_time_diff\").filter(pl.col(\"event_name\")==c).quantile(0.9, \"nearest\").alias(f\"{c}_ET_quantile9_{feature_suffix}\") for c in event_name_feature],      \n        *[pl.col(\"elapsed_time_diff\").filter(pl.col(\"event_name\")==c).mean().alias(f\"{c}_ET_mean_{feature_suffix}\") for c in event_name_feature],\n        *[pl.col(\"elapsed_time_diff\").filter(pl.col(\"event_name\")==c).std().alias(f\"{c}_ET_std_{feature_suffix}\") for c in event_name_feature],\n        *[pl.col(\"elapsed_time_diff\").filter(pl.col(\"event_name\")==c).max().alias(f\"{c}_ET_max_{feature_suffix}\") for c in event_name_feature],\n        *[pl.col(\"elapsed_time_diff\").filter(pl.col(\"event_name\")==c).min().alias(f\"{c}_ET_min_{feature_suffix}\") for c in event_name_feature],\n     \n        *[pl.col(\"name\").filter(pl.col(\"name\") == c).count().alias(f\"{c}_name_counts{feature_suffix}\")for c in name_feature],   \n        *[pl.col(\"elapsed_time_diff\").filter(pl.col(\"name\")==c).mean().alias(f\"{c}_ET_mean_{feature_suffix}\") for c in name_feature],\n        *[pl.col(\"elapsed_time_diff\").filter(pl.col(\"name\")==c).max().alias(f\"{c}_ET_max_{feature_suffix}\") for c in name_feature],\n        *[pl.col(\"elapsed_time_diff\").filter(pl.col(\"name\")==c).min().alias(f\"{c}_ET_min_{feature_suffix}\") for c in name_feature],\n        *[pl.col(\"elapsed_time_diff\").filter(pl.col(\"name\")==c).std().alias(f\"{c}_ET_std_{feature_suffix}\") for c in name_feature],  \n        \n        *[pl.col(\"room_fqid\").filter(pl.col(\"room_fqid\") == c).count().alias(f\"{c}_room_fqid_counts{feature_suffix}\")for c in room_lists],\n        *[pl.col(\"elapsed_time_diff\").filter(pl.col(\"room_fqid\") == c).std().alias(f\"{c}_ET_std_{feature_suffix}\") for c in room_lists],\n        *[pl.col(\"elapsed_time_diff\").filter(pl.col(\"room_fqid\") == c).mean().alias(f\"{c}_ET_mean_{feature_suffix}\") for c in room_lists],\n        *[pl.col(\"elapsed_time_diff\").filter(pl.col(\"room_fqid\") == c).max().alias(f\"{c}_ET_max_{feature_suffix}\") for c in room_lists],\n        *[pl.col(\"elapsed_time_diff\").filter(pl.col(\"room_fqid\") == c).min().alias(f\"{c}_ET_min_{feature_suffix}\") for c in room_lists],\n        *[pl.col(\"elapsed_time_diff\").filter(pl.col(\"room_fqid\") == c).sum().alias(f\"{c}_ET_sum_{feature_suffix}\") for c in room_lists],\n                \n        *[pl.col(\"fqid\").filter(pl.col(\"fqid\") == c).count().alias(f\"{c}_fqid_counts{feature_suffix}\")for c in fqid_lists],\n        *[pl.col(\"elapsed_time_diff\").filter(pl.col(\"fqid\") == c).std().alias(f\"{c}_ET_std_{feature_suffix}\") for c in fqid_lists],\n        *[pl.col(\"elapsed_time_diff\").filter(pl.col(\"fqid\") == c).mean().alias(f\"{c}_ET_mean_{feature_suffix}\") for c in fqid_lists],\n        *[pl.col(\"elapsed_time_diff\").filter(pl.col(\"fqid\") == c).max().alias(f\"{c}_ET_max_{feature_suffix}\") for c in fqid_lists],\n        *[pl.col(\"elapsed_time_diff\").filter(pl.col(\"fqid\") == c).min().alias(f\"{c}_ET_min_{feature_suffix}\") for c in fqid_lists],\n        *[pl.col(\"elapsed_time_diff\").filter(pl.col(\"fqid\") == c).sum().alias(f\"{c}_ET_sum_{feature_suffix}\") for c in fqid_lists],\n       \n        *[pl.col(\"text_fqid\").filter(pl.col(\"text_fqid\") == c).count().alias(f\"{c}_text_fqid_counts{feature_suffix}\") for c in text_lists],\n        *[pl.col(\"elapsed_time_diff\").filter(pl.col(\"text_fqid\") == c).std().alias(f\"{c}_ET_std_{feature_suffix}\") for c in text_lists],\n        *[pl.col(\"elapsed_time_diff\").filter(pl.col(\"text_fqid\") == c).mean().alias(f\"{c}_ET_mean_{feature_suffix}\") for c in text_lists],\n        *[pl.col(\"elapsed_time_diff\").filter(pl.col(\"text_fqid\") == c).max().alias(f\"{c}_ET_max_{feature_suffix}\") for c in text_lists],\n        *[pl.col(\"elapsed_time_diff\").filter(pl.col(\"text_fqid\") == c).min().alias(f\"{c}_ET_min_{feature_suffix}\") for c in text_lists],\n        *[pl.col(\"elapsed_time_diff\").filter(pl.col(\"text_fqid\") == c).sum().alias(f\"{c}_ET_sum_{feature_suffix}\") for c in text_lists],\n         \n        *[pl.col(\"location_x_diff\").filter(pl.col(\"event_name\")==c).mean().alias(f\"{c}_ET_mean_x{feature_suffix}\") for c in event_name_feature],\n        *[pl.col(\"location_x_diff\").filter(pl.col(\"event_name\")==c).std().alias(f\"{c}_ET_std_x{feature_suffix}\") for c in event_name_feature],\n        *[pl.col(\"location_x_diff\").filter(pl.col(\"event_name\")==c).max().alias(f\"{c}_ET_max_x{feature_suffix}\") for c in event_name_feature],\n        *[pl.col(\"location_x_diff\").filter(pl.col(\"event_name\")==c).min().alias(f\"{c}_ET_min_x{feature_suffix}\") for c in event_name_feature],\n        ]\n    \n    df = x.groupby([\"session_id\"], maintain_order=True).agg(aggs).sort(\"session_id\")\n  \n    if use_extra:\n        if grp=='5-12':\n            aggs = [\n                pl.col(\"elapsed_time\").filter((pl.col(\"text\")==\"Here's the log book.\")|(pl.col(\"fqid\")=='logbook.page.bingo')).apply(lambda s: s.max()-s.min()).alias(\"logbook_bingo_duration\"),\n                pl.col(\"index\").filter((pl.col(\"text\")==\"Here's the log book.\")|(pl.col(\"fqid\")=='logbook.page.bingo')).apply(lambda s: s.max()-s.min()).alias(\"logbook_bingo_indexCount\"),\n                pl.col(\"elapsed_time\").filter(((pl.col(\"event_name\")=='navigate_click')&(pl.col(\"fqid\")=='reader'))|(pl.col(\"fqid\")==\"reader.paper2.bingo\")).apply(lambda s: s.max()-s.min()).alias(\"reader_bingo_duration\"),\n                pl.col(\"index\").filter(((pl.col(\"event_name\")=='navigate_click')&(pl.col(\"fqid\")=='reader'))|(pl.col(\"fqid\")==\"reader.paper2.bingo\")).apply(lambda s: s.max()-s.min()).alias(\"reader_bingo_indexCount\"),\n                pl.col(\"elapsed_time\").filter(((pl.col(\"event_name\")=='navigate_click')&(pl.col(\"fqid\")=='journals'))|(pl.col(\"fqid\")==\"journals.pic_2.bingo\")).apply(lambda s: s.max()-s.min()).alias(\"journals_bingo_duration\"),\n                pl.col(\"index\").filter(((pl.col(\"event_name\")=='navigate_click')&(pl.col(\"fqid\")=='journals'))|(pl.col(\"fqid\")==\"journals.pic_2.bingo\")).apply(lambda s: s.max()-s.min()).alias(\"journals_bingo_indexCount\"),\n            ]\n            tmp = x.groupby([\"session_id\"], maintain_order=True).agg(aggs).sort(\"session_id\")\n            df = df.join(tmp, on=\"session_id\", how='left')\n\n        if grp=='13-22':\n            aggs = [\n                pl.col(\"elapsed_time\").filter(((pl.col(\"event_name\")=='navigate_click')&(pl.col(\"fqid\")=='reader_flag'))|(pl.col(\"fqid\")==\"tunic.library.microfiche.reader_flag.paper2.bingo\")).apply(lambda s: s.max()-s.min() if s.len()>0 else 0).alias(\"reader_flag_duration\"),\n                pl.col(\"index\").filter(((pl.col(\"event_name\")=='navigate_click')&(pl.col(\"fqid\")=='reader_flag'))|(pl.col(\"fqid\")==\"tunic.library.microfiche.reader_flag.paper2.bingo\")).apply(lambda s: s.max()-s.min() if s.len()>0 else 0).alias(\"reader_flag_indexCount\"),\n                pl.col(\"elapsed_time\").filter(((pl.col(\"event_name\")=='navigate_click')&(pl.col(\"fqid\")=='journals_flag'))|(pl.col(\"fqid\")==\"journals_flag.pic_0.bingo\")).apply(lambda s: s.max()-s.min() if s.len()>0 else 0).alias(\"journalsFlag_bingo_duration\"),\n                pl.col(\"index\").filter(((pl.col(\"event_name\")=='navigate_click')&(pl.col(\"fqid\")=='journals_flag'))|(pl.col(\"fqid\")==\"journals_flag.pic_0.bingo\")).apply(lambda s: s.max()-s.min() if s.len()>0 else 0).alias(\"journalsFlag_bingo_indexCount\"),\n            ]\n            tmp = x.groupby([\"session_id\"], maintain_order=True).agg(aggs).sort(\"session_id\")\n            df = df.join(tmp, on=\"session_id\", how='left')\n    df = df.to_pandas()\n    if grp == '5-12':\n        df = df[col512[1:-10]]\n    elif grp == '13-22':\n        df = df[col1322[1:-5]]\n    else: \n        df = df[col04[1:-3]]\n        \n    for col in df.columns:\n        if df[col].isnull().sum() > 0:\n            df[col].fillna(-1, inplace=True)\n        \n    return df","metadata":{"execution":{"iopub.status.busy":"2023-06-27T11:29:18.096038Z","iopub.execute_input":"2023-06-27T11:29:18.096448Z","iopub.status.idle":"2023-06-27T11:29:18.330610Z","shell.execute_reply.started":"2023-06-27T11:29:18.096415Z","shell.execute_reply":"2023-06-27T11:29:18.329365Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_cont_values(dataframe, cont_count=10, target=10):\n    cont = dataframe.iloc[:, -cont_count:]\n    return cont.to_numpy()\n\ndef get_categ_values(dataframe, cont_count=10):\n    categ = dataframe.iloc[:, 0:-cont_count]\n\n    for i in range(categ.shape[1]):\n        categ.loc[:, categ.columns[i]] = categ[categ.columns[i]].astype(\"category\").cat.codes\n\n    return categ.to_numpy()\n\ndef get_categ_cont_values(dataframe, positive_class=1):\n    cate_counter = 0\n    for col in dataframe.columns:\n        if 'unique' in col:\n            cate_counter += 1\n    cont_count = len(dataframe.columns) - cate_counter\n\n    cont = get_cont_values(dataframe, cont_count)\n    categ = get_categ_values(dataframe, cont_count)\n\n    return cont, categ","metadata":{"execution":{"iopub.status.busy":"2023-06-27T11:29:18.493486Z","iopub.execute_input":"2023-06-27T11:29:18.493868Z","iopub.status.idle":"2023-06-27T11:29:18.502008Z","shell.execute_reply.started":"2023-06-27T11:29:18.493838Z","shell.execute_reply":"2023-06-27T11:29:18.501117Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def validate(model, cont_data, categ_data, target_dim, device=\"cpu\", val_batch_size=1, coefficients=None):\n    model = model.eval()\n    results = []\n\n    for i in range(categ_data.shape[0] // val_batch_size):\n        x_categ = torch.tensor(categ_data[val_batch_size*i:val_batch_size*i+val_batch_size]).to(dtype=torch.int64, device=device)\n        x_cont = torch.tensor(cont_data[val_batch_size*i:val_batch_size*i+val_batch_size]).to(dtype=torch.float32, device=device)\n\n        pred = model(x_categ, x_cont)\n        \n        results.append(torch.sigmoid(pred).squeeze().cpu().detach().numpy())\n    \n    results = np.array(results)    \n#     print('Target dim', target_dim)\n#     print(results)\n    for i, coefficient in enumerate(coefficients):\n        results[:, i] = (results[:, i] >= coefficient).astype(np.int64)\n    results = results.astype(np.int64)\n#     print(results)\n    return results","metadata":{"execution":{"iopub.status.busy":"2023-06-27T11:29:18.881758Z","iopub.execute_input":"2023-06-27T11:29:18.882917Z","iopub.status.idle":"2023-06-27T11:29:18.893077Z","shell.execute_reply.started":"2023-06-27T11:29:18.882876Z","shell.execute_reply":"2023-06-27T11:29:18.891738Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def infer(tf, df, grp, target_dim, coef):\n    columns = [\n        pl.col(\"page\").cast(pl.Float32),\n        (\n            (pl.col(\"elapsed_time\") - pl.col(\"elapsed_time\").shift(1)) \n             .fill_null(0)\n             .clip(0, 1e9)\n             .over([\"session_id\", \"level_group\"])\n             .alias(\"elapsed_time_diff\")\n        ),\n        (\n            (pl.col(\"screen_coor_x\") - pl.col(\"screen_coor_x\").shift(1)) \n             .abs()\n             .over([\"session_id\", \"level_group\"])\n            .alias(\"location_x_diff\") \n        ),\n        (\n            (pl.col(\"screen_coor_y\") - pl.col(\"screen_coor_y\").shift(1)) \n             .abs()\n             .over([\"session_id\", \"level_group\"])\n            .alias(\"location_y_diff\") \n        ),\n        pl.col(\"fqid\").fill_null(\"fqid_None\"),\n        pl.col(\"text_fqid\").fill_null(\"text_fqid_None\")\n    ]\n    \n    df = (pl.from_pandas(df)\n          .drop([\"fullscreen\", \"hq\", \"music\"])\n          .with_columns(columns))\n    \n    df = feature_engineer_pl(df, grp, use_extra=True, feature_suffix='')\n    \n    cont, categ = get_categ_cont_values(df)\n    \n    results = validate(tf, cont, categ, target_dim, coefficients = coef)\n    \n    return results\n    ","metadata":{"execution":{"iopub.status.busy":"2023-06-27T11:29:19.389600Z","iopub.execute_input":"2023-06-27T11:29:19.389967Z","iopub.status.idle":"2023-06-27T11:29:19.399245Z","shell.execute_reply.started":"2023-06-27T11:29:19.389937Z","shell.execute_reply":"2023-06-27T11:29:19.398267Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"coef04 = [0.741252,   0.39964452, 0.37312032]\ncoef512 = [0.50184932, 0.62682469, 0.55683839, 0.5986551,  0.57346862, 0.5688148,\n 0.65741359, 0.61575163, 0.33960375, 0.54991444]\ncoef1322 = [0.36813991, 0.36142655, 0.51881254, 0.50101977, 0.66000616]\n\ndef main(df, lv):\n    if lv == '0-4':\n        results = infer(tf04, df, lv, 3, coef04)\n    elif lv == '5-12':\n        results = infer(tf512, df, lv, 10, coef512) \n    else:\n        results = infer(tf1322, df, lv, 5, coef1322)\n    \n    return results","metadata":{"execution":{"iopub.status.busy":"2023-06-27T11:29:20.324460Z","iopub.execute_input":"2023-06-27T11:29:20.325100Z","iopub.status.idle":"2023-06-27T11:29:20.331123Z","shell.execute_reply.started":"2023-06-27T11:29:20.325066Z","shell.execute_reply":"2023-06-27T11:29:20.330225Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submit","metadata":{}},{"cell_type":"code","source":"import jo_wilder_310\ntry:\n    jo_wilder_310.make_env.__called__ = False\n    env.__called__ = False\n    type(env)._state = type(type(env)._state).__dict__['INIT']\nexcept:\n    pass\n\nenv = jo_wilder_310.make_env()\niter_test = env.iter_test() \n\n# gc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-06-27T11:29:21.265355Z","iopub.execute_input":"2023-06-27T11:29:21.266106Z","iopub.status.idle":"2023-06-27T11:29:21.313452Z","shell.execute_reply.started":"2023-06-27T11:29:21.266068Z","shell.execute_reply":"2023-06-27T11:29:21.312242Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"limits = {'0-4':(1,4), '5-12':(4,14), '13-22':(14,19)}\n\nfor (test, sample_submission) in iter_test:\n    try:\n    #     INFER TEST DATA\n        test = test.sort_values('index').reset_index(drop=True)\n\n        sample_submission['question'] = [int(label.split('_')[1][1:]) for label in sample_submission['session_id']]\n        sample_submission = sample_submission.sort_values('question').reset_index(drop=True)\n\n        grp = test.level_group.values[0]\n        a,b = limits[grp]\n        \n        output = main(test, grp)\n        for i,t in enumerate(range(a,b)):\n            mask = sample_submission.session_id.str.contains(f'q{t}')\n            sample_submission.loc[mask,'correct'] = int(output[0][i])\n\n    finally:\n        env.predict(sample_submission[['session_id', 'correct']])","metadata":{"execution":{"iopub.status.busy":"2023-06-27T11:29:21.843857Z","iopub.execute_input":"2023-06-27T11:29:21.844656Z","iopub.status.idle":"2023-06-27T11:29:25.601153Z","shell.execute_reply.started":"2023-06-27T11:29:21.844622Z","shell.execute_reply":"2023-06-27T11:29:25.598704Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv('submission.csv')\nprint( df.shape )\ndf.head(50)","metadata":{"execution":{"iopub.status.busy":"2023-06-27T11:29:31.641761Z","iopub.execute_input":"2023-06-27T11:29:31.642210Z","iopub.status.idle":"2023-06-27T11:29:31.673245Z","shell.execute_reply.started":"2023-06-27T11:29:31.642174Z","shell.execute_reply":"2023-06-27T11:29:31.672363Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}