# %% [markdown]
# **LOADING LIBRARIES**

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:04:26.382537Z","iopub.execute_input":"2023-11-02T03:04:26.382930Z","iopub.status.idle":"2023-11-02T03:04:26.388075Z","shell.execute_reply.started":"2023-11-02T03:04:26.382899Z","shell.execute_reply":"2023-11-02T03:04:26.387025Z"}}
import pandas as pd
import numpy as np
import os
import torch
from torch import optim
from torch import nn
from torch.utils.data import Dataset, DataLoader
import torchvision.transforms as transforms

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:04:26.389856Z","iopub.execute_input":"2023-11-02T03:04:26.390167Z","iopub.status.idle":"2023-11-02T03:04:26.408398Z","shell.execute_reply.started":"2023-11-02T03:04:26.390134Z","shell.execute_reply":"2023-11-02T03:04:26.407098Z"}}
if torch.cuda.is_available():
    device = torch.device('cuda')
else:
    device = torch.device('cpu')
device

# %% [markdown]
# **LOADING DATA**

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:04:26.411576Z","iopub.execute_input":"2023-11-02T03:04:26.412235Z","iopub.status.idle":"2023-11-02T03:04:29.442470Z","shell.execute_reply.started":"2023-11-02T03:04:26.412155Z","shell.execute_reply":"2023-11-02T03:04:29.441520Z"}}
for dirname, _, filenames in os.walk('/kaggle/input/'):
    if len(filenames) != 0:
        if filenames[0] != "sample_submission.csv":
            avg = np.array([os.path.getsize(os.path.join(dirname, filename)) for filename in filenames]).mean()
            # os.path.getsize returns the size of the dictionary passed in bytes
            print(dirname, len(os.listdir(dirname)))
            print("Size: {:.3f} KB".format(avg/1024)) # 1024 is 1KB

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:04:29.443781Z","iopub.execute_input":"2023-11-02T03:04:29.445010Z","iopub.status.idle":"2023-11-02T03:04:29.454356Z","shell.execute_reply.started":"2023-11-02T03:04:29.444969Z","shell.execute_reply":"2023-11-02T03:04:29.452968Z"}}
tile = np.load('/kaggle/input/predict-ai-model-runtime/npz_all/npz/tile/xla/train/retinanet.4x4.fp32_-431a58cc30e72ec6.npz')
tile.files

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:04:29.457555Z","iopub.execute_input":"2023-11-02T03:04:29.458698Z","iopub.status.idle":"2023-11-02T03:04:35.584866Z","shell.execute_reply.started":"2023-11-02T03:04:29.458660Z","shell.execute_reply":"2023-11-02T03:04:35.583558Z"}}
basic_structure = tile.files
for dirname, _, filenames in os.walk('/kaggle/input'):
    flag = False
    for filename in filenames:
        if filename != "sample_submission.csv" and filename[-4:] == '.npz':
            
            if np.load(os.path.join(dirname, filename)).files != basic_structure and not flag:
                print(dirname)
                print(np.load(os.path.join(dirname, filename)).files)
                flag = True

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:04:35.586462Z","iopub.execute_input":"2023-11-02T03:04:35.587095Z","iopub.status.idle":"2023-11-02T03:04:35.593693Z","shell.execute_reply.started":"2023-11-02T03:04:35.587060Z","shell.execute_reply":"2023-11-02T03:04:35.592433Z"}}
def load_data(directory):
    splits = ['train', 'valid', 'test']
    dfs = dict()
    for split in splits:
        path = os.path.join(directory, split)
        files = os.listdir(path)
        list_df = []
        
        for file in files:
            list_df.append(dict(np.load(os.path.join(path, file))))
        dfs[split] = pd.DataFrame.from_dict(list_df)
    return dfs

