{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":58266,"databundleVersionId":6641124,"sourceType":"competition"},{"sourceId":6433176,"sourceType":"datasetVersion","datasetId":3712282},{"sourceId":6434726,"sourceType":"datasetVersion","datasetId":3713282},{"sourceId":6434849,"sourceType":"datasetVersion","datasetId":3713371},{"sourceId":6438594,"sourceType":"datasetVersion","datasetId":3715784},{"sourceId":6439679,"sourceType":"datasetVersion","datasetId":3716464},{"sourceId":6441547,"sourceType":"datasetVersion","datasetId":3717788},{"sourceId":6444900,"sourceType":"datasetVersion","datasetId":3720080},{"sourceId":6573856,"sourceType":"datasetVersion","datasetId":3796859},{"sourceId":6578151,"sourceType":"datasetVersion","datasetId":3798685},{"sourceId":6630935,"sourceType":"datasetVersion","datasetId":3827878},{"sourceId":144043388,"sourceType":"kernelVersion"},{"sourceId":144045966,"sourceType":"kernelVersion"}],"dockerImageVersionId":30559,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nfrom pathlib import Path\nfrom typing import Dict, Optional, List, Union, Tuple\nfrom dataclasses import dataclass\nimport math\nimport numpy as np\nimport pandas as pd\nfrom datasets import Dataset\nfrom tqdm import tqdm\nimport torch\nfrom torch import nn\nfrom torch.nn import functional as F\nfrom torch.nn.utils.rnn import pad_sequence\nfrom torch.utils.data import  DataLoader\n\nfrom transformers.modeling_outputs import BaseModelOutputWithPastAndCrossAttentions\nfrom transformers.pytorch_utils import apply_chunking_to_forward\nfrom transformers.activations import ACT2FN\nimport pytorch_lightning as pl\nimport torchmetrics as tm\n# import bitsandbytes as bnb","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T05:12:44.17944Z","iopub.execute_input":"2025-11-12T05:12:44.180243Z","iopub.status.idle":"2025-11-12T05:12:57.64941Z","shell.execute_reply.started":"2025-11-12T05:12:44.180213Z","shell.execute_reply":"2025-11-12T05:12:57.648721Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"NODE_OP_CODES = 120\nNODE_FEATS = 140\nCONFIG_FEATS = 24\nNODE_CONFIG_FEATS = 18","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T05:12:57.650831Z","iopub.execute_input":"2025-11-12T05:12:57.651079Z","iopub.status.idle":"2025-11-12T05:12:57.655085Z","shell.execute_reply.started":"2025-11-12T05:12:57.651058Z","shell.execute_reply":"2025-11-12T05:12:57.654259Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# TILE","metadata":{}},{"cell_type":"code","source":"DATA_DIR = \"../input/predict-ai-model-runtime/npz_all/npz\"\n\n\ndef generate_tile_df() -> pd.DataFrame:\n    tile_df = pd.DataFrame({'paths': [elem for elem in (Path(DATA_DIR) / 'tile').rglob(\"*\") if elem.is_file()]}).assign(\n        split=lambda df: df.paths.apply(lambda x: x.parent.name),\n        configuration=lambda df: df.paths.apply(lambda x: x.parent.parent.name),\n        extra=lambda df: df.paths.apply(lambda x: x.parent.parent.parent.name),\n        model_name=lambda df: df.paths.apply(lambda x: x.stem),\n        collection=lambda df: df.extra + ':' + df.configuration ,\n        ID=lambda df: df.collection + ':' + df.model_name ,\n        paths = lambda df: df.paths.apply(lambda x: str(x))\n    )\n    return tile_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T05:12:57.65618Z","iopub.execute_input":"2025-11-12T05:12:57.656725Z","iopub.status.idle":"2025-11-12T05:12:57.667536Z","shell.execute_reply.started":"2025-11-12T05:12:57.656694Z","shell.execute_reply":"2025-11-12T05:12:57.666649Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"tile_df = generate_tile_df()\ntile_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T05:12:57.669413Z","iopub.execute_input":"2025-11-12T05:12:57.669665Z","iopub.status.idle":"2025-11-12T05:13:20.687774Z","shell.execute_reply.started":"2025-11-12T05:12:57.669645Z","shell.execute_reply":"2025-11-12T05:13:20.686899Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataset\n* Create an Adjacency matrix for masking the attention\n* Creates a virtual first node equivalent to the [CLS] token which contains the global config for tile cases, while layout node configuration goes to the corresponding node position","metadata":{}},{"cell_type":"code","source":"def edges_adjacency(edges: torch.Tensor, add_diagonal=True) -> torch.Tensor:\n    \"\"\"\n    Generate an adjacency matrix from the edges\n    Args:\n        edges: Tensor of shape (num_edges, 2) with the edges\n        add_diagonal: Boolean indicating if the diagonal should be added to the adjacency matrix\n    Returns:\n        adjacency_matrix: Tensor of shape (num_nodes, num_nodes) with the adjacency matrix\n    \"\"\"\n    adjacency_matrix = torch.zeros((edges.max() + 1, edges.max() + 1))\n    adjacency_matrix[edges[:, 0], edges[:, 1]] = 1\n    if add_diagonal:\n        diag_idx = torch.arange(adjacency_matrix.shape[0])\n        adjacency_matrix[diag_idx, diag_idx] = 1\n    return adjacency_matrix\n\ndef tile_loader(path):\n    tile_dict =  dict(np.load(path))\n    tile_dict = {k: torch.from_numpy(v) for k, v in tile_dict.items()}\n    tile_dict['edges_adjecency'] = edges_adjacency(tile_dict['edge_index'])\n    return tile_dict\n\ndef node_cls_token(elem_dict, shift_node_config_ids:bool=True):\n    \"\"\"\n    Add a cls token to the node opcode, features, edges adjacency matrix, shift node_config_ids by 1 to account for the cls token\n    Args:\n        elem_dict: Dictionary with the elements of the tile\n    Returns:\n        elem_dict: Dictionary with the elements of the tile with the cls token\n    \"\"\"\n    elem_dict['node_opcode'] = torch.cat([torch.tensor([0]), elem_dict['node_opcode']]) # Introduce [CLS] node\n    elem_dict['node_feat'] = torch.cat([torch.zeros((1, elem_dict['node_feat'].shape[1])), elem_dict['node_feat']])\n    elem_dict['edges_adjecency'] = F.pad(elem_dict['edges_adjecency'], (1,0,1,0), value=1)\n    if 'node_config_ids' in elem_dict and shift_node_config_ids:\n        elem_dict['node_config_ids'] = elem_dict['node_config_ids'] + 1 # Shift Node Config IDs to take in to account [CLS] node\n    return elem_dict\n\n\nclass TileDataset(torch.utils.data.Dataset):\n    \n    def __init__(self, df:pd.DataFrame ,add_cls_token:bool=True, num_configs:int=10,  max_configs:Optional[int]=None):\n        self.df = df\n        self.add_cls_token = add_cls_token\n        self.num_configs = num_configs\n        self.max_configs = max_configs  \n        \n    def __len__(self) -> int:\n        return len(self.df)\n    \n    def select_configs(self, total_configs:int):\n        if self.max_configs is not None:\n            total_configs = min(total_configs, self.max_configs)\n        if self.num_configs == -1:\n            return np.arange(total_configs)\n        if total_configs < self.num_configs:\n            return np.random.choice(total_configs, self.num_configs, replace=True)\n        return  np.random.choice(total_configs, self.num_configs, replace=False)\n    \n    def __getitem__(self, idx:int, selected_configs:List[int]=None):\n        tile_dict = tile_loader(self.df.paths[idx])\n        if selected_configs is None:\n            selected_configs = self.select_configs(tile_dict['config_feat'].shape[0])\n        tile_dict['node_config_feat'] = tile_dict.pop('config_feat')[selected_configs]\n        tile_dict['node_config_feat'] = F.pad(tile_dict['node_config_feat'].unsqueeze(1), (0,NODE_CONFIG_FEATS))\n        tile_dict['config_runtime'] = tile_dict['config_runtime'][selected_configs].float()\n        tile_dict['config_runtime'] /= tile_dict['config_runtime_normalizers'][selected_configs].float()\n        tile_dict['node_config_ids'] = torch.zeros((1,))\n        tile_dict['selected_idxs'] = selected_configs\n        if self.add_cls_token:\n            tile_dict = node_cls_token(tile_dict, False)\n        return tile_dict","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T05:13:20.689168Z","iopub.execute_input":"2025-11-12T05:13:20.689443Z","iopub.status.idle":"2025-11-12T05:13:20.70256Z","shell.execute_reply.started":"2025-11-12T05:13:20.68942Z","shell.execute_reply":"2025-11-12T05:13:20.701534Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"tile_dataset = TileDataset(tile_df)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T05:13:20.703651Z","iopub.execute_input":"2025-11-12T05:13:20.70391Z","iopub.status.idle":"2025-11-12T05:13:20.72458Z","shell.execute_reply.started":"2025-11-12T05:13:20.703882Z","shell.execute_reply":"2025-11-12T05:13:20.72375Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"elem = tile_dataset[0]\nfor k,v in elem.items():\n    print(k, v.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T05:13:20.7266Z","iopub.execute_input":"2025-11-12T05:13:20.726872Z","iopub.status.idle":"2025-11-12T05:13:20.807165Z","shell.execute_reply.started":"2025-11-12T05:13:20.726845Z","shell.execute_reply":"2025-11-12T05:13:20.806286Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"elem['edges_adjecency']","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T05:13:20.808002Z","iopub.execute_input":"2025-11-12T05:13:20.80824Z","iopub.status.idle":"2025-11-12T05:13:20.827904Z","shell.execute_reply.started":"2025-11-12T05:13:20.80822Z","shell.execute_reply":"2025-11-12T05:13:20.827144Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Collator","metadata":{}},{"cell_type":"code","source":"def pad_edge_adjacency(edges_adjacency_list):\n    max_len = max([elem.shape[0] for elem in edges_adjacency_list])\n    return torch.stack([F.pad(elem, (0, max_len-elem.shape[0], 0, max_len-elem.shape[0]), value=0) for elem in edges_adjacency_list], dim=0)\n\n@dataclass\nclass LayoutCollator:\n    pad_to_multiple_of: int = 64\n    targets:bool = True\n    padding_idx:int = 120\n    node_padding_idx:int = 0\n    \n    def __call__(self, batch):\n        output = {}\n        max_node_len = max([elem['node_opcode'].shape[0] for elem in batch])\n        node_pad_amount = self.pad_to_multiple_of - max_node_len % max(self.pad_to_multiple_of, 1)\n        output['node_opcode'] = F.pad(pad_sequence([elem['node_opcode'] for elem in batch], batch_first=True, padding_value=self.padding_idx),\n                                      (0, node_pad_amount), value=self.padding_idx).long()\n        output['node_feat'] = F.pad(pad_sequence([elem['node_feat'] for elem in batch], batch_first=True),\n                                    (0,0,0, node_pad_amount), value=0)\n        output['edges_adjecency'] = F.pad(pad_edge_adjacency([elem['edges_adjecency'] for elem in batch]),\n                                          (0, node_pad_amount, 0, node_pad_amount), value=0)\n        output['node_attn_mask'] = F.pad(pad_sequence([torch.ones(len(elem['node_opcode'])) for elem in batch], batch_first=True),\n                                         (0, node_pad_amount), value=0)\n\n        max_node_config_len = max([elem['node_config_ids'].shape[0] for elem in batch])\n        node_config_pad_amount = self.pad_to_multiple_of - max_node_config_len % max(self.pad_to_multiple_of, 1)\n        output['node_config_ids'] = F.pad(pad_sequence([elem['node_config_ids'] for elem in batch], batch_first=True),\n                                         (0, node_config_pad_amount), value=0).long()\n        padded_node_config_feat = pad_sequence([elem['node_config_feat'].permute(1,0,2) for elem in batch], batch_first=True, padding_value=-1)\n        padded_node_config_feat = F.pad(padded_node_config_feat.permute(0,2,1,3),\n                                           (0,0,0, node_config_pad_amount,0,0), value=-1)\n        \n        output['node_config_feat'] = torch.where(padded_node_config_feat!=-1, padded_node_config_feat, self.node_padding_idx)\n                                      \n        output['config_idxs'] = torch.stack([torch.from_numpy(elem['selected_idxs']) for elem in batch])\n        \n        if self.targets:\n            output['config_runtime'] = pad_sequence([elem['config_runtime'].float() for elem in batch], batch_first=True)\n        return output","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T05:13:20.828817Z","iopub.execute_input":"2025-11-12T05:13:20.829044Z","iopub.status.idle":"2025-11-12T05:13:20.83993Z","shell.execute_reply.started":"2025-11-12T05:13:20.829025Z","shell.execute_reply":"2025-11-12T05:13:20.839059Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"collate_fn = LayoutCollator(64)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T05:13:20.840774Z","iopub.execute_input":"2025-11-12T05:13:20.840979Z","iopub.status.idle":"2025-11-12T05:13:20.852638Z","shell.execute_reply.started":"2025-11-12T05:13:20.840962Z","shell.execute_reply":"2025-11-12T05:13:20.851973Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"batch = collate_fn([tile_dataset[0], tile_dataset[1]])\nfor k,v in batch.items():\n    print(k,v.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T05:13:20.853703Z","iopub.execute_input":"2025-11-12T05:13:20.854301Z","iopub.status.idle":"2025-11-12T05:13:20.88184Z","shell.execute_reply.started":"2025-11-12T05:13:20.854272Z","shell.execute_reply":"2025-11-12T05:13:20.880928Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model - Config","metadata":{}},{"cell_type":"code","source":"@dataclass\nclass GraphConfig:\n    num_hidden_layers: int = 8\n    hidden_size: int = 256\n    num_attention_heads: int = 16\n    intermediate_size: int = 64\n    chunk_size_feed_forward: int = 64\n    attention_probs_dropout_prob: float = 0.0\n    max_position_embeddings: int = 512\n    hidden_dropout_prob: float = 0.0\n    layer_norm_eps: float = 1e-12\n    hidden_act: str = 'gelu'\n    initializer_range: float = 0.02\n    output_hidden_states: bool = False\n    output_attentions: bool = False\n    gradient_checkpointing: bool = False\n    margin: float = 0.1\n    number_permutations: int = 10\n    \n    def __post_init__(self):\n        self.embedding_size = self.hidden_size\n    \n    def validate(self):\n        if self.hidden_size % self.num_attention_heads != 0 and not hasattr(self, \"embedding_size\"):\n            raise ValueError(\n                f\"The hidden size ({self.hidden_size}) is not a multiple of the number of attention \"\n                f\"heads ({self.num_attention_heads})\"\n            )\n            \n    def save_config(self, path):\n        config = asdict(self)\n        with open(path, 'w') as f:\n            json.dump(config, f)\n            \n    @classmethod\n    def load_config(cls, path):\n        with open(path, 'r') as f:\n            config = json.load(f)\n        return cls(**config)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T05:13:20.883164Z","iopub.execute_input":"2025-11-12T05:13:20.883534Z","iopub.status.idle":"2025-11-12T05:13:20.894354Z","shell.execute_reply.started":"2025-11-12T05:13:20.883502Z","shell.execute_reply":"2025-11-12T05:13:20.893365Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Loss\n* Uses Ranking loss to compare different configuration\n* Compares does configurations with different indexes, masks those cases where the permutation returns the same element\n* Compares multiple configurations in each run","metadata":{}},{"cell_type":"code","source":"class MultiElementRankLoss(nn.Module):\n    \"\"\"\n    Loss function that compares the output of the model with the output of the model with a permutation of the elements\n    \"\"\"\n    \n    def __init__(self, margin:float=0.0, number_permutations:int = 1) -> None:\n        super().__init__()\n        self.loss_fn = torch.nn.MarginRankingLoss(margin=margin, reduction = 'none')\n        self.number_permutations = number_permutations\n    \n    def calculate_rank_loss(self,\n                            outputs: torch.Tensor,\n                            config_runtime: torch.Tensor,\n                            config_idxs: torch.Tensor\n                            ):\n        \"\"\"\n        Generates a permutation of the predictions and targets and calculates the loss MarginRankingLoss against the permutation\n        Args:\n            outputs: Tensor of shape (bs, seq_len) with the outputs of the model\n            config_runtime: Tensor of shape (bs, seq_len) with the runtime of the model\n            config_mask: Tensor of shape (bs, seq_len) with 1 in the positions of the elements\n            and 0 in the positions of the padding\n        Returns:\n            loss: Tensor of shape (bs, seq_len) with the loss for each element in the batch\n        \"\"\"\n        bs, num_configs = outputs.shape\n        permutation = torch.randperm(num_configs) \n        permuted_idxs = config_idxs[:, permutation]\n        # We mask those cases where we compare the same configuration\n        config_mask = torch.where(config_idxs != permuted_idxs, 1, 0)\n        permuted_runtime = config_runtime[:, permutation]\n        labels = 2*((config_runtime - permuted_runtime) > 0) -1\n        permuted_output = outputs[:, permutation]\n        loss = self.loss_fn(outputs.view(-1,1), permuted_output.view(-1,1), labels.view(-1,1))\n        loss = loss.view(bs, num_configs) * config_mask\n        return loss.mean()\n                \n    \n    def forward(self,\n                outputs: torch.Tensor,\n                config_runtime: torch.Tensor,\n                config_idxs: torch.Tensor\n                ):\n        loss = 0 \n        for _ in range(self.number_permutations):\n            loss += self.calculate_rank_loss(outputs, config_runtime, config_idxs)\n        return loss/ self.number_permutations","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T05:13:20.895681Z","iopub.execute_input":"2025-11-12T05:13:20.895993Z","iopub.status.idle":"2025-11-12T05:13:20.90607Z","shell.execute_reply.started":"2025-11-12T05:13:20.895965Z","shell.execute_reply":"2025-11-12T05:13:20.905267Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Metric","metadata":{}},{"cell_type":"code","source":"\nclass TileTopK(tm.Metric):\n    \n    higher_is_better = True\n    \n    def __init__(self, k:int=5) -> None:\n        super().__init__()\n        self.add_state(\"runtimes\", default=[], dist_reduce_fx=None)\n        self.k = k\n        \n    def update(self, preds: torch.Tensor, target: torch.Tensor, config_attn_mask:torch.Tensor) -> None:\n        \"\"\"\n        Update the metric state\n        Args:\n            preds: Tensor of shape (bs, seq_len) with the predicted runtimes orders\n            target: Tensor of shape (bs, seq_len) with the target runtimes\n            config_attn_mask: Tensor of shape (bs, seq_len) with 1 in the positions of the elements\n        \"\"\"\n        best_runtimes = torch.where(config_attn_mask==1, target, torch.tensor(float('inf'))).min(1).values\n        masked_preds = torch.where(config_attn_mask==1, preds, torch.tensor(float('inf')))\n        pred_bottomk_indices = torch.topk(masked_preds, k=self.k, largest=False).indices\n        bs = preds.shape[0]\n        bottom_k_positions = torch.stack([torch.arange(bs).repeat_interleave(self.k).to(config_attn_mask.device), pred_bottomk_indices.view(-1)])\n        predicted_runtimes = target[bottom_k_positions[0], bottom_k_positions[1]].view(bs,self.k)\n        best_predicted_runtimes = predicted_runtimes.min(1).values\n        self.runtimes.append(best_predicted_runtimes/ best_runtimes)\n        \n    def compute(self) -> torch.Tensor:\n        return (2-torch.cat(self.runtimes)).mean()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T05:13:20.906999Z","iopub.execute_input":"2025-11-12T05:13:20.907231Z","iopub.status.idle":"2025-11-12T05:13:20.920977Z","shell.execute_reply.started":"2025-11-12T05:13:20.907205Z","shell.execute_reply":"2025-11-12T05:13:20.920041Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model\nModified version of 🤗 Bert implementation to take in to account [Graph Attention](https://arxiv.org/abs/1710.10903)\n* Removed the parts corresponding to Cross-attention\n* Made layer_head_mask the same for all layers, heads\n* The Head mask corresponds to the edge adjacency \n","metadata":{}},{"cell_type":"code","source":"# Modified from https://github.com/huggingface/transformers/blob/main/src/transformers/models/bert/modeling_bert.py\nclass BertEncoder(nn.Module):\n    def __init__(self, config:GraphConfig):\n        super().__init__()\n        self.config = config\n        self.layer = nn.ModuleList([BertLayer(config) for _ in range(config.num_hidden_layers)])\n        self.gradient_checkpointing = False\n\n    def forward(\n        self,\n        hidden_states: torch.Tensor,\n        attention_mask: Optional[torch.FloatTensor] = None,\n        head_mask: Optional[torch.FloatTensor] = None,\n        output_attentions: Optional[bool] = False,\n        output_hidden_states: Optional[bool] = False,\n        return_dict: Optional[bool] = True,\n    ) -> Union[Tuple[torch.Tensor], BaseModelOutputWithPastAndCrossAttentions]:\n        all_hidden_states = () if output_hidden_states else None\n        all_self_attentions = () if output_attentions else None\n\n        for i, layer_module in enumerate(self.layer):\n            if output_hidden_states:\n                all_hidden_states = all_hidden_states + (hidden_states,)\n\n            layer_head_mask = head_mask #DONE: Same Head Mask for all layers\n\n            if self.gradient_checkpointing and self.training:\n\n                def create_custom_forward(module):\n                    def custom_forward(*inputs):\n                        return module(*inputs,  output_attentions)\n\n                    return custom_forward\n\n                layer_outputs = torch.utils.checkpoint.checkpoint(\n                    create_custom_forward(layer_module),\n                    hidden_states,\n                    attention_mask,\n                    layer_head_mask,\n                )\n            else:\n                layer_outputs = layer_module(\n                    hidden_states,\n                    attention_mask,\n                    layer_head_mask,\n                    output_attentions,\n                )\n\n            hidden_states = layer_outputs[0]\n            if output_attentions:\n                all_self_attentions = all_self_attentions + (layer_outputs[1],)\n\n        if output_hidden_states:\n            all_hidden_states = all_hidden_states + (hidden_states,)\n\n        if not return_dict:\n            return tuple(\n                v\n                for v in [\n                    hidden_states,\n                    all_hidden_states,\n                    all_self_attentions,\n                ]\n                if v is not None\n            )\n        return BaseModelOutputWithPastAndCrossAttentions(\n            last_hidden_state=hidden_states,\n            past_key_values=None,\n            hidden_states=all_hidden_states,\n            attentions=all_self_attentions,\n            cross_attentions=None,\n        )\n        \n        \nclass BertLayer(nn.Module):\n    def __init__(self, config:GraphConfig):\n        super().__init__()\n        self.chunk_size_feed_forward = config.chunk_size_feed_forward\n        self.seq_len_dim = 1\n        self.attention = BertAttention(config)\n        self.intermediate = BertIntermediate(config)\n        self.output = BertOutput(config)\n\n    def forward(\n        self,\n        hidden_states: torch.Tensor,\n        attention_mask: Optional[torch.FloatTensor] = None,\n        head_mask: Optional[torch.FloatTensor] = None,\n        output_attentions: Optional[bool] = False,\n    ) -> Tuple[torch.Tensor]:\n        # decoder uni-directional self-attention cached key/values tuple is at positions 1,2\n        self_attention_outputs = self.attention(\n            hidden_states,\n            attention_mask,\n            head_mask,\n            output_attentions=output_attentions,\n        )\n        attention_output = self_attention_outputs[0]\n        outputs = self_attention_outputs[1:]  # add self attentions if we output attention weights\n        layer_output = apply_chunking_to_forward(\n            self.feed_forward_chunk, self.chunk_size_feed_forward, self.seq_len_dim, attention_output\n        )\n        outputs = (layer_output,) + outputs\n\n\n        return outputs\n\n    def feed_forward_chunk(self, attention_output):\n        intermediate_output = self.intermediate(attention_output)\n        layer_output = self.output(intermediate_output, attention_output)\n        return layer_output\n    \nclass BertIntermediate(nn.Module):\n    def __init__(self, config:GraphConfig):\n        super().__init__()\n        self.dense = nn.Linear(config.hidden_size, config.intermediate_size)\n        if isinstance(config.hidden_act, str):\n            self.intermediate_act_fn = ACT2FN[config.hidden_act]\n        else:\n            self.intermediate_act_fn = config.hidden_act\n\n    def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:\n        hidden_states = self.dense(hidden_states)\n        hidden_states = self.intermediate_act_fn(hidden_states)\n        return hidden_states\n    \nclass BertOutput(nn.Module):\n    def __init__(self, config:GraphConfig):\n        super().__init__()\n        self.dense = nn.Linear(config.intermediate_size, config.hidden_size)\n        self.LayerNorm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)\n        self.dropout = nn.Dropout(config.hidden_dropout_prob)\n\n    def forward(self, hidden_states: torch.Tensor, input_tensor: torch.Tensor) -> torch.Tensor:\n        hidden_states = self.dense(hidden_states)\n        hidden_states = self.dropout(hidden_states)\n        hidden_states = self.LayerNorm(hidden_states + input_tensor)\n        return hidden_states\n    \nclass BertAttention(nn.Module):\n    def __init__(self, config:GraphConfig, position_embedding_type=None):\n        super().__init__()\n        self.self = BertSelfAttention(config, position_embedding_type=position_embedding_type)\n        self.output = BertSelfOutput(config)\n        self.pruned_heads = set()\n\n\n    def forward(\n        self,\n        hidden_states: torch.Tensor,\n        attention_mask: Optional[torch.FloatTensor] = None,\n        head_mask: Optional[torch.FloatTensor] = None,\n        output_attentions: Optional[bool] = False,\n    ) -> Tuple[torch.Tensor]:\n        self_outputs = self.self(\n            hidden_states,\n            attention_mask,\n            head_mask,\n            output_attentions,\n        )\n        attention_output = self.output(self_outputs[0], hidden_states)\n        outputs = (attention_output,) + self_outputs[1:]  # add attentions if we output them\n        return outputs\n    \n    \nclass BertSelfAttention(nn.Module):\n    def __init__(self, config:GraphConfig, position_embedding_type=None):\n        super().__init__()\n        if config.hidden_size % config.num_attention_heads != 0 and not hasattr(config, \"embedding_size\"):\n            raise ValueError(\n                f\"The hidden size ({config.hidden_size}) is not a multiple of the number of attention \"\n                f\"heads ({config.num_attention_heads})\"\n            )\n\n        self.num_attention_heads = config.num_attention_heads\n        self.attention_head_size = int(config.hidden_size / config.num_attention_heads)\n        self.all_head_size = self.num_attention_heads * self.attention_head_size\n\n        self.query = nn.Linear(config.hidden_size, self.all_head_size)\n        self.key = nn.Linear(config.hidden_size, self.all_head_size)\n        self.value = nn.Linear(config.hidden_size, self.all_head_size)\n\n        self.dropout = nn.Dropout(config.attention_probs_dropout_prob)\n        self.position_embedding_type = position_embedding_type or getattr(\n            config, \"position_embedding_type\", \"absolute\"\n        )\n        if self.position_embedding_type == \"relative_key\" or self.position_embedding_type == \"relative_key_query\":\n            self.max_position_embeddings = config.max_position_embeddings\n            self.distance_embedding = nn.Embedding(2 * config.max_position_embeddings - 1, self.attention_head_size)\n\n\n    def transpose_for_scores(self, x: torch.Tensor) -> torch.Tensor:\n        new_x_shape = x.size()[:-1] + (self.num_attention_heads, self.attention_head_size)\n        x = x.view(new_x_shape)\n        return x.permute(0, 2, 1, 3)\n\n    def forward(\n        self,\n        hidden_states: torch.Tensor,\n        attention_mask: Optional[torch.FloatTensor] = None,\n        head_mask: Optional[torch.FloatTensor] = None,\n        output_attentions: Optional[bool] = False,\n    ) -> Tuple[torch.Tensor]:\n        \n        mixed_query_layer = self.query(hidden_states)\n        key_layer = self.transpose_for_scores(self.key(hidden_states))\n        value_layer = self.transpose_for_scores(self.value(hidden_states))\n        query_layer = self.transpose_for_scores(mixed_query_layer)\n\n\n        # Take the dot product between \"query\" and \"key\" to get the raw attention scores.\n        attention_scores = torch.matmul(query_layer, key_layer.transpose(-1, -2))\n\n        if self.position_embedding_type == \"relative_key\" or self.position_embedding_type == \"relative_key_query\":\n            query_length, key_length = query_layer.shape[2], key_layer.shape[2]\n            position_ids_l = torch.arange(query_length, dtype=torch.long, device=hidden_states.device).view(-1, 1)\n            position_ids_r = torch.arange(key_length, dtype=torch.long, device=hidden_states.device).view(1, -1)\n            distance = position_ids_l - position_ids_r\n\n            positional_embedding = self.distance_embedding(distance + self.max_position_embeddings - 1)\n            positional_embedding = positional_embedding.to(dtype=query_layer.dtype)  # fp16 compatibility\n\n            if self.position_embedding_type == \"relative_key\":\n                relative_position_scores = torch.einsum(\"bhld,lrd->bhlr\", query_layer, positional_embedding)\n                attention_scores = attention_scores + relative_position_scores\n            elif self.position_embedding_type == \"relative_key_query\":\n                relative_position_scores_query = torch.einsum(\"bhld,lrd->bhlr\", query_layer, positional_embedding)\n                relative_position_scores_key = torch.einsum(\"bhrd,lrd->bhlr\", key_layer, positional_embedding)\n                attention_scores = attention_scores + relative_position_scores_query + relative_position_scores_key\n\n        attention_scores = attention_scores / math.sqrt(self.attention_head_size)\n        if attention_mask is not None:\n            # Apply the attention mask is (precomputed for all layers in BertModel forward() function)\n            attention_scores = attention_scores + attention_mask\n\n        # Normalize the attention scores to probabilities.\n        attention_probs = nn.functional.softmax(attention_scores, dim=-1)\n\n        # This is actually dropping out entire tokens to attend to, which might\n        # seem a bit unusual, but is taken from the original Transformer paper.\n        attention_probs = self.dropout(attention_probs)\n\n        # Mask heads if we want to\n        if head_mask is not None:\n            attention_probs = attention_probs * head_mask #DONE: Same Head Mask for all Heads\n\n        context_layer = torch.matmul(attention_probs, value_layer)\n\n        context_layer = context_layer.permute(0, 2, 1, 3).contiguous()\n        new_context_layer_shape = context_layer.size()[:-2] + (self.all_head_size,)\n        context_layer = context_layer.view(new_context_layer_shape)\n\n        outputs = (context_layer, attention_probs) if output_attentions else (context_layer,)\n\n        return outputs\n\n\nclass BertSelfOutput(nn.Module):\n    def __init__(self, config:GraphConfig):\n        super().__init__()\n        self.dense = nn.Linear(config.hidden_size, config.hidden_size)\n        self.LayerNorm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)\n        self.dropout = nn.Dropout(config.hidden_dropout_prob)\n\n    def forward(self, hidden_states: torch.Tensor, input_tensor: torch.Tensor) -> torch.Tensor:\n        hidden_states = self.dense(hidden_states)\n        hidden_states = self.dropout(hidden_states)\n        hidden_states = self.LayerNorm(hidden_states + input_tensor)\n        return hidden_states\n    \n    \nclass NodeEncoder(nn.Module):\n    \n    def __init__(self, config:GraphConfig):\n        super().__init__()\n        self.node_opcode_embeddings = nn.Embedding(NODE_OP_CODES+1 , config.embedding_size, padding_idx=NODE_OP_CODES)\n        self.linear = nn.Linear(NODE_FEATS, config.embedding_size, bias=False)\n        self.layer_norm = nn.LayerNorm(config.embedding_size, eps=config.layer_norm_eps)\n        \n        \n    def forward(self,\n                node_opcode: torch.Tensor,\n                node_feat: torch.Tensor\n                ) -> torch.Tensor:\n        opcode_embeddings = self.node_opcode_embeddings(node_opcode) \n        node_feats =  self.linear(node_feat)\n        features = opcode_embeddings + node_feats\n        features = self.layer_norm(features)\n        return features\n    \n    \nclass BertNodeEncoder(nn.Module):\n    \n    def __init__(self, config:GraphConfig) -> None:\n        super().__init__()\n        self.config = config\n        self.node_embeddings = NodeEncoder(config)\n        self.node_encoder = BertEncoder(config)\n        \n    def forward(self,\n                node_opcode: torch.Tensor,\n                node_feat: torch.Tensor,\n                edges_adjecency: torch.Tensor,\n                node_attn_mask: torch.Tensor\n                ):\n        node_embeddings = self.node_embeddings(node_opcode, node_feat)\n        node_attn_mask = node_attn_mask.unsqueeze(1).unsqueeze(-1)\n        node_encoder_outputs = self.node_encoder(node_embeddings,\n                                                 attention_mask=node_attn_mask,\n                                                 head_mask=edges_adjecency.unsqueeze(0).repeat(self.config.num_hidden_layers, 1, 1, 1).unsqueeze(2),\n                                                 output_attentions=True)\n        return node_encoder_outputs\n    \ndef transform_node_positional_embeddings(embeddings_output:torch.Tensor,\n                                         node_config_ids:torch.Tensor,\n                                         num_nodes:int\n                                         ) -> torch.Tensor:\n    bs, num_configs, _, dim = embeddings_output.shape\n    idxs = node_config_ids.unsqueeze(1).repeat(1,num_configs,1)\n    zeros = torch.zeros(bs, num_configs, num_nodes, dim, device=embeddings_output.device, dtype=embeddings_output.dtype)\n    idxs = idxs.unsqueeze(-1).repeat(1,1,1,dim)\n    zeros.scatter_reduce_(2, idxs, embeddings_output, reduce='sum')\n    return zeros\n\nclass NodeFeatEmbeddings(nn.Module):\n    def __init__(self, config:GraphConfig):\n        super().__init__()\n        self.config = config\n        self.node_feat_embeddings = nn.Linear(NODE_CONFIG_FEATS + CONFIG_FEATS, config.embedding_size, bias=False)\n        self.layer_norm = nn.LayerNorm(config.embedding_size, eps=config.layer_norm_eps)\n        \n    def forward(self, node_config_feat: torch.Tensor, node_config_ids: torch.Tensor, num_nodes:int) -> torch.Tensor:\n        node_config_feat_embeddings = self.node_feat_embeddings(node_config_feat)\n        node_config_feat_embeddings = self.layer_norm(node_config_feat_embeddings)\n        node_config_feat_embeddings = transform_node_positional_embeddings(node_config_feat_embeddings, node_config_ids, num_nodes)\n        return node_config_feat_embeddings\n        \n    \nclass BertGraphEncoder(nn.Module):\n    def __init__(self, config:GraphConfig) -> None:\n        super().__init__()\n        self.config = config\n        self.node_embeddings = NodeEncoder(config)\n        self.node_encoder = BertEncoder(config)\n        self.node_feat_embeddings = NodeFeatEmbeddings(config)\n        \n    def forward(self,\n                node_opcode: torch.Tensor, # (bs, num_nodes)\n                node_feat: torch.Tensor, # (bs, num_nodes, num_node_feats)\n                edges_adjecency: torch.Tensor, # (bs, num_nodes, num_nodes)\n                node_attn_mask: torch.Tensor, # (bs, num_nodes)\n                node_config_feat: torch.Tensor, # (bs, num_configs, num_config_nodes, num_node_feats)\n                node_config_ids: torch.Tensor, # (bs, num_configs, num_config_nodes)\n                ):\n        bs, num_nodes = node_opcode.shape\n        num_configs = node_config_feat.shape[1]\n        node_embeddings = self.node_embeddings(node_opcode, node_feat)\n        node_config_feat_embeddings = self.node_feat_embeddings(node_config_feat, node_config_ids, num_nodes)\n        \n        node_embeddings = node_embeddings.unsqueeze(1).repeat(1, num_configs, 1, 1)\n        node_embeddings += node_config_feat_embeddings\n        node_attn_mask = node_attn_mask.unsqueeze(1).repeat(1, num_configs, 1)\n        node_embeddings = node_embeddings.reshape(bs *num_configs, num_nodes, -1)\n        node_attn_mask = node_attn_mask.reshape(bs *num_configs, num_nodes)\n        node_attn_mask = node_attn_mask.unsqueeze(1).unsqueeze(-1)\n        edges_adjecency = edges_adjecency.unsqueeze(1).repeat(1, num_configs, 1, 1).reshape(bs *num_configs, num_nodes, num_nodes)\n        edges_adjecency = edges_adjecency.unsqueeze(1)\n        \n\n        node_encoder_outputs = self.node_encoder(node_embeddings,\n                                                 attention_mask=node_attn_mask,\n                                                 head_mask=edges_adjecency,\n                                                 output_attentions=True)\n        \n        return node_encoder_outputs.last_hidden_state.reshape(bs, num_configs, num_nodes, -1)\n    \n    \nclass GraphEncoder(nn.Module):\n    \n    config_class = GraphConfig\n    \n    def __init__(self, config:GraphConfig):\n        super().__init__()\n        self.config = config\n        self.node_encoder = BertGraphEncoder(config)\n        self.head = nn.Linear(config.hidden_size, 1)\n        self.loss_fn = MultiElementRankLoss(margin=config.margin, number_permutations=config.number_permutations)\n        \n        \n    def forward(self,\n                node_opcode: torch.Tensor, # (bs, num_nodes)\n                node_feat: torch.Tensor, # (bs, num_nodes, num_node_feats)\n                edges_adjecency: torch.Tensor, # (bs, num_nodes, num_nodes)\n                node_attn_mask: torch.Tensor, # (bs, num_nodes)\n                node_config_feat: torch.Tensor, # (bs, num_configs, num_config_nodes, num_node_feats)\n                node_config_ids: torch.Tensor, # (bs, num_configs, num_config_nodes)\n                config_idxs: Optional[torch.Tensor] = None, # (bs, num_configs)\n                config_runtime: Optional[torch.Tensor] = None,):\n        \n        last_hidden_state = self.node_encoder(node_opcode,\n                                    node_feat,\n                                    edges_adjecency,\n                                    node_attn_mask,\n                                    node_config_feat,\n                                    node_config_ids)\n        \n        output = self.head(last_hidden_state[:,:,0]).squeeze(-1)\n        outputs = {'outputs': output, 'order': torch.argsort(output, dim=1)}\n        if config_runtime is not None:\n            loss = 0\n            loss += self.loss_fn(output, config_runtime, config_idxs)\n            outputs['loss'] = loss\n        return outputs","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T05:13:20.923755Z","iopub.execute_input":"2025-11-12T05:13:20.923981Z","iopub.status.idle":"2025-11-12T05:13:20.965691Z","shell.execute_reply.started":"2025-11-12T05:13:20.923962Z","shell.execute_reply":"2025-11-12T05:13:20.964803Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class LightningWrapper(pl.LightningModule):\n    def __init__(self, model:nn.Module):\n        super().__init__()\n        self.model = model\n        self.topk = TileTopK()\n        \n    def forward(self, x):\n        return self.model(x)\n\n    def training_step(self, batch, batch_idx):\n        outputs = self.model(**batch)\n        return outputs['loss']\n\n    def validation_step(self, batch, batch_idx):\n        outputs = self.model(**batch)\n        loss = outputs['loss']\n        self.log(\"val_loss\", loss, prog_bar=True)\n        config_attn_mask = torch.ones_like(batch['config_runtime'], device=batch['config_runtime'].device)\n        self.topk.update(outputs['outputs'], batch['config_runtime'], config_attn_mask)\n        return loss\n    \n    def on_validation_end(self) -> None:\n        topk = self.topk.compute()\n        self.print(f\"topk {topk:.3f}\")\n        self.topk.reset()\n        return super().on_validation_end()\n\n    def test_step(self, batch, batch_idx):\n        x, y = batch\n        y_hat = self.model(x)\n        loss = self.model.loss(y_hat, y)\n        self.log(\"test_loss\", loss, prog_bar=True)\n        return loss\n\n    def configure_optimizers(self):\n        optimizer = torch.optim.AdamW(self.trainer.model.parameters(), lr=1e-3)\n        return optimizer","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T05:13:20.966582Z","iopub.execute_input":"2025-11-12T05:13:20.966788Z","iopub.status.idle":"2025-11-12T05:13:20.979363Z","shell.execute_reply.started":"2025-11-12T05:13:20.96677Z","shell.execute_reply":"2025-11-12T05:13:20.978741Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"config_kwargs = dict(hidden_size= 128,\n    num_attention_heads= 4,\n    num_hidden_layers= 2,\n    intermediate_size= 64,\n    gradient_checkpointing= True,\n    margin= 0.1,\n    number_permutations= 4,\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T05:13:20.98035Z","iopub.execute_input":"2025-11-12T05:13:20.980623Z","iopub.status.idle":"2025-11-12T05:13:20.989567Z","shell.execute_reply.started":"2025-11-12T05:13:20.980602Z","shell.execute_reply":"2025-11-12T05:13:20.988809Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"config = GraphConfig(**config_kwargs)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T05:13:20.990373Z","iopub.execute_input":"2025-11-12T05:13:20.990605Z","iopub.status.idle":"2025-11-12T05:13:21.052275Z","shell.execute_reply.started":"2025-11-12T05:13:20.990586Z","shell.execute_reply":"2025-11-12T05:13:21.051373Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = GraphEncoder(config)\nmodel = LightningWrapper(model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T05:13:21.053524Z","iopub.execute_input":"2025-11-12T05:13:21.053764Z","iopub.status.idle":"2025-11-12T05:13:21.070808Z","shell.execute_reply.started":"2025-11-12T05:13:21.053744Z","shell.execute_reply":"2025-11-12T05:13:21.070167Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"tile_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T05:13:21.071578Z","iopub.execute_input":"2025-11-12T05:13:21.071776Z","iopub.status.idle":"2025-11-12T05:13:21.087423Z","shell.execute_reply.started":"2025-11-12T05:13:21.071759Z","shell.execute_reply":"2025-11-12T05:13:21.08653Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df = tile_df.query(\"split == 'train'\").reset_index(drop=True)\nvalid_df = tile_df.query(\"split == 'valid'\").reset_index(drop=True)\ntrain_dataset = TileDataset(train_df, num_configs=24)\nvalid_dataset = TileDataset(valid_df, num_configs=24)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T05:13:21.088712Z","iopub.execute_input":"2025-11-12T05:13:21.089309Z","iopub.status.idle":"2025-11-12T05:13:21.109228Z","shell.execute_reply.started":"2025-11-12T05:13:21.08928Z","shell.execute_reply":"2025-11-12T05:13:21.108527Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataloader = DataLoader(train_dataset, collate_fn=collate_fn, batch_size=8, num_workers=4, shuffle=True, persistent_workers=True)\nvalid_dataloader = DataLoader(valid_dataset, collate_fn=collate_fn, batch_size=8, num_workers=4)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T05:13:21.110366Z","iopub.execute_input":"2025-11-12T05:13:21.110693Z","iopub.status.idle":"2025-11-12T05:13:21.115545Z","shell.execute_reply.started":"2025-11-12T05:13:21.110663Z","shell.execute_reply":"2025-11-12T05:13:21.114652Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trainer_config = dict(\n    max_epochs= 50,\n    precision= 32,\n    gradient_clip_val= 1.0,\n    accumulate_grad_batches= 4,\n    check_val_every_n_epoch= 10)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T05:13:21.117552Z","iopub.execute_input":"2025-11-12T05:13:21.117812Z","iopub.status.idle":"2025-11-12T05:13:21.126345Z","shell.execute_reply.started":"2025-11-12T05:13:21.117786Z","shell.execute_reply":"2025-11-12T05:13:21.125447Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pytorch_lightning as pl\nimport sys\ntorch.set_float32_matmul_precision(\"medium\") \n\nclass PrintCallback(pl.Callback):\n    def on_validation_epoch_end(self, trainer, pl_module):\n        metrics = trainer.callback_metrics\n        epoch = trainer.current_epoch\n        def safe(x): return f\"{x:.4f}\" if x is not None else \"N/A\"\n        print(f\"[Epoch {epoch+1}] train_loss={safe(metrics.get('train_loss'))}, \"\n              f\"val_loss={safe(metrics.get('val_loss'))}, \"\n              f\"topk={safe(metrics.get('topk'))}\", flush=True)\n\ntrainer = pl.Trainer(\n    **trainer_config,\n    accelerator=\"gpu\",\n    devices=1,\n    enable_progress_bar=False,\n    callbacks=[PrintCallback()],\n    num_sanity_val_steps=0,\n)\n\ntrainer.fit(model, train_dataloader, valid_dataloader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T05:13:21.489418Z","iopub.execute_input":"2025-11-12T05:13:21.489813Z","iopub.status.idle":"2025-11-12T05:35:02.39318Z","shell.execute_reply.started":"2025-11-12T05:13:21.489786Z","shell.execute_reply":"2025-11-12T05:35:02.391985Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = 'cuda:0' if torch.cuda.is_available() else 'cpu'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T05:35:02.395532Z","iopub.execute_input":"2025-11-12T05:35:02.396267Z","iopub.status.idle":"2025-11-12T05:35:02.400423Z","shell.execute_reply.started":"2025-11-12T05:35:02.396238Z","shell.execute_reply":"2025-11-12T05:35:02.399524Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"split = 'test'\ntest_tile_df = tile_df.query(\"split == @split\").reset_index(drop=True)\ntest_tile_ds = TileDataset(test_tile_df, num_configs=-1)\ncollate_fn = LayoutCollator(64, targets=split!=\"test\")\ntest_dataloader = DataLoader(test_tile_ds, batch_size=1, shuffle=False, num_workers=0, collate_fn=collate_fn)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T05:35:02.401425Z","iopub.execute_input":"2025-11-12T05:35:02.401733Z","iopub.status.idle":"2025-11-12T05:35:02.419304Z","shell.execute_reply.started":"2025-11-12T05:35:02.401711Z","shell.execute_reply":"2025-11-12T05:35:02.418667Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.to(device)\nmodel = model.eval()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T05:35:02.421094Z","iopub.execute_input":"2025-11-12T05:35:02.421325Z","iopub.status.idle":"2025-11-12T05:35:02.436386Z","shell.execute_reply.started":"2025-11-12T05:35:02.421305Z","shell.execute_reply":"2025-11-12T05:35:02.435775Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def chunk_batch(batch, start_idx, end_idx):\n    output = {k:batch[k] for k in ['node_opcode', 'node_feat', 'edges_adjecency', 'node_attn_mask', 'node_config_ids']}\n    output['node_config_feat'] = batch['node_config_feat'][:, start_idx: end_idx]\n    return output\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T05:35:02.43732Z","iopub.execute_input":"2025-11-12T05:35:02.437571Z","iopub.status.idle":"2025-11-12T05:35:02.445902Z","shell.execute_reply.started":"2025-11-12T05:35:02.437549Z","shell.execute_reply":"2025-11-12T05:35:02.445113Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pred_order = []\nfor batch in tqdm(test_dataloader):\n    batch.pop('config_idxs')\n    batch = {k: v.to(device) for k, v in batch.items()}\n    num_configs = batch['node_config_feat'].shape[1]\n    # Chunk the configs to avoid OOM errors\n    configs_cut_points = list(range(0,num_configs, 100)) + [num_configs]\n    chunk_order = []\n    for start, end in zip(configs_cut_points, configs_cut_points[1:]):\n        chunked_batch = chunk_batch(batch, start, end)\n        with torch.no_grad():\n            output = model.model(**chunked_batch)\n        chunk_order.extend(output['outputs'].cpu().numpy())\n    pred_order.append(np.argsort(np.concatenate(chunk_order))[:5])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T05:35:02.447096Z","iopub.execute_input":"2025-11-12T05:35:02.447632Z","iopub.status.idle":"2025-11-12T05:36:31.497773Z","shell.execute_reply.started":"2025-11-12T05:35:02.447601Z","shell.execute_reply":"2025-11-12T05:36:31.496873Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"idxs_string = [\";\".join(map(str,elem)) for elem in pred_order]\ntest_tile_df['TopConfigs'] = idxs_string\ntest_tile_df = test_tile_df[['ID', 'TopConfigs']]\n\ntest_tile_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T06:06:17.616514Z","iopub.execute_input":"2025-11-12T06:06:17.617457Z","iopub.status.idle":"2025-11-12T06:06:17.638831Z","shell.execute_reply.started":"2025-11-12T06:06:17.617418Z","shell.execute_reply":"2025-11-12T06:06:17.637863Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_tile_df.shape[0]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T06:03:42.947581Z","iopub.execute_input":"2025-11-12T06:03:42.947881Z","iopub.status.idle":"2025-11-12T06:03:42.95324Z","shell.execute_reply.started":"2025-11-12T06:03:42.94786Z","shell.execute_reply":"2025-11-12T06:03:42.952404Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# LAYOUT","metadata":{}},{"cell_type":"code","source":"!pip install -U dill\n!pip install -U tensorflow_gnn --pre\n!pip install -U tensorflow_ranking","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T06:06:38.083759Z","iopub.execute_input":"2025-11-12T06:06:38.084429Z","iopub.status.idle":"2025-11-12T06:07:19.508394Z","shell.execute_reply.started":"2025-11-12T06:06:38.084402Z","shell.execute_reply":"2025-11-12T06:07:19.507495Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import datetime, os, time, sys\nfrom pathlib import Path\nfrom typing import Dict, Optional, List, Union, Tuple\nfrom dataclasses import dataclass\nimport math\nimport numpy as np\nimport pandas as pd\nfrom datasets import Dataset\nimport time, gc\nimport joblib\nimport optuna\nfrom tqdm import tqdm\n\n\nimport tensorflow as tf\nimport tensorflow_gnn as tfgnn\nimport tensorflow_ranking as tfr\n\nimport tpugraphsv1_layout_data_py as layout_data\nimport tpugraphsv1_implicit_py as implicit\n\nprint(tf.__version__)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T06:22:31.945125Z","iopub.execute_input":"2025-11-12T06:22:31.946013Z","iopub.status.idle":"2025-11-12T06:22:50.740278Z","shell.execute_reply.started":"2025-11-12T06:22:31.945982Z","shell.execute_reply":"2025-11-12T06:22:50.739052Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Load Data","metadata":{}},{"cell_type":"code","source":"layout_npz_dataset = None  # declare as global","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T06:24:22.357221Z","iopub.execute_input":"2025-11-12T06:24:22.357603Z","iopub.status.idle":"2025-11-12T06:24:22.362119Z","shell.execute_reply.started":"2025-11-12T06:24:22.357573Z","shell.execute_reply":"2025-11-12T06:24:22.361132Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# `MAX_KEEP_NODES` is (or, is not) useful for Segment Dropout, if model uses\n# edges \"sampled_config\" and \"sampled_feed\" (or, \"config\" and \"feed\")\ndef split_layout_dataset(source, search):\n    global layout_npz_dataset\n    layout_data_root_dir = os.path.join(os.path.expanduser(LAYOUT_DATA_ROOT), source, search)\n    # Load layout dataset\n    layout_npz_dataset = layout_data.get_npz_dataset(\n        layout_data_root_dir,\n        min_train_configs=CONFIGS_PER_GRAPH,\n        max_train_configs=MAX_NUM_CONFIGS,  # Default: 500 If any graph has more than this configurations, it will be filtered [speeds up loading + training]\n        cache_dir=f'cache_{source}_{search}'\n    )\n    \n    ## Layout training dataset\n    layout_train_ds = (layout_npz_dataset.train.get_graph_tensors_dataset(CONFIGS_PER_GRAPH, max_nodes=MAX_KEEP_NODES)\n                            .shuffle(100, reshuffle_each_iteration=True)\n                            .batch(BATCH_SIZE, drop_remainder=False)\n                            .map(tfgnn.GraphTensor.merge_batch_to_components)\n                            .map(pair_layout_graph_with_label))\n    # Layout valid dataset\n    layout_valid_ds = (layout_npz_dataset.validation.get_graph_tensors_dataset(CONFIGS_PER_GRAPH)\n                            .batch(BATCH_SIZE, drop_remainder=False) \n                            .map(tfgnn.GraphTensor.merge_batch_to_components)\n                            .map(pair_layout_graph_with_label))\n                       \n#     print(next(iter(layout_train_ds)))\n    return layout_train_ds, layout_valid_ds","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T06:24:23.510186Z","iopub.execute_input":"2025-11-12T06:24:23.510758Z","iopub.status.idle":"2025-11-12T06:24:23.517138Z","shell.execute_reply.started":"2025-11-12T06:24:23.510724Z","shell.execute_reply":"2025-11-12T06:24:23.516209Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"def _mlp(dims, hidden_activation, l2reg=1e-4, use_bias=True):\n    \"\"\"Helper function for multi-layer perceptron (MLP).\"\"\"\n    layers = []\n    for i, dim in enumerate(dims):\n        if i > 0:\n            layers.append(tf.keras.layers.Activation(hidden_activation))\n        layers.append(tf.keras.layers.Dense(dim, kernel_regularizer=tf.keras.regularizers.l2(l2reg),\n                                            use_bias=use_bias))\n    return tf.keras.Sequential(layers)\n\n\nclass _OpEmbedding(tf.keras.Model):\n    \"\"\"Embeds GraphTensor.node_sets['op']['op'] nodes into feature 'op_e'.\"\"\"\n\n    def __init__(self, num_ops: int, embed_d: int, l2reg: float = 1e-4):\n        super().__init__()\n        self.embedding_layer = tf.keras.layers.Embedding(num_ops,\n                                                         embed_d,\n                                                         activity_regularizer=tf.keras.regularizers.l2(l2reg))\n\n    def call(self, graph: tfgnn.GraphTensor, training: bool = False) -> tfgnn.GraphTensor:\n        op_features = dict(graph.node_sets['op'].features)\n        op_features['op_e'] = self.embedding_layer(tf.cast(graph.node_sets['op']['op'], tf.int32))\n        return graph.replace_features(node_sets={'op': op_features})\n\n\ndef pair_layout_graph_with_label(graph: tfgnn.GraphTensor):\n    \"\"\"Extracts label from graph (`tfgnn.GraphTensor`) and returns a pair of `(graph, label)`\"\"\"\n    # Return runtimes divded over large number: only ranking is required. The\n    # runtimes are in the 100K range\n    label = tf.cast(graph.node_sets['g']['runtimes'], tf.float32) / 1e7\n    return graph, label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T06:24:27.321214Z","iopub.execute_input":"2025-11-12T06:24:27.321602Z","iopub.status.idle":"2025-11-12T06:24:27.329713Z","shell.execute_reply.started":"2025-11-12T06:24:27.321574Z","shell.execute_reply":"2025-11-12T06:24:27.32864Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ResModel(tf.keras.Model):\n    \"\"\"GNN with residual connections.\"\"\"\n    def __init__(self, num_configs: int, num_ops: int, op_embed_dim: int = 32,\n                 num_gnns: int = 2, mlp_layers: int = 2,\n                 hidden_activation: str = 'leaky_relu',\n                 hidden_dim: int = 32, reduction: str = 'sum'):\n        super().__init__()\n        self._num_configs = num_configs\n        self._num_ops = num_ops\n        self._op_embedding = _OpEmbedding(num_ops, op_embed_dim)\n        self._prenet = _mlp([hidden_dim] * mlp_layers, hidden_activation)\n        self._gc_layers = []\n        for _ in range(num_gnns):\n            self._gc_layers.append(_mlp([hidden_dim] * mlp_layers, hidden_activation))\n        self._postnet = _mlp([hidden_dim, 1], hidden_activation, use_bias=False)\n\n    def call(self, graph: tfgnn.GraphTensor, training: bool = False):\n        del training\n        return self.forward(graph, self._num_configs)\n\n    def _node_level_forward(self, node_features: tf.Tensor,\n                            config_features: tf.Tensor,\n                            graph: tfgnn.GraphTensor, num_configs: int,\n                            edgeset_prefix='') -> tf.Tensor:\n        adj_op_op = implicit.AdjacencyMultiplier(graph, edgeset_prefix+'feed')  # op->op\n        adj_config = implicit.AdjacencyMultiplier(graph, edgeset_prefix+'config')  # nconfig->op\n\n        adj_op_op_hat = (adj_op_op + adj_op_op.transpose()).add_eye()\n        adj_op_op_hat = adj_op_op_hat.normalize_symmetric()\n\n        x = node_features\n\n        x = tf.stack([x] * num_configs, axis=1)\n        config_features = 100 * (adj_config @ config_features)\n        x = tf.concat([config_features, x], axis=-1)\n        x = self._prenet(x)\n        x = tf.nn.leaky_relu(x)\n\n        for layer in self._gc_layers:\n            y = x\n            y = tf.concat([config_features, y], axis=-1)\n            y = tf.nn.leaky_relu(layer(adj_op_op_hat @ y))\n            x += y\n        return x\n\n    def forward(self, graph: tfgnn.GraphTensor, num_configs: int, backprop=True) -> tf.Tensor:\n        graph = self._op_embedding(graph)\n\n        config_features = graph.node_sets['nconfig']['feats']\n        node_features = tf.concat([ graph.node_sets['op']['feats'],\n                                    graph.node_sets['op']['op_e']], axis=-1)\n\n        x_full = self._node_level_forward(node_features=tf.stop_gradient(node_features),\n                                          config_features=tf.stop_gradient(config_features),\n                                          graph=graph, num_configs=num_configs)\n\n        if backprop:\n            x_backprop = self._node_level_forward(\n                                node_features=node_features,\n                                config_features=config_features,\n                                graph=graph, num_configs=num_configs,\n                                edgeset_prefix='sampled_')\n\n            is_selected = graph.node_sets['op']['selected']\n            # Need to expand twice as `is_selected` is a vector (num_nodes) but\n            # x_{backprop, full} are 3D tensors (num_nodes, num_configs, num_feats).\n            is_selected = tf.expand_dims(is_selected, axis=-1)\n            is_selected = tf.expand_dims(is_selected, axis=-1)\n            x = tf.where(is_selected, x_backprop, x_full)\n        else:\n            x = x_full\n        # Multiplication of adjacency matrix of two graphs\n        adj_config = implicit.AdjacencyMultiplier(graph, 'config')\n\n        # Features for configurable nodes.\n        config_feats = (adj_config.transpose() @ x)\n\n        # Global pooling\n        adj_pool_op_sum = implicit.AdjacencyMultiplier(graph, 'g_op').transpose()\n        adj_pool_op_mean = adj_pool_op_sum.normalize_right()\n        adj_pool_config_sum = implicit.AdjacencyMultiplier(graph, 'g_config').transpose()\n        x = self._postnet(tf.concat([\n            # (A D^-1) @ Features\n            adj_pool_op_mean @ x,\n            # l2_normalize( A @ Features )\n            tf.nn.l2_normalize(adj_pool_op_sum @ x, axis=-1),\n            # l2_normalize( A @ Features )\n            tf.nn.l2_normalize(adj_pool_config_sum @ config_feats, axis=-1),\n        ], axis=-1))\n\n        x = tf.squeeze(x, -1)\n\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T06:24:28.665428Z","iopub.execute_input":"2025-11-12T06:24:28.665781Z","iopub.status.idle":"2025-11-12T06:24:28.678618Z","shell.execute_reply.started":"2025-11-12T06:24:28.665757Z","shell.execute_reply":"2025-11-12T06:24:28.677695Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def build_model(num_ops):    \n    # Create a ResModel\n    model = ResModel(CONFIGS_PER_GRAPH, num_ops)\n\n    loss = tfr.keras.losses.ListMLELoss()  # (temperature=10)\n    # opt = tf.keras.optimizers.Adam(learning_rate=1e-3, clipnorm=0.5)\n    opt = tf.keras.optimizers.AdamW(learning_rate=1e-3, clipnorm=0.5)\n    # opt = tf.keras.optimizers.SGD(lr=1e-3)\n\n    model.compile(loss=loss, optimizer=opt,\n                  metrics=[tfr.keras.metrics.OPAMetric(name='opa_metric')],\n                  steps_per_execution=32)\n    return model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T06:24:33.569011Z","iopub.execute_input":"2025-11-12T06:24:33.569347Z","iopub.status.idle":"2025-11-12T06:24:33.574364Z","shell.execute_reply.started":"2025-11-12T06:24:33.569321Z","shell.execute_reply":"2025-11-12T06:24:33.57352Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training","metadata":{}},{"cell_type":"code","source":"AUTO = tf.data.experimental.AUTOTUNE\n\n# Detect TPU, return appropriate distribution strategy\ntry:\n    tpu = tf.distribute.cluster_resolver.TPUClusterResolver() \n    print('Running on TPU ', tpu.master())\nexcept ValueError:\n    tpu = None\n\nif tpu:\n    tf.config.experimental_connect_to_cluster(tpu)\n    tf.tpu.experimental.initialize_tpu_system(tpu)\n    strategy = tf.distribute.experimental.TPUStrategy(tpu)\nelse:\n    strategy = tf.distribute.get_strategy() ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T06:24:35.009865Z","iopub.execute_input":"2025-11-12T06:24:35.010686Z","iopub.status.idle":"2025-11-12T06:24:35.015504Z","shell.execute_reply.started":"2025-11-12T06:24:35.010657Z","shell.execute_reply":"2025-11-12T06:24:35.014437Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Batch size information.\nBATCH_SIZE = 16 * strategy.num_replicas_in_sync # Number of graphs per batch.\n\nCONFIGS_PER_GRAPH = 5  # Number of configurations (features and target values) per graph.\nMAX_NUM_CONFIGS = 350 # Maximal Number of configurations used for filter. Default value = 500 \nMAX_KEEP_NODES = 500  # Useful for dropout.\nBUFFER_SIZE = 10000\n\nprint(BATCH_SIZE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T06:24:35.993353Z","iopub.execute_input":"2025-11-12T06:24:35.994172Z","iopub.status.idle":"2025-11-12T06:24:35.998838Z","shell.execute_reply.started":"2025-11-12T06:24:35.994143Z","shell.execute_reply":"2025-11-12T06:24:35.997929Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def clear_memory():\n    gc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T06:24:36.882672Z","iopub.execute_input":"2025-11-12T06:24:36.88322Z","iopub.status.idle":"2025-11-12T06:24:36.88845Z","shell.execute_reply.started":"2025-11-12T06:24:36.883185Z","shell.execute_reply":"2025-11-12T06:24:36.887175Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_model_by_epoch(epoch, model, layout_train_ds, layout_valid_ds):\n    global best_params, best_val_opa, best_val_at_epoch, early_stop\n    print(f\"Starting training the model with EPOCH {epoch}\")\n    start = time.time()\n\n    # Train the model \n    history = model.fit(layout_train_ds, epochs=epoch+1, verbose=\"auto\", \n                        batch_size = BATCH_SIZE,\n                        workers=4, validation_data=layout_valid_ds) # Frequency of validation\n\n    # get the training result\n    print(f\"epoch = {epoch} history = {history.history}\")\n#             train_loss = history.history['loss'][-1]\n#             train_opa = history.history['opa_metric'][-1]\n#             val_loss = history.history['val_loss'][-1]\n    val_opa = history.history['val_opa_metric'][-1]\n    if val_opa > best_val_opa:\n        best_val_opa = val_opa\n        best_val_at_epoch = epoch\n        best_params = {v.ref: v + 0 for v in model.trainable_variables}\n        print(f' * [{epoch}] Validation (NEW BEST): {val_opa}')\n        \n    print(f\"Finish the training for EPOCH {epoch} in {time.time() - start}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T06:24:38.503164Z","iopub.execute_input":"2025-11-12T06:24:38.503746Z","iopub.status.idle":"2025-11-12T06:24:38.509561Z","shell.execute_reply.started":"2025-11-12T06:24:38.503719Z","shell.execute_reply":"2025-11-12T06:24:38.508514Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"LAYOUT_DATA_ROOT = '/kaggle/input/predict-ai-model-runtime/npz_all/npz/layout'\nTOTAL_EPOCHS = 1  # Total number of epochs\n\ndef training_model(source, search):\n    start = time.time()\n\n    # Split layout dataset into 'train' and 'valid' datasets\n    layout_train_ds, layout_valid_ds = split_layout_dataset(source, search)\n    num_ops = layout_npz_dataset.num_ops\n\n    # Create a ResModel\n    model = build_model(num_ops)\n\n    # Initialize model variables\n    model.fit(layout_train_ds, epochs=0, verbose=\"auto\")\n\n    # ### Train for a few epochs.\n    global best_params, best_val_opa, best_val_at_epoch, early_stop \n    early_stop = 5  # Stop if validation OPA did not improve in this many epochs\n    best_params = None\n    best_val_opa = -1\n    best_val_at_epoch = -1\n    epochs = TOTAL_EPOCHS\n\n    for epoch in range(epochs):\n        try:\n            train_model_by_epoch(epoch, model, layout_train_ds, layout_valid_ds)\n            if early_stop > 0 and epoch - best_val_at_epoch >= early_stop:\n                print(f'[{epoch}] Best accuracy was attained at epoch {best_val_at_epoch}. Stopping.')\n                break\n        except Exception as error:\n            print(\"An exception occurred during training the model:\", error)\n\n    # Restore best parameters\n    print('Restoring parameters corresponding to the best validation OPA.')\n    assert best_params is not None\n    for v in model.trainable_variables:\n        v.assign(best_params[v.ref])\n\n    # Save weights\n    model.save_weights(f'/kaggle/working/layout_{source}_{search}/best_model_{source}_{search}')\n\n    # Clean up\n    del layout_train_ds, layout_valid_ds\n    print(f\"Training time of {source}-{search}: {time.time() - start}\")\n    \n    return model\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T06:24:39.989313Z","iopub.execute_input":"2025-11-12T06:24:39.989672Z","iopub.status.idle":"2025-11-12T06:24:39.996749Z","shell.execute_reply.started":"2025-11-12T06:24:39.989642Z","shell.execute_reply":"2025-11-12T06:24:39.995889Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Take a cut of the configs.\ndef infer_worker(i, graph, model, num_configs):\n    end_i = min(i + INFERENCE_CONFIGS_BATCH_SIZE, num_configs)\n    # Take a cut of the configs.\n    node_set_g = graph.node_sets['g']\n    subconfigs_graph = tfgnn.GraphTensor.from_pieces(\n        edge_sets=graph.edge_sets, ## Edges\n        ## Node set\n        node_sets={'op': graph.node_sets['op'],\n                   'nconfig': tfgnn.NodeSet.from_fields(\n                               sizes=graph.node_sets['nconfig'].sizes,\n                               features={'feats': graph.node_sets['nconfig']['feats'][:, i:end_i]}\n                                                       ),\n                   'g': tfgnn.NodeSet.from_fields(\n                               sizes=tf.constant([1]),\n                               features={'graph_id': node_set_g['graph_id'],\n                                         'runtimes': node_set_g['runtimes'][:, i:end_i],\n                                         'kept_node_ratio': node_set_g['kept_node_ratio']}\n                                                  ) # End of 'g'\n\n                    } # End of 'node_sets'\n        )\n    h = model.forward(subconfigs_graph, num_configs=(end_i - i), backprop=False) # Don't update model weights\n#     return h[0]\n    global all_scores  # needed to modify the global value\n    all_scores.append(h[0]) \n    # print(f\"Total number of scores = {len(all_scores)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T06:24:42.513286Z","iopub.execute_input":"2025-11-12T06:24:42.513688Z","iopub.status.idle":"2025-11-12T06:24:42.520114Z","shell.execute_reply.started":"2025-11-12T06:24:42.513648Z","shell.execute_reply":"2025-11-12T06:24:42.519293Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import concurrent\nfrom concurrent.futures import ThreadPoolExecutor\nfrom functools import partial\n\nINFERENCE_CONFIGS_BATCH_SIZE = 50\n# Source can be \"xla\" or \"nlp\". Search can be \"random\" or \"default\"\ndef infer_layout(source, search, is_pretrain=True):\n    # Training the model\n    model = training_model(source, search)\n    # Infer the results using the model\n    start = time.time()\n    # Create the submission file\n    output_csv_filename = f'inference_layout_{source}_{search}.csv'\n    print('\\n\\n   Running inference on test set ...\\n\\n')\n    # Store the results of test dataset\n    test_rankings = []\n    assert layout_npz_dataset.test.graph_id is not None\n    for graph in tqdm(layout_npz_dataset.test.iter_graph_tensors(),\n                      total=layout_npz_dataset.test.graph_id.shape[-1],\n                      desc='Inference'):\n        num_configs = graph.node_sets['g']['runtimes'].shape[-1]\n        print(f\"num_configs = {num_configs}\")\n        global all_scores # declare as a global value\n        all_scores = []\n#         func = partial(infer_worker, graph=graph, model=model, num_configs=num_configs)\n#         with concurrent.futures.ThreadPoolExecutor(max_workers=3) as executor:\n#             # execute tasks concurrently and process results in order\n#             for result in executor.map(func, list(range(0, num_configs, INFERENCE_CONFIGS_BATCH_SIZE))):\n#                 all_scores.append(result) \n        for i in tqdm(range(0, num_configs, INFERENCE_CONFIGS_BATCH_SIZE)):\n            infer_worker(i, graph, model, num_configs)\n        all_scores = tf.concat(all_scores, axis=0)\n        graph_id = graph.node_sets['g']['graph_id'][0].numpy().decode()\n        sorted_indices = tf.strings.join(tf.strings.as_string(tf.argsort(all_scores)), ';').numpy().decode()\n        test_rankings.append((f\"layout:{source}:{search}:{graph_id}\", sorted_indices))\n    # Write the test_ranking \n    df = pd.DataFrame(test_rankings, columns=['ID', 'TopConfigs'])\n    df.to_csv(output_csv_filename)\n\n    del model\n    clear_memory()\n    print(f\"Total inference time of {source}-{search}: {time.time() - start}\")\n    return test_rankings","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T06:42:03.24072Z","iopub.execute_input":"2025-11-12T06:42:03.241404Z","iopub.status.idle":"2025-11-12T06:42:03.249246Z","shell.execute_reply.started":"2025-11-12T06:42:03.241375Z","shell.execute_reply":"2025-11-12T06:42:03.248387Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# --- Step 1: Run Layout inference and collect predictions ---\ntest_rankings = []\nfor source, search in [('nlp', 'default'), ('nlp', 'random'), ('xla', 'default'), ('xla', 'random')]:\n    test_rankings.extend(infer_layout(source, search))\n\nlayout_test_df = pd.DataFrame(test_rankings, columns=['ID', 'TopConfigs'])\nlayout_test_df.to_csv('/kaggle/working/inference_layout_all.csv', index=False)\n\n# --- Step 2: Merge Tile and Layout predictions into final submission ---\nsubmission_df = pd.read_csv('../input/predict-ai-model-runtime/sample_submission.csv')\n\n# Remove IDs already predicted by Tile model\nlayout_only_df = layout_test_df[~layout_test_df['ID'].isin(test_tile_df['ID'])]\n\n# Concatenate Tile + Layout\nfinal_submission_df = pd.concat([test_tile_df, layout_only_df], axis=0)\n\n# Ensure order matches sample_submission\nfinal_submission_df = submission_df[['ID']].merge(final_submission_df, on='ID', how='left')\n\n# Save submission\nfinal_submission_df.to_csv('/kaggle/working/submission.csv', index=False)\nfinal_submission_df.head(3)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T06:42:05.140934Z","iopub.execute_input":"2025-11-12T06:42:05.141763Z","iopub.status.idle":"2025-11-12T13:19:27.870277Z","shell.execute_reply.started":"2025-11-12T06:42:05.141733Z","shell.execute_reply":"2025-11-12T13:19:27.869354Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_tile_df.shape[0]\n#final_submission_df.shape[0]\n#layout_only_df.shape[0]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T13:29:24.271273Z","iopub.execute_input":"2025-11-12T13:29:24.271587Z","iopub.status.idle":"2025-11-12T13:29:24.276924Z","shell.execute_reply.started":"2025-11-12T13:29:24.271565Z","shell.execute_reply":"2025-11-12T13:29:24.276074Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_tile_df.to_csv('/kaggle/working/test_tile_results.csv', index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-12T13:35:48.647076Z","iopub.execute_input":"2025-11-12T13:35:48.647409Z","iopub.status.idle":"2025-11-12T13:35:48.655242Z","shell.execute_reply.started":"2025-11-12T13:35:48.647384Z","shell.execute_reply":"2025-11-12T13:35:48.654548Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}