# %% [markdown]
# **DATA PREPROCESSING**

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:04:35.595599Z","iopub.execute_input":"2023-11-02T03:04:35.595993Z","iopub.status.idle":"2023-11-02T03:04:35.609169Z","shell.execute_reply.started":"2023-11-02T03:04:35.595960Z","shell.execute_reply":"2023-11-02T03:04:35.607664Z"}}
def make_datum(datum, i):
    configs = datum.config_feat[i].shape[0]
    graph = pd.DataFrame(columns=['node_feat', 'node_opcode', 'edge_index', 'config_feat', 'config_runtime', 'config_runtime_normalizers'])
    for config in range(0, configs):
        Sample = pd.Series()
        Sample['node_feat'] = datum.node_feat[i]
        Sample['node_opcode'] = datum.node_opcode[i]
        Sample['edge_index'] = datum.edge_index[i]
        Sample['config_feat'] = datum.config_feat[i][config]
        Sample['config_runtime'] = datum.config_runtime[i][config]
        Sample['config_runtime_normalizers'] = datum.config_runtime_normalizers[i][config]
        Sample['topo_order'] = datum.topo_order.values[0]
        graph = pd.concat([graph, Sample.to_frame().T], ignore_index=True)
    return graph


# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:04:35.611098Z","iopub.execute_input":"2023-11-02T03:04:35.611718Z","iopub.status.idle":"2023-11-02T03:04:35.623688Z","shell.execute_reply.started":"2023-11-02T03:04:35.611684Z","shell.execute_reply":"2023-11-02T03:04:35.621673Z"}}
def make_data(data, b=0):
    df = pd.DataFrame(columns = data.columns)
    for i in range(len(data)):
        datum = data.iloc[i].to_frame().T
        graph = make_datum(datum, i+b)
        df = pd.concat([df, graph], axis=0, ignore_index=True)
    df['avg_runtime'] = df['config_runtime'] / (df['config_runtime_normalizers'] + 1e-5)
    df.drop(['config_runtime', 'config_runtime_normalizers'], axis=1, inplace=True)
    return df

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:04:35.625205Z","iopub.execute_input":"2023-11-02T03:04:35.625922Z","iopub.status.idle":"2023-11-02T03:04:35.641145Z","shell.execute_reply.started":"2023-11-02T03:04:35.625890Z","shell.execute_reply":"2023-11-02T03:04:35.640241Z"}}
class dataset(Dataset):
    def __init__(self, data):
        self.feats = data[data.columns[:-1]]
        self.labels = data[data.columns[-1]]
        self.transform = transforms.Compose([transforms.ToTensor()])
    
    def __len__(self):
        return len(self.feats)
    
    def __getitem__(self, idx):
        # early join of config features
        op_code = torch.tensor(self.feats.node_opcode.iloc[idx])
        node_f = torch.tensor(self.feats.node_feat.iloc[idx])
        config_f = torch.tensor(self.feats.config_feat.iloc[idx])
        adj_mat = torch.tensor(self.feats.edge_index.iloc[idx])
        topo_o = torch.tensor(self.feats.topo_order.iloc[idx])
        config_f_broadcasted = config_f.unsqueeze(0).expand(node_f.size(0), -1)
        node_config_f = torch.cat((node_f, config_f_broadcasted), axis=1)
        # padding for compatibility with the model GNN
        to_add_rows = 500 - node_config_f.size(0)
        node_config_row_zeros = torch.zeros((to_add_rows, node_config_f.size(1)), dtype=node_config_f.dtype)
        node_config_f = torch.cat((node_config_f, node_config_row_zeros), dim=0)
        op_top_zeros = torch.zeros(to_add_rows, dtype=op_code.dtype)
        op_code = torch.cat((op_code, op_top_zeros))
        topo_o = torch.cat((topo_o, op_top_zeros))
        adj_row_zeros = torch.zeros((to_add_rows, adj_mat.size(1)), dtype=adj_mat.dtype)
        adj_mat = torch.cat((adj_mat, adj_row_zeros), dim=0)
        adj_col_zeros = torch.zeros((adj_mat.size(0), to_add_rows), dtype=adj_mat.dtype)
        adj_mat = torch.cat((adj_mat, adj_col_zeros), dim=1)
        # padding done (nodes == 100)
        inputs = [op_code, node_config_f, adj_mat, topo_o]
        return inputs, torch.tensor(self.labels.iloc[idx])

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:04:35.642414Z","iopub.execute_input":"2023-11-02T03:04:35.643304Z","iopub.status.idle":"2023-11-02T03:04:35.660289Z","shell.execute_reply.started":"2023-11-02T03:04:35.643268Z","shell.execute_reply":"2023-11-02T03:04:35.658985Z"}}
test = False
if(test):
    custom_dataset = dataset(made_tilex_train)
    loader = DataLoader(custom_dataset, batch_size=2, shuffle=True)
    for batch in loader:
        break

# %% [markdown]
# **BUILDING MODEL GNN**

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:04:35.665593Z","iopub.execute_input":"2023-11-02T03:04:35.665945Z","iopub.status.idle":"2023-11-02T03:04:35.676645Z","shell.execute_reply.started":"2023-11-02T03:04:35.665917Z","shell.execute_reply":"2023-11-02T03:04:35.675268Z"}}
test = False
if(test):
    gnn = GNN(174)
    batch_feats, batch_true = batch
    batch_preds = gnn(batch_feats)
    print(batch_preds)

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:04:35.678217Z","iopub.execute_input":"2023-11-02T03:04:35.678527Z","iopub.status.idle":"2023-11-02T03:04:35.694411Z","shell.execute_reply.started":"2023-11-02T03:04:35.678501Z","shell.execute_reply":"2023-11-02T03:04:35.692927Z"}}
def adj_mat(edges, n):
    # edges are 0-indexed
    adj_matrix = np.zeros((n, n))
    for u, v in edges:
        adj_matrix[u][v] = 1
        adj_matrix[v][u] = 1
    return adj_matrix

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:04:35.696405Z","iopub.execute_input":"2023-11-02T03:04:35.696930Z","iopub.status.idle":"2023-11-02T03:04:35.713133Z","shell.execute_reply.started":"2023-11-02T03:04:35.696885Z","shell.execute_reply":"2023-11-02T03:04:35.711743Z"}}
adj_matrix = adj_mat([[2, 3], [3, 1]], 4)
print(type(adj_matrix))
adj_matrix

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:04:35.714371Z","iopub.execute_input":"2023-11-02T03:04:35.714985Z","iopub.status.idle":"2023-11-02T03:04:35.727129Z","shell.execute_reply.started":"2023-11-02T03:04:35.714939Z","shell.execute_reply":"2023-11-02T03:04:35.725605Z"}}
def topo_sort(edges):
    indeg = {}
    for edge in edges:
        if(edge[0] not in indeg):
            indeg[edge[0]] = 0
        if(edge[1] not in indeg):
            indeg[edge[1]] = 0
        indeg[edge[1]] += 1
    
    queue = []
    for node in indeg:
        if(indeg[node] == 0):
            queue.append(node)
    topo_order = []
    while(queue):
        node = queue.pop()
        topo_order.append(node)
        for neighbour in [edge[1] for edge in edges if edge[0] == node]:
            indeg[neighbour] -= 1
            if(indeg[neighbour] == 0): 
                queue.append(neighbour)
    # let's reorder the topological sort to be compatible with topo-order aware downsampling
    topo_embd = []
    for i, node in enumerate(topo_order):
        topo_embd.append((node, i))
    topo_embd = sorted(topo_embd)
    topo_order = []
    for node, i in topo_embd:
        topo_order.append(i)
    return topo_order

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:04:35.729083Z","iopub.execute_input":"2023-11-02T03:04:35.729769Z","iopub.status.idle":"2023-11-02T03:04:35.743247Z","shell.execute_reply.started":"2023-11-02T03:04:35.729732Z","shell.execute_reply":"2023-11-02T03:04:35.742189Z"}}
edges = [[2, 3], [3, 1], [1, 4], [3, 4]]
topo_sort(edges)

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:04:35.744747Z","iopub.execute_input":"2023-11-02T03:04:35.745239Z","iopub.status.idle":"2023-11-02T03:04:35.753580Z","shell.execute_reply.started":"2023-11-02T03:04:35.745194Z","shell.execute_reply":"2023-11-02T03:04:35.752512Z"}}
def kipf_norm(adj_mats):
    norm_mats = []
    for adj_mat in adj_mats:
        max_degree = torch.max(torch.sum(adj_mat, dim=1))
        # add inv-deg term if convergence is slow
        norm_mat = adj_mat / max_degree
        norm_mats.append(norm_mat)
    return torch.stack(norm_mats, axis=0)

# %% [markdown]
# **TILE CONFIGURATION**

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:04:35.754745Z","iopub.execute_input":"2023-11-02T03:04:35.755141Z","iopub.status.idle":"2023-11-02T03:05:05.170528Z","shell.execute_reply.started":"2023-11-02T03:04:35.755103Z","shell.execute_reply":"2023-11-02T03:05:05.169450Z"}}
tile_xla = load_data('/kaggle/input/predict-ai-model-runtime/npz_all/npz/tile/xla/')

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:05:05.172124Z","iopub.execute_input":"2023-11-02T03:05:05.172463Z","iopub.status.idle":"2023-11-02T03:05:05.180829Z","shell.execute_reply.started":"2023-11-02T03:05:05.172434Z","shell.execute_reply":"2023-11-02T03:05:05.179748Z"}}
tilex_train = tile_xla['train']
tilex_valid = tile_xla['valid']
tilex_test = tile_xla['test']


# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:05:05.182536Z","iopub.execute_input":"2023-11-02T03:05:05.182939Z","iopub.status.idle":"2023-11-02T03:05:05.630458Z","shell.execute_reply.started":"2023-11-02T03:05:05.182903Z","shell.execute_reply":"2023-11-02T03:05:05.629282Z"}}
print(tilex_train.shape)
# convert all the edge_index to adjacency matrix
# also add topo_order of edge_index
tilex_train.head()

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:05:05.632111Z","iopub.execute_input":"2023-11-02T03:05:05.632478Z","iopub.status.idle":"2023-11-02T03:05:05.638978Z","shell.execute_reply.started":"2023-11-02T03:05:05.632450Z","shell.execute_reply":"2023-11-02T03:05:05.637853Z"}}
# prepare train data 1
test = False
if test:
    tile_train = tilex_train.head(len(tilex_train)//2).copy(deep=True)
    tile_train['topo_order'] = tile_train['edge_index'].map(lambda e: topo_sort(e))
    tile_train['edge_index'] = tile_train['edge_index'].map(lambda e: adj_mat(e, max([max(edge) for edge in e])+1))
    tile_train = make_data(tile_train)
    tile_train.to_pickle("/kaggle/working/tiler_train1.pkl")

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:05:05.640492Z","iopub.execute_input":"2023-11-02T03:05:05.640805Z","iopub.status.idle":"2023-11-02T03:05:05.660284Z","shell.execute_reply.started":"2023-11-02T03:05:05.640778Z","shell.execute_reply":"2023-11-02T03:05:05.659009Z"}}
# prepare train data 2
test = False
if test:
    b = len(tilex_train)//2
    tile_train = tilex_train.tail(len(tilex_train)-b).copy(deep=True)
    tile_train['topo_order'] = tile_train['edge_index'].map(lambda e: topo_sort(e))
    tile_train['edge_index'] = tile_train['edge_index'].map(lambda e: adj_mat(e, max([max(edge) for edge in e])+1))
    tile_train = make_data(tile_train, b)
    tile_train.to_pickle('/kaggle/working/tiler_train2.pkl')

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:05:05.661948Z","iopub.execute_input":"2023-11-02T03:05:05.662295Z","iopub.status.idle":"2023-11-02T03:05:05.674621Z","shell.execute_reply.started":"2023-11-02T03:05:05.662267Z","shell.execute_reply":"2023-11-02T03:05:05.673144Z"}}
# prepare valid data
test = False
if test:
    tile_val = tilex_valid.copy(deep=True)
    tile_val['topo_order'] = tile_val['edge_index'].map(lambda e: topo_sort(e))
    tile_val['edge_index'] = tile_val['edge_index'].map(lambda e: adj_mat(e, max([max(edge) for edge in e])+1))
    tile_val = make_data(tile_val)
    tile_val.to_pickle('/kaggle/working/tiler_val.pkl')

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:05:05.675678Z","iopub.execute_input":"2023-11-02T03:05:05.677236Z","iopub.status.idle":"2023-11-02T03:05:05.688236Z","shell.execute_reply.started":"2023-11-02T03:05:05.677201Z","shell.execute_reply":"2023-11-02T03:05:05.686955Z"}}
# test GNN, gConv, data preprocessing steps
toy = tilex_train.head(1).copy(deep=True)
n = toy.node_feat[0].shape[0]
toy['topo_order'] = toy['edge_index'].map(lambda e: topo_sort(e))

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:05:05.689860Z","iopub.execute_input":"2023-11-02T03:05:05.690431Z","iopub.status.idle":"2023-11-02T03:05:05.791775Z","shell.execute_reply.started":"2023-11-02T03:05:05.690403Z","shell.execute_reply":"2023-11-02T03:05:05.790678Z"}}
# edge_index -> adj_matrix
toy['edge_index'] = toy['edge_index'].map(lambda e: adj_mat(e, n))
toy

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:05:05.793092Z","iopub.execute_input":"2023-11-02T03:05:05.793429Z","iopub.status.idle":"2023-11-02T03:05:05.897040Z","shell.execute_reply.started":"2023-11-02T03:05:05.793401Z","shell.execute_reply":"2023-11-02T03:05:05.895543Z"}}
print(toy.edge_index[0].shape)
toy

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:05:05.898741Z","iopub.execute_input":"2023-11-02T03:05:05.899254Z","iopub.status.idle":"2023-11-02T03:05:13.770225Z","shell.execute_reply.started":"2023-11-02T03:05:05.899220Z","shell.execute_reply":"2023-11-02T03:05:13.768541Z"}}
# apply GNN, gConv on toy
import warnings
warnings.filterwarnings('ignore')
made_toy = make_data(toy)
made_toy

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:05:13.772436Z","iopub.execute_input":"2023-11-02T03:05:13.772775Z","iopub.status.idle":"2023-11-02T03:05:13.800506Z","shell.execute_reply.started":"2023-11-02T03:05:13.772748Z","shell.execute_reply":"2023-11-02T03:05:13.799224Z"}}
# make data into batchs and test it on gConv and GNN
custom_dataset = dataset(made_toy)
loader = DataLoader(custom_dataset, batch_size=2, shuffle=True)
for batch in loader:
    break
batch

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:05:13.802273Z","iopub.execute_input":"2023-11-02T03:05:13.802597Z","iopub.status.idle":"2023-11-02T03:05:13.809350Z","shell.execute_reply.started":"2023-11-02T03:05:13.802569Z","shell.execute_reply":"2023-11-02T03:05:13.807896Z"}}
# use batch to test the model
test = False
if test:
    batch_feat, batch_labels = batch
    model = GNN(174, 2)
    model = model.to(device)
    batch_preds = model(batch_feat)
    batch_preds

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:05:13.814786Z","iopub.execute_input":"2023-11-02T03:05:13.815210Z","iopub.status.idle":"2023-11-02T03:05:13.825019Z","shell.execute_reply.started":"2023-11-02T03:05:13.815165Z","shell.execute_reply":"2023-11-02T03:05:13.824107Z"}}
avg = np.array([os.path.getsize(os.path.join(dirname, filename)) for filename in filenames]).mean()
# what does this mean, something related to the size of the data?

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:05:13.826692Z","iopub.execute_input":"2023-11-02T03:05:13.827080Z","iopub.status.idle":"2023-11-02T03:05:13.839496Z","shell.execute_reply.started":"2023-11-02T03:05:13.827047Z","shell.execute_reply":"2023-11-02T03:05:13.838235Z"}}
sample = tilex_train.head(1)
print('no of nodes: (#nodes) in computation graph ', len(sample.node_feat[0]))
print('node_feat: (#features_of_a_node) ', len(sample.node_feat[0][0]))
print('node_opcode: (#operation_codes) ', len(sample.node_opcode[0]))
print('edge_index: (#edges) ', len(sample.edge_index[0]))
print('configs: (#configurations) ', len(sample.config_feat[0]))
print('config_feat: (#config_features) ', len(sample.config_feat[0][0]))
print('config_runtime: (#runtimes == # configs) ', len(sample.config_runtime[0]))
print('config_runtime_normalizers: (#runtimes_normalized == #configs) ', len(sample.config_runtime_normalizers[0]))

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:05:13.841395Z","iopub.execute_input":"2023-11-02T03:05:13.842520Z","iopub.status.idle":"2023-11-02T03:05:13.853617Z","shell.execute_reply.started":"2023-11-02T03:05:13.842473Z","shell.execute_reply":"2023-11-02T03:05:13.852517Z"}}
sample = tilex_train.head(2)
print('no of nodes: (#nodes) in computation graph ', len(sample.node_feat[1]))
print('node_feat: (#features_of_a_node) ', len(sample.node_feat[1][0]))
print('node_opcode: (#operation_codes) ', len(sample.node_opcode[1]))
print('edge_index: (#edges) ', len(sample.edge_index[1]))
print('configs: (#configurations) ', len(sample.config_feat[1]))
print('config_feat: (#config_features) ', len(sample.config_feat[1][0]))
print('config_runtime: (#runtimes == # configs) ', len(sample.config_runtime[1]))
print('config_runtime_normalizers: (#runtimes_normalized == #configs) ', len(sample.config_runtime_normalizers[1]))

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:05:13.854968Z","iopub.execute_input":"2023-11-02T03:05:13.855315Z","iopub.status.idle":"2023-11-02T03:05:13.867466Z","shell.execute_reply.started":"2023-11-02T03:05:13.855286Z","shell.execute_reply":"2023-11-02T03:05:13.866063Z"}}
class gConv(nn.Module):
    def __init__(self, in_feat, out_feat, batch_size):
        super(gConv, self).__init__()
        self.in_feat = in_feat
        self.out_feat = out_feat
        # try nn.init xavier_normal_ or xavier_uniform_ if convergence is slow
        self.w = nn.Parameter(torch.stack([torch.randn(self.in_feat, self.out_feat, device=device)]*batch_size, axis=0), requires_grad=True)
        
    def forward(self, adj_matrix, feats):
        adj_matrix = adj_matrix.to(device)
        feats = feats.to(device)
        id_matrix = torch.eye(feats.shape[1], device=device)
        id_matrix = torch.stack([id_matrix]*feats.shape[0], axis=0)
        adj_matrix = adj_matrix + id_matrix
        adj_matrix = kipf_norm(adj_matrix)
        agg_feat = torch.matmul(torch.matmul(adj_matrix, feats), self.w)
        return torch.relu(agg_feat)

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:05:13.869395Z","iopub.execute_input":"2023-11-02T03:05:13.869901Z","iopub.status.idle":"2023-11-02T03:05:13.889207Z","shell.execute_reply.started":"2023-11-02T03:05:13.869856Z","shell.execute_reply":"2023-11-02T03:05:13.888323Z"}}
gconv = gConv(in_feat=3, out_feat=2, batch_size=2) # feats in, feats out
gconv = gconv.to(device) # need re-assign
adj_matrix = torch.FloatTensor([[0, 1, 1], 
              [1, 0, 0], 
              [1, 0, 0]])
feats = torch.FloatTensor([[1, 2, 3], 
                           [2, 1, 2], 
                           [3, 1, 5]])
x = gconv(torch.unsqueeze(adj_matrix, 0), torch.unsqueeze(feats, 0))
x

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:05:13.891165Z","iopub.execute_input":"2023-11-02T03:05:13.891681Z","iopub.status.idle":"2023-11-02T03:05:13.911107Z","shell.execute_reply.started":"2023-11-02T03:05:13.891638Z","shell.execute_reply":"2023-11-02T03:05:13.909902Z"}}
class GNN(nn.Module):
    def __init__(self, input_feat, batch_size):
        super(GNN, self).__init__()
        self.input_feat = input_feat
        self.gConv1 = gConv(input_feat, 256, batch_size)
        self.gConv1.to(device)
        self.gConv2 = gConv(256, 512, batch_size)
        self.gConv2.to(device)
        self.gConv3 = gConv(512, 128, batch_size)
        self.gConv3.to(device)
        self.gConv4 = gConv(128, 64, batch_size)
        self.gConv4.to(device)
        self.gConv5 = gConv(64, 32, batch_size)
        self.gConv5.to(device)
    
    def forward(self, inputs):
        op_code, node_config_feat, adj_matrix, topo_order = inputs
        op_code, node_config_feat, adj_matrix, topo_order = op_code, node_config_feat, adj_matrix, topo_order
        embd_layer = nn.Embedding(120, 10)
        # type casting
        op_code = op_code.clone().detach().to(torch.int64)
        adj_matrix = adj_matrix.clone().detach().to(torch.float)
        topo_order = topo_order.clone().detach().to(torch.int64)
        # type casting
        op_embd = embd_layer(op_code)
        feats = torch.cat((op_embd, node_config_feat), dim=-1)
        # type casting
        feats = feats.clone().detach().to(torch.float)
        # type casting
        assert (feats.shape[-1] == (10+140+24))
        # graph convolutions
        x = self.gConv1(adj_matrix, feats)
        x = self.gConv2(adj_matrix, x)
        x = self.gConv3(adj_matrix, x)
        x = self.gConv4(adj_matrix, x)
        x = self.gConv5(adj_matrix, x)
        assert(x.shape[1] == feats.shape[1])
        assert(x.shape[2] == 32)
        # topological order aware downsampling
        topo_embd_layer = nn.Embedding(feats.shape[1], 10)
        topo_embd = topo_embd_layer(topo_order)
        topo_embd = torch.transpose(topo_embd, 1, 2)
        assert(topo_embd.shape[1] == 10)
        assert(topo_embd.shape[2] == feats.shape[1])
        latent_feat = torch.matmul(topo_embd.to(device), x.to(device))
        flat_feat = torch.flatten(latent_feat, start_dim=1)
        flat_dim = flat_feat.shape[1]
        
        # linear layers init
        self.dense1 = nn.Linear(flat_dim, 256).to(device)
        self.bn1 = nn.BatchNorm1d(256).to(device)
        self.dense2 = nn.Linear(256, 128).to(device)
        self.bn2 = nn.BatchNorm1d(128).to(device)
        self.dense3 = nn.Linear(128, 32).to(device)
        self.bn3 = nn.BatchNorm1d(32).to(device)
        self.dense4 = nn.Linear(32, 8).to(device)
        self.bn4 = nn.BatchNorm1d(8).to(device)
        self.dense5 = nn.Linear(8, 1).to(device)
        self.act = nn.ReLU().to(device)
        # linear layers build
        x = self.act(self.dense1(flat_feat))
        x = self.bn1(x)
        x = self.act(self.dense2(x))
        x = self.bn2(x)
        x = self.act(self.dense3(x))
        x = self.bn3(x)
        x = self.act(self.dense4(x))
        x = self.bn4(x)
        x = self.act(self.dense5(x))
        return x

# %% [markdown]
# **TRAINING TILE**

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:05:42.249657Z","iopub.execute_input":"2023-11-02T03:05:42.250055Z","iopub.status.idle":"2023-11-02T03:05:42.299709Z","shell.execute_reply.started":"2023-11-02T03:05:42.250024Z","shell.execute_reply":"2023-11-02T03:05:42.298297Z"}}
batch_size = 64 # determine the batch_size closely (also frame a algorithmic question related to this and post in on LeetCode)
model = GNN(174, batch_size)
model = model.to(device)

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:06:49.466240Z","iopub.execute_input":"2023-11-02T03:06:49.466694Z","iopub.status.idle":"2023-11-02T03:06:49.545642Z","shell.execute_reply.started":"2023-11-02T03:06:49.466658Z","shell.execute_reply":"2023-11-02T03:06:49.543448Z"}}
train1_pkl = "/kaggle/input/processed-tile-data-1/tiler_train1.pkl"
train2_pkl = "/kaggle/input/processed-tile-data-2/tiler_train2.pkl"
val_pkl = "/kaggle/input/valid-fast-slow/tiler_val.pkl"
train1 = True
if train1:
    tile_train = pd.read_pickle(train1_pkl)
else:
    tile_train = pd.read_pickle(train2_pkl)
tile_val = pd.read_pickle(val_pkl)

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:05:14.064060Z","iopub.status.idle":"2023-11-02T03:05:14.064593Z","shell.execute_reply.started":"2023-11-02T03:05:14.064392Z","shell.execute_reply":"2023-11-02T03:05:14.064413Z"}}
print("train data rows: ", len(tile_train))
# slice the data to make it compatible with the batch_size
train_just = tile_train.iloc[0:5546240].copy(deep=True)
train_just.head()

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:05:14.065956Z","iopub.status.idle":"2023-11-02T03:05:14.066814Z","shell.execute_reply.started":"2023-11-02T03:05:14.066586Z","shell.execute_reply":"2023-11-02T03:05:14.066608Z"}}
print("valid data rows: ", len(tile_val))
# slice the data to make it compatible with the batch_size
val_just = tile_val.iloc[0:1042048].copy(deep=True)
val_just.head()

# %% [markdown]
# **TRAINING MODEL**

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:05:14.068477Z","iopub.status.idle":"2023-11-02T03:05:14.068858Z","shell.execute_reply.started":"2023-11-02T03:05:14.068677Z","shell.execute_reply":"2023-11-02T03:05:14.068695Z"}}
test = True
if test:
    criterion = nn.MSELoss()
    optimizer = optim.Adam(model.parameters(), lr=0.01)
    num_epochs = 0 # change the no of epochs to train it
    train_batchs = dataset(train_just)
    val_batchs = dataset(val_just)
    train_batchs = DataLoader(train_batchs, batch_size=batch_size, shuffle=True)
    val_batchs = DataLoader(val_batchs, batch_size=batch_size, shuffle=True) 
    for epoch in range(num_epochs):
        tot_loss = 20
        for batch in train_batchs:
            batch_feats, batch_true = batch
            optimizer.zero_grad()
            batch_preds = model(batch_feats)
            batch_preds = batch_preds.view(batch_size,)
            loss = criterion(batch_preds.to(device), batch_true.to(device))
            tot_loss += loss.item()
            # add high weightage to top 5 fastest runtime config (ranking loss)
            loss.backward()
            optimizer.step()
        print('train tot_loss: ', tot_loss, end='  ')
        tot_loss = 0
        with torch.no_grad():
            for batch in val_batchs:
                batch_feats, batch_true = batch
                batch_preds = model(batch_feats)
                batch_preds = batch_preds.view(batch_size,)
                tot_loss += criterion(batch_preds.to(device), batch_true.to(device)).item()
        print('valid tot_loss: ', tot_loss)

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:05:14.070159Z","iopub.status.idle":"2023-11-02T03:05:14.070905Z","shell.execute_reply.started":"2023-11-02T03:05:14.070701Z","shell.execute_reply":"2023-11-02T03:05:14.070723Z"}}
checkpoint = {
    'model_state_dict': model.state_dict(),
    'optimizer_state_dict': optimizer.state_dict(),
}
model_path = '/kaggle/working/trained_model.pth'
torch.save(checkpoint, model_path)

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:05:14.072265Z","iopub.status.idle":"2023-11-02T03:05:14.073103Z","shell.execute_reply.started":"2023-11-02T03:05:14.072894Z","shell.execute_reply":"2023-11-02T03:05:14.072919Z"}}
load = False
if load:
    model = GNN()
    checkpoint = torch.load(model_path)
    model.load_state_dict(checkpoint["model_state_dict"])
    optimizer = optim.Adam(model.parameters(), lr=0.01)
    optimizer.load_state_dict(checkpoint["optimizer_state_dict"])

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:05:14.075324Z","iopub.status.idle":"2023-11-02T03:05:14.075731Z","shell.execute_reply.started":"2023-11-02T03:05:14.075543Z","shell.execute_reply":"2023-11-02T03:05:14.075562Z"}}
test_data = tile_xla["test"]
print(len(test_data))
test_data.head()

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:05:14.077306Z","iopub.status.idle":"2023-11-02T03:05:14.077969Z","shell.execute_reply.started":"2023-11-02T03:05:14.077763Z","shell.execute_reply":"2023-11-02T03:05:14.077783Z"}}
if False: # Change it to True to run the cell
    model.eval()
    result = []
    for i in range(2): # len(test_data)
        row = test_data.iloc[i].to_frame().T
        row = row.copy(deep=True)
        row['topo_order'] = row['edge_index'].map(lambda e: topo_sort(e))
        row['edge_index'] = row['edge_index'].map(lambda e: adj_mat(e, max([max(edge) for edge in e])+1))
        made_row = make_data(row, i)
        rank = []
        for config in range(2): # len(made_row)
            econ = made_row.iloc[config].to_frame().T
            econ_batch = dataset(econ)
            econ_data = DataLoader(econ_batch, batch_size=1, shuffle=False)
            for batch in econ_data:
                batch_feats, batch_labels = batch
                pred = model(batch_feats)
                pred_mean = torch.mean(pred).item()
                print(pred_mean, batch_labels.item())
            rank.append((config, pred_mean))
        rank.sort(key=lambda x: x[1])
        # take top 5 and append to sample_submissions.csv
        result.append([x[0] for x in rank[:5]])

# %% [code] {"execution":{"iopub.status.busy":"2023-11-02T03:05:14.079360Z","iopub.status.idle":"2023-11-02T03:05:14.080232Z","shell.execute_reply.started":"2023-11-02T03:05:14.079988Z","shell.execute_reply":"2023-11-02T03:05:14.080011Z"}}

if False: # Change it to True to get the inference
    assert(len(result) == 844)
    sub = pd.read_csv("/kaggle/input/predict-ai-model-runtime/sample_submission.csv")
    for i, res in enumerate(result):
        s = ""
        for config in res:
            s += str(config)+';'
        s = s[:-1]
        sub.iloc[i].TopConfigs = s
    sub.to_csv("/kaggle/working/sub.csv", index=False)