{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"%%writefile submission.py\nimport os\n\nos.system(\"pip install --no-index --find-links=../input/torchgeometric/ torch-geometric torch-scatter torch-sparse torch-cluster torch-spline-conv\")\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch_geometric.nn import GCNConv\nfrom torch_geometric.nn.norm.graph_norm import GraphNorm\nfrom torch_geometric.nn.glob.glob import global_mean_pool\n\nclass GNNBlock(nn.Module):\n    def __init__(self, dim_in, dim):\n        super().__init__()\n        self.fc = nn.Linear(3, 1, bias=False)\n        self.conv1 = GCNConv(dim_in, dim, bias=False)\n        self.conv2 = GCNConv(dim, 4 * dim, bias=False)\n        self.conv3 = GCNConv(4 * dim, dim, bias=False)\n        self.gnorm = GraphNorm(dim)\n    \n    def forward(self, x, edge_index, edge_attr, batch):\n        edge_weight = torch.sigmoid(self.fc(edge_attr))\n        y = self.gnorm(self.conv1(x, edge_index, edge_weight), batch)\n        y = F.gelu(self.conv2(y, edge_index, edge_weight))\n        return self.conv3(y, edge_index, edge_weight)\n\nclass ResGNNBlock(nn.Module):\n    def __init__(self, dim, edge_conv=True, num_heads=6):\n        super().__init__()\n        self.convs = nn.ModuleList([GNNBlock(dim, dim // num_heads) for _ in range(num_heads)])\n        if edge_conv:\n            self.head_q = nn.Linear(dim, 3, bias=False)\n            self.head_k = nn.Linear(dim, 3, bias=False)\n        self.edge_conv = edge_conv\n    \n    def forward(self, x, edge_index, edge_attr, batch):\n        w = torch.cat([conv(x, edge_index, edge_attr, batch) for conv in self.convs], dim=1)\n        if self.edge_conv:\n            q = self.head_q(w)\n            k = self.head_k(w)\n            row, col = edge_index\n            edge_attr = edge_attr + torch.tanh(q[row] * k[col])\n        return (x + w), edge_attr\n\nclass Encoder(nn.Module):\n    def __init__(self, dim_in, dim, depth):\n        super().__init__()\n        self.in_gate = nn.Linear(dim_in, dim)\n        self.blocks = nn.ModuleList([ResGNNBlock(dim, edge_conv=(k < depth - 1)) for k in range(depth)])\n\n    def forward(self, x, edge_index, edge_attr, batch=None):\n        x = self.in_gate(x)\n        for b in self.blocks:\n            x, edge_attr = b(x, edge_index, edge_attr, batch)\n        return x\n\nclass KoreNet(nn.Module):\n    def __init__(self, dim_in=22, dim=48, depth=4):\n        super().__init__()\n        self.encoder = Encoder(dim_in, dim, depth)\n        self.lstm = nn.LSTMCell(dim + 16, dim)\n        self.head_q = nn.Linear(dim, dim, bias=False) # query\n        self.head_k = nn.Linear(dim, dim, bias=False) # key\n        # 0 -> do-nothing\n        # 1 ... 14 -> trip-length\n        # 15 -> spawn\n        self.head_p = nn.Linear(dim, 16) # policy head\n        self.head_fs = nn.Linear(dim + 16, 1) # fraction of ships to spawn\n        self.head_fl = nn.Linear(dim + 16, 1) # fraction of ships to launch\n\n    def similarity(self, x, h, c, y):\n        h, c = self.lstm(x, (h, c))\n        q = self.head_q(h).unsqueeze(0)\n        k = self.head_k(y)\n        scores = torch.bmm(q, k.transpose(1, 2).contiguous())\n        return scores.reshape(-1, 5), (h, c) # (N, 5)\n\n    def get_mask(self, x):\n        mask = torch.zeros((x.shape[0], 16)).bool().to(x.device)\n        for k in range(1, 15):\n            mask[:, k] = x[:, 10] < 0.1 * k # filter on max trip-length\n        mask[:, 15] = (x[:, 11] == 0) # filter on kore capacity\n        return mask\n\n    def forward(self, batch):\n        edge_attr = torch.abs(batch.edge_attr).float()\n        z = self.encoder(batch.x, batch.edge_index, edge_attr, batch.batch)\n        avg = global_mean_pool(z, batch.batch)\n        w = avg[batch.batch]\n        p = self.head_p(z)\n        mask = self.get_mask(batch.x)\n        p.masked_fill_(mask, -float('inf'))\n        return p, (z, w)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-07-02T12:42:22.338839Z","iopub.execute_input":"2022-07-02T12:42:22.339303Z","iopub.status.idle":"2022-07-02T12:42:22.379682Z","shell.execute_reply.started":"2022-07-02T12:42:22.339181Z","shell.execute_reply":"2022-07-02T12:42:22.378488Z"},"_kg_hide-input":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile -a submission.py\n\nimport numpy as np\n\ndef get_max_trip_length(num_ships):\n    return int(np.floor(2 * np.log(num_ships)) + 1) if num_ships > 0 else 0\n\ndef get_min_ship_num_for_trip(trip_length):\n    return int(np.ceil(np.exp(0.5 * (trip_length - 1))))\n\ndef get_mining_percentage(num_ships):\n    return np.floor(np.log(num_ships) / 20) if num_ships > 0 else 0","metadata":{"execution":{"iopub.status.busy":"2022-07-02T12:42:22.38187Z","iopub.execute_input":"2022-07-02T12:42:22.38241Z","iopub.status.idle":"2022-07-02T12:42:22.389513Z","shell.execute_reply.started":"2022-07-02T12:42:22.382349Z","shell.execute_reply":"2022-07-02T12:42:22.38864Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile -a submission.py\n\nimport networkx as nx\nfrom kaggle_environments.envs.kore_fleets.helpers import *\nfrom kaggle_environments.envs.kore_fleets.helpers import *\nfrom torch_geometric.utils.convert import from_networkx\nfrom torch_geometric.data import Data\nimport torch\nfrom torch_geometric.utils import add_remaining_self_loops\nfrom torch_sparse import SparseTensor\n\nGraphPos = Tuple[int, int]\nGraphPath = List[GraphPos]\n\n\ndef expand_plan(flight_plan: str, direction: str=None) -> str:\n    if len(flight_plan) == 0:\n        return direction\n    plan_expanded = ''\n    prev_c = direction\n    buffer = '0'\n    for c in flight_plan:\n        if c in '0123456789':\n            assert (prev_c is not None)\n            assert (prev_c in 'NEWS')\n            buffer += c\n        else:\n            plan_expanded += ''.join([prev_c] * int(buffer))\n            plan_expanded += c\n            prev_c = c\n            buffer = '0'\n    if (buffer != '0') and (plan_expanded == ''): # all characters are digits\n        assert (direction is not None)\n        plan_expanded = [direction]\n    return plan_expanded\n\ndef get_next_pos(pos: GraphPos, c: str) -> GraphPos:\n    if c == 'W':\n        return ((pos[0] + 20) % 21, pos[1])\n    elif c == 'N':\n        return (pos[0], (pos[1] + 1) % 21)\n    elif c == 'E':\n        return ((pos[0] + 1) % 21, pos[1])\n    elif c == 'S':\n        return (pos[0], (pos[1] + 20) % 21)\n    return (pos[0], pos[1])\n\ndef action2path(action: str,\n                shipyard: Shipyard,\n                board: Board) -> Tuple[float, GraphPath]:\n    if action.startswith('SPAWN_'):\n        frac = int(action.split('_')[1]) / shipyard.max_spawn if shipyard.max_spawn > 0 else 0\n        return frac, [shipyard.position, shipyard.position]\n    if action.startswith('LAUNCH_'):\n        all_shipyards_pos = [board.shipyards[key].position for key in board.shipyards]\n        frac = int(action.split('_')[1]) / shipyard.ship_count if shipyard.ship_count > 0 else 0\n        return frac, create_path(shipyard.position, action.split('_')[2], None, all_shipyards_pos)\n    raise ValueError(f\"Action {action} is not recognized\")\n\ndef create_path(start: GraphPos,\n                flight_plan: str,\n                direction: str,\n                shipyards_pos: List[GraphPos]) -> GraphPath:\n    plan = expand_plan(flight_plan, direction)\n    pos = (start[0], start[1])\n    path = [pos]\n    for c in plan:\n        next_pos = get_next_pos(pos, c)\n        path.append(next_pos)\n        if next_pos in shipyards_pos:\n            return path\n        if pos == next_pos:\n            assert (c == 'C')\n            return path\n        pos = next_pos\n    end_pos = pos\n    while pos not in shipyards_pos:\n        next_pos = get_next_pos(pos, c)\n        assert (pos != next_pos)\n        path.append(next_pos)\n        pos = next_pos\n        if pos == end_pos:\n            return path\n    return path\n\ndef get_adjacent(i, j, n=21):\n    return [((i + 1) % n,     j),\n            ((i + n - 1) % n, j),\n            (i,     (j + 1) % n),\n            (i, (j + n - 1) % n),]\n\ndef get_edge_attr(edge):\n    attr = []\n    for e0, e1 in zip(edge[0], edge[1]):\n        if e0 == e1:\n            attr.append(0.)\n        elif np.abs(e1 - e0) == 1:\n            attr.append(e1 - e0)\n        else:\n            attr.append((e0 - e1) / np.abs(e1 - e0))\n    return attr\n\ndef build_empty_graph(d=4, n=21) -> nx.Graph:\n    graphs = []\n    for t in range(d):\n        G_part = nx.grid_2d_graph(n, n, periodic=True, create_using=nx.DiGraph)\n        nx.relabel_nodes(G_part, lambda node: (t, node[0], node[1]), copy=False)\n        graphs.append(G_part)\n    G = nx.compose_all(graphs)\n    edges = set()\n    for t in range(d - 1):\n        for k in range(n):\n            for j in range(n):\n                edges.add(((t + 1, k, j), (t, k, j)))\n    G.add_edges_from(edges)\n    # separate spatial and temporal edges\n    nx.set_edge_attributes(G, {edge: get_edge_attr(edge) for edge in G.edges()}, \"edge_attr\")\n    return G\n\ndef get_position_nbrs(position, n=21):\n    return [((position.x + 1) % n,     position.y),\n            ((position.x + n - 1) % n, position.y),\n            (position.x, (position.y + 1) % n),\n            (position.x, (position.y + n - 1) % n),]\n\ndef add_player_features(G, board, key, t):\n    if key == 'plr':\n        player, opponent = board.current_player, board.opponents[0]\n    else:\n        opponent, player  = board.current_player, board.opponents[0]\n    # add player fleet features\n    for f in player.fleets:\n        node = (t, f.position[0], f.position[1])\n        G.nodes[node][key + '_fleets'] = [1.,\n                                          0.01 * f.ship_count,\n                                          f.collection_rate,\n                                          0.01 * f.kore,\n                                          0.]\n    # add damage from enemy fleet\n    for f in opponent.fleets:\n        nbrs = get_position_nbrs(f.position)\n        for nbr in nbrs:\n            node = (t, nbr[0], nbr[1])\n            G.nodes[node][key + '_fleets'][-1] += 0.01 * f.ship_count\n\n    # add shipyards\n    for s in player.shipyards:\n        node = (t, s.position[0], s.position[1])\n        G.nodes[node][key + '_shipyards'] = [1, \n                                                0.01 * s.ship_count,\n                                                0.1 * get_max_trip_length(s.ship_count),\n                                                0.1 * min(s.max_spawn, player.kore // board.configuration.spawn_cost)]\n    return G\n\ndef build_graph(board: Board, d=4, n=21) -> nx.Graph:\n    G = build_empty_graph(d, n)\n\n    nx.set_node_attributes(G, 0., \"step\")\n    nx.set_node_attributes(G, 0., \"kore\")\n    nx.set_node_attributes(G, {node: node[0] for node in G.nodes}, 't')\n    for key in ['plr', 'opn']:\n        nx.set_node_attributes(G, {node: [0., 0., 0., 0., 0.] for node in G.nodes}, key + '_fleets')\n        nx.set_node_attributes(G, {node: [0., 0., 0., 0.] for node in G.nodes}, key + '_shipyards')\n        nx.set_node_attributes(G, {node: 0. for node in G.nodes}, key + '_kore_left')\n\n    for t in range(d):\n        for node in G.nodes():\n            if node[0] == t:\n                G.nodes[node]['step'] = 0.0025 * board.step\n                G.nodes[node]['kore'] = 0.002 * min(board.cells[node[1:]].kore, 500)\n                G.nodes[node]['plr_kore_left'] = 0.001 * board.current_player.kore\n                G.nodes[node]['opn_kore_left'] = 0.001 * board.opponents[0].kore\n        add_player_features(G, board, 'plr', t)\n        add_player_features(G, board, 'opn', t)\n        board = board.next()\n    return G\n\ndef compress_plan(directions):\n    if len(directions) == 0:\n        return ''\n    compressed_plan = [directions[0]]\n    last_num = None\n    for k, c in enumerate(directions[1:], start=1):\n        if directions[k - 1] == c:\n            if last_num:\n                last_num += 1\n            else:\n                last_num = 1\n        else:\n            if last_num:\n                compressed_plan.append(str(last_num))\n                last_num = None\n            compressed_plan.append(c)\n    if last_num:\n        compressed_plan.append(str(last_num))\n    return ''.join(compressed_plan)\n\ndef actions2features(G: nx.Graph, actions: dict, board: Board, data: dict) -> torch.tensor:\n    H = nx.MultiDiGraph()\n    H.add_nodes_from(G.nodes(data=False))\n    \n    nx.set_node_attributes(H, {node: [0.] * 16 for node in H.nodes()},  \"log-p\")\n    nx.set_node_attributes(H, 0., \"log-pth\")\n    nx.set_node_attributes(H, 0., \"log-fs\")\n    nx.set_node_attributes(H, 0., \"log-fl\")\n    nx.set_node_attributes(H, 0,  \"action\")\n\n    flag = False\n    for k, key in enumerate(actions.keys()):\n        if key not in board.shipyards:\n            print(f'Action-key not in board: {key}')\n            continue\n        shipyard = board.shipyards[key]\n        frac, path = action2path(actions[key], shipyard, board)\n        node_id = 0, shipyard.position.x, shipyard.position.y\n        if actions[key].startswith('SPAWN'):\n            H.nodes[node_id][\"action\"] = 15\n        else:\n            bin = get_max_trip_length(shipyard.ship_count)\n            H.nodes[node_id][\"action\"] = min(max(1, bin), 14)\n\n        if data is not None:\n            for logit in ['log-p', 'log-pth', 'log-fs', 'log-fl']:\n                H.nodes[node_id][logit] = data[key][logit]\n        else:\n            H.nodes[node_id]['log-p'] = [-float('inf')] * 16\n            H.nodes[node_id]['log-p'][H.nodes[node_id][\"action\"]] = 0.\n            H.nodes[node_id]['log-pth'] = 0.\n            H.nodes[node_id]['log-fs'] = -float('inf')\n            H.nodes[node_id]['log-fl'] = -float('inf')\n            if frac > 0:\n                if H.nodes[node_id][\"action\"] == 15:\n                    H.nodes[node_id]['log-fs'] = np.log(frac)\n                elif 0 < H.nodes[node_id][\"action\"] < 15:\n                    H.nodes[node_id]['log-fl'] = np.log(frac)\n\n        graph_path = [(0, p[0], p[1]) for p in path]\n\n        # nx.add_path(H, graph_path, action_id=k)\n        if actions[key].startswith('SPAWN_'):\n            space_left = [0.]\n        else:\n            max_trip_len = get_max_trip_length(int(actions[key].split('_')[1]))\n            expanded_plan = expand_plan(actions[key].split('_')[2])\n            space_left = [float(max_trip_len - len(compress_plan(expanded_plan[:i]))) for i in range(len(expanded_plan))]\n            space_left += [np.min(space_left) - 1] * (len(graph_path) - len(space_left) - 1)\n        assert(len(space_left) >= len(graph_path) - 1)\n        for i, (p0, p1) in enumerate(zip(graph_path[:-1], graph_path[1:])):\n            H.add_edge(p0, p1, action_id=k, order=i, space_left=space_left[i])\n            flag = True\n\n    h = from_networkx(H)\n    for key in ['log-p', 'log-pth', 'log-fs', 'log-fl', 'action']:\n        h[key] = h[key].to_sparse()\n\n    if not flag:\n        h.action_id = torch.tensor([])\n        h.order = torch.tensor([])\n        h.space_left = torch.tensor([])\n    return h\n\ndef to_torch_graph(G: nx.Graph) -> Data:\n    # create graph\n    g = from_networkx(G)\n    # all keys -> x\n    if g.num_edges > 0:\n        lst = []\n        for key in ['step', 'kore',\n                    'plr_fleets', 'plr_kore_left', 'plr_shipyards',\n                    'opn_fleets', 'opn_kore_left', 'opn_shipyards']:\n            tensor = g[key].unsqueeze(1) if len(g[key].shape) == 1 else g[key]\n            lst.append(tensor.float())\n            del g[key]\n        g.x = torch.concat(lst, dim=1)\n    else:\n        g.x = None\n        g.direction = None\n    return g\n","metadata":{"execution":{"iopub.status.busy":"2022-07-02T12:42:22.5883Z","iopub.execute_input":"2022-07-02T12:42:22.588683Z","iopub.status.idle":"2022-07-02T12:42:22.603819Z","shell.execute_reply.started":"2022-07-02T12:42:22.588652Z","shell.execute_reply":"2022-07-02T12:42:22.602557Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile -a submission.py\n\nimport torch\nimport torch.nn.functional as F\nfrom torch.distributions.categorical import Categorical\n\ndef get_logits(node, adj, z, hc, net, ships_to_launch, space_left):\n    (h, c) = hc\n    nbrs, _, _ = adj[node.item()].t().coo()\n    x = torch.cat([z[node].unsqueeze(0), \n                   F.one_hot(torch.tensor(space_left).long(), num_classes=16).unsqueeze(0).to(z.device)], dim=1)\n    y = z[nbrs].unsqueeze(0)\n    scores, (h, c) = net.similarity(x, h, c, y)\n    scores = scores.squeeze()\n    build_shipyard_fltr = lambda nbr: (ships_to_launch < 50) and (nbr == node)\n    mask = torch.tensor([build_shipyard_fltr(nbr) for nbr in nbrs]).to(scores.device)\n    scores.masked_fill_(mask, -float('inf'))\n    logits = F.log_softmax(scores, dim=0)\n    return logits, nbrs, (h, c)\n\ndef choose_next(logits, n=1, train=False):\n    if train:\n        return Categorical(logits=logits).sample().item()\n    return torch.argmax(logits) if n == 1 else torch.topk(logits, n).indices\n\ndef get_direction(node, next_node, geopos):\n    if (geopos[next_node][0] == geopos[node][0] + 1) or (geopos[next_node][0] == geopos[node][0] - 20):\n        return 'E'\n    if (geopos[next_node][0] == geopos[node][0] - 1) or (geopos[next_node][0] == geopos[node][0] + 20):\n        return 'W'\n    if (geopos[next_node][1] == geopos[node][1] + 1) or (geopos[next_node][1] == geopos[node][1] - 20):\n        return 'N'\n    if (geopos[next_node][1] == geopos[node][1] - 1) or (geopos[next_node][1] == geopos[node][1] + 20):\n        return 'S'\n    if (torch.abs(geopos[next_node] - geopos[node]).sum() != 0):\n        raise ValueError(f\"Direction cannot be determined for {geopos[node]}->{geopos[next_node]}\")\n    return 'C'\n    \ndef compress_plan(directions, add_last_num=True):\n    compressed_plan = [directions[0]]\n    last_num = None\n    for k, c in enumerate(directions[1:], start=1):\n        if directions[k - 1] == c:\n            if last_num:\n                last_num += 1\n            else:\n                last_num = 1\n        else:\n            if last_num:\n                compressed_plan.append(str(last_num))\n                last_num = None\n            compressed_plan.append(c)\n    if last_num and add_last_num:\n        compressed_plan.append(str(last_num))\n    return ''.join(compressed_plan)\n\ndef greedy_search(start_node, adj, z, h, net, geopos, ships_to_launch, shipyards, train=False):\n    limit = get_max_trip_length(ships_to_launch)\n    if limit == 0:\n        return torch.stack([start_node, start_node]), 'SPAWN', torch.tensor(0.)\n    node = start_node\n    path = [start_node]\n    plan_letters = []\n    plan = ''\n    log_prob = torch.tensor(0.)\n    c = torch.zeros(1, 48).to(z.device)\n    while limit > len(plan):\n        logits, nbrs, (h, c) = get_logits(node, adj, z, (h, c), net, ships_to_launch, limit-len(plan))\n        logit_id = choose_next(logits, n=1, train=train)\n        next_node, logit = nbrs[logit_id], logits[logit_id]\n        log_prob += logit\n        direction = get_direction(node, next_node, geopos)\n        if (next_node in shipyards) or (direction == 'C'): # build shipyard or finish cycle\n            if len(plan_letters) == 0:\n                return torch.stack([start_node, start_node]), 'SPAWN', torch.tensor(0.)\n            return torch.stack(path + [next_node]), compress_plan(plan_letters + [direction]), log_prob\n        plan_letters.append(direction)\n        path.append(next_node)\n        node = next_node\n        plan = compress_plan(plan_letters)\n    return torch.stack(path), plan, log_prob\n\ndef beam_search(start_node, adj, z, h, net, geopos, ships_to_launch, shipyards, train=False, beam_width=4):\n    limit = get_max_trip_length(ships_to_launch)\n    if limit == 0:\n        return torch.stack([start_node, start_node]), 'SPAWN', torch.tensor(0.)\n    paths = [[start_node] for _ in range(beam_width)]\n    plans_letters = [[] for _ in range(beam_width)]\n    plans = [''] * beam_width\n    c = torch.zeros(1, 48).to(z.device)\n    finished_ids = set()\n    zero_space_ids = set()\n    t = 0\n    finished_ids_properly = set()\n    while len(finished_ids) != beam_width:\n        if max([len(plan) for plan in plans]) == 0:\n            # if just started -> get initial arrays\n            logits, nbrs, (h, c) = get_logits(start_node, adj, z, (h, c), net, ships_to_launch, limit)\n            logit_ids = choose_next(logits, n=beam_width, train=train)\n            next_nodes, log_probs = [], []\n            hs, cs = [], []\n            for logit_id in logit_ids:\n                next_nodes.append(nbrs[logit_id])\n                log_probs.append(logits[logit_id])\n                hs.append(h)\n                cs.append(c)\n        else:\n            # otherwise extend every path\n            all_logits, all_logprobs, all_nbrs = [], [], []\n            for i in range(beam_width):\n                if i in finished_ids:\n                    logits = torch.tensor([-float('inf')] * 5)\n                    nbrs = torch.tensor([-1] * 5)\n                    hs[i], cs[i] = None, None\n                else:\n                    if limit - len(plans[i]) < 0:\n                        # must not happen - will raise error\n                        print(i, plans, limit)\n                        print(finished_ids, zero_space_ids)\n\n                    logits, nbrs, (hs[i], cs[i]) = get_logits(next_nodes[i], adj, z, (hs[i], cs[i]), net,\n                                                              ships_to_launch, limit - len(plans[i]))\n                    if i in zero_space_ids:\n                        direction = get_direction(paths[i][-2], paths[i][-1], geopos)\n                        for k, nbr in enumerate(nbrs):\n                            nbr_dir = get_direction(paths[i][-1], nbr, geopos)\n                            if nbr_dir != direction:\n                                logits[k] = torch.tensor(-float('inf'))\n\n                all_logits.append(logits)\n                all_logprobs.append(logits + log_probs[i])\n                all_nbrs.append(nbrs)\n            \n            all_logits = torch.cat(all_logits)\n            all_logprobs = torch.cat(all_logprobs)\n            all_nbrs = torch.cat(all_nbrs)\n            logit_ids = choose_next(all_logprobs, n=beam_width - len(finished_ids), train=train)\n            if len(finished_ids) == beam_width - 1:\n                logit_ids = [logit_ids]\n\n            new_plans_letters, new_paths = [], []\n            new_plans = plans.copy()\n            new_hs, new_cs = [], []\n            new_log_probs = log_probs.copy()\n            new_zero_spaces_ids = set()\n            for i in range(beam_width):\n                new_plans_letters.append(plans_letters[i].copy())\n                new_paths.append(paths[i].copy())\n                new_hs.append(hs[i].clone() if hs[i] is not None else None)\n                new_cs.append(cs[i].clone() if cs[i] is not None else None)\n\n            k_id = 0\n            for i in range(beam_width):\n                if i in finished_ids:\n                    continue\n                logit_id = logit_ids[k_id]\n                row_id = torch.div(logit_id, 5, rounding_mode='floor')\n                if row_id.item() in zero_space_ids:\n                    new_zero_spaces_ids.add(i)\n                next_nodes[i] = all_nbrs[logit_id]\n                new_log_probs[i] = log_probs[row_id] + all_logits[logit_id]\n                # update hidden states\n                new_hs[i] = hs[row_id]\n                new_cs[i] = cs[row_id]\n                # updates plans and paths\n                new_plans_letters[i] = plans_letters[row_id].copy()\n                new_plans[i] = plans[row_id]\n                new_paths[i] = paths[row_id].copy()\n                k_id += 1\n\n            plans_letters = new_plans_letters.copy()\n            plans = new_plans.copy()\n            paths = new_paths.copy()\n            log_probs = new_log_probs.copy()\n            hs = new_hs.copy()\n            cs = new_cs.copy()\n            zero_space_ids = new_zero_spaces_ids.copy()\n            # if len(zero_space_ids) > 0:\n            #     print(\"ZEOO:\", zero_space_ids, finished_ids, [compress_plan(pl) for pl in plans_letters], plans)\n\n        # extend plans\n        for i in range(beam_width):\n            # check that we can continue this path\n            if i in finished_ids:\n                continue\n            \n            if i not in zero_space_ids:\n                direction = get_direction(paths[i][-1], next_nodes[i], geopos)\n                plans_letters[i].append(direction)\n                plans[i] = compress_plan(plans_letters[i])\n\n            paths[i].append(next_nodes[i])\n            t_key = min(t, max(shipyards.keys()))\n            # if ((next_nodes[i] in shipyards)) or (direction == 'C'): # build shipyard or finish cycle\n            if (next_nodes[i] in shipyards[t_key]) or (direction == 'C'): # build shipyard or finish cycle\n                # print('FINISHED:', plans[i], torch.stack(paths[i]), log_probs[i])\n                if i in zero_space_ids:\n                    zero_space_ids.remove(i)\n                finished_ids.add(i)\n                finished_ids_properly.add(i)\n            elif len(plans[i]) == limit:\n                # print('ZERO-SPACED:', plans[i], torch.stack(paths[i]))\n                zero_space_ids.add(i)\n\n        # print('mini-beams:', plans, limit)\n        # print('log-probs:', log_probs)\n        # print('-------------------------------')\n        \n        # this is needed for us not to go into infinity\n        if min(len(finished_ids), len(zero_space_ids)) > 0:\n            max_logprob = max([log_probs[i] for i in finished_ids])\n            for i in zero_space_ids:\n                if log_probs[i] < max_logprob:\n                    finished_ids.add(i)\n            for i in finished_ids:\n                if i in zero_space_ids:\n                    zero_space_ids.remove(i)\n\n        t += 1\n        if t > 100:\n            # print('COUNTER!!!', plans)\n            break\n\n    # now choose the best\n    # print('BEAM PLANS:', plans)\n    for i in range(beam_width):\n        if i not in finished_ids_properly:\n            log_probs[i] = -torch.tensor(float('inf'))\n    logit_id = choose_next(torch.stack(log_probs), n=1, train=train)\n    \n    if plans[logit_id] == 'C':\n        return torch.stack([start_node, start_node]), 'SPAWN', log_probs[logit_id]\n    return paths[logit_id], plans[logit_id], log_probs[logit_id]","metadata":{"execution":{"iopub.status.busy":"2022-07-02T12:42:22.60614Z","iopub.execute_input":"2022-07-02T12:42:22.607044Z","iopub.status.idle":"2022-07-02T12:42:22.627583Z","shell.execute_reply.started":"2022-07-02T12:42:22.606931Z","shell.execute_reply":"2022-07-02T12:42:22.626136Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile -a submission.py\n\nimport torch\nfrom torch_geometric.data import Data, Batch\nimport numpy as np\nimport networkx as nx\nfrom torch_geometric.utils import add_remaining_self_loops, add_self_loops\nfrom torch_geometric.utils.convert import from_networkx\nfrom torch_geometric.utils.undirected import to_undirected\nfrom torch_sparse import SparseTensor\nfrom torch_geometric.nn.glob.glob import global_mean_pool\n\nimport base64\nimport bz2\nimport pickle\n\nfrom kaggle_environments.envs.kore_fleets.helpers import *\n\nclass RLAgent():\n    def __init__(self, rl_net, train=False):\n        self.rl_net = rl_net.eval()\n        self.train = train\n    \n    def logits2actions(self, logits):\n        if self.train:\n            # torch logits -> numpy probabilities\n            probs = torch.exp(logits).detach().numpy()\n            action = np.random.choice(range(len(probs)), p=probs)\n        else:\n            action = torch.argmax(logits)\n        return action\n\n    def raw_outputs(self, board):\n        me = board.current_player\n\n        if not self.train:\n            # check if we can act\n            must_skip = (me.kore < board.configuration.spawn_cost)\n            for shipyard in me.shipyards:\n                must_skip &= (shipyard.ship_count == 0)\n            if must_skip:\n                return (None, None), {}\n            \n        if board.step < 10:\n            data = {}\n            for shipyard in me.shipyards:\n                can_spawn = min(shipyard.max_spawn, me.kore // board.configuration.spawn_cost)\n                ships_to_spawn = int(np.ceil(can_spawn))\n                ships_to_launch = 0\n                logp, logpth = torch.tensor(0.), torch.tensor(0.)\n                data[shipyard.id] = {'plan': 'SPAWN',\n                                     'ships_to_spawn': ships_to_spawn, 'ships_to_launch': ships_to_launch,\n                                     'log-p': logp, 'log-pth': logpth,\n                                     'log-fs': torch.tensor(0.), 'log-fl': torch.tensor(0.)}\n            return (None, None), data\n        \n        d = min(12, 400 - board.step)\n        G = build_graph(board, d)\n        # fix time-spatial positions\n        nx.set_node_attributes(G, {node: node[1:] for node in G.nodes()}, \"pos\")\n        g = to_torch_graph(G)\n        \n        data = {}\n        shipyards = {t: torch.where((g.t == t) & (g.x[:, 8] + g.x[:, 18] == 1))[0] - t * 441 for t in range(d)}\n\n        with torch.no_grad():\n            # get rl-network outputs\n            policy, (z, w) = self.rl_net(Batch.from_data_list([g]))\n\n            # v = torch.tanh(self.rl_net.head_v(global_mean_pool(w, g.batch))).squeeze()\n            # print('VALUE:', board.step, v)\n\n            for shipyard in me.shipyards:\n                mask = (g.t == 0) & (g.pos[:, 0] == shipyard.position.x) & (g.pos[:, 1] == shipyard.position.y)\n                start_node = torch.where(mask)[0][0]\n                if not self.train:\n                    action_type = torch.argmax(policy[start_node])\n                else:\n                    action_type = Categorical(logits=policy[start_node]).sample()\n                logp = F.log_softmax(policy[start_node], dim=0).tolist()\n                head_inpt = torch.cat([z[start_node].unsqueeze(0),\n                                       F.one_hot(torch.tensor([action_type]), num_classes=16)], dim=1)\n                fs = torch.sigmoid(self.rl_net.head_fs(head_inpt)).squeeze()\n                fl = torch.sigmoid(self.rl_net.head_fl(head_inpt)).squeeze()\n                \n                if action_type == 0: # do nothing\n                    plan = ''\n                    ships_to_spawn = 0\n                    ships_to_launch = 0\n                    logpth = torch.tensor(0.)\n                elif action_type == 15:\n                    plan = 'SPAWN'\n                    can_spawn = min(shipyard.max_spawn, me.kore // board.configuration.spawn_cost)\n                    ships_to_spawn = int(np.ceil(fs * can_spawn))\n                    ships_to_launch = 0\n                    logpth = torch.tensor(0.)\n                else:\n                    # create sparse adjacency tensor\n                    adj = SparseTensor.from_edge_index(add_self_loops(g.edge_index)[0])\n                    ships_to_launch = int(np.ceil(fl * shipyard.ship_count))\n                    ships_to_launch = max(ships_to_launch, get_min_ship_num_for_trip(action_type.item()))\n                    ships_to_launch = min(ships_to_launch, get_min_ship_num_for_trip(action_type.item() + 1))\n                    ships_to_launch = min(ships_to_launch, shipyard.ship_count)\n                    ships_to_spawn = 0\n                    \n                    if self.train:\n                        path, plan, logpth = greedy_search(start_node=start_node,\n                                                        adj=adj,\n                                                        z=z,\n                                                        h=w[start_node].unsqueeze(0),\n                                                        net=self.rl_net,\n                                                        geopos=g.pos,\n                                                        ships_to_launch=ships_to_launch,\n                                                        shipyards=shipyards,\n                                                        train=True)\n                    else:                                                        \n                        path, plan, logpth = beam_search(start_node=start_node,\n                                                        adj=adj,\n                                                        z=z,\n                                                        h=w[start_node].unsqueeze(0),\n                                                        net=self.rl_net,\n                                                        geopos=g.pos,\n                                                        ships_to_launch=ships_to_launch,\n                                                        shipyards=shipyards,\n                                                        train=False)\n\n                data[shipyard.id] = {'plan': plan,\n                                     'ships_to_spawn': ships_to_spawn, 'ships_to_launch': ships_to_launch,\n                                     'log-p': logp, 'log-pth': logpth,\n                                     'log-fs': torch.log(fs), 'log-fl': torch.log(fl)}\n                \n        return (g, G), data\n\n    def raw_data_to_actions(self, board, data):\n        me = board.current_player\n        for shipyard in me.shipyards:\n            if shipyard.id in data:\n                plan = data[shipyard.id]['plan']\n                if plan == '':\n                    continue\n                elif plan == 'SPAWN':\n                    ships_to_spawn = data[shipyard.id]['ships_to_spawn']\n                    shipyard.next_action = ShipyardAction.spawn_ships(ships_to_spawn)\n                else:\n                    ships_to_launch = data[shipyard.id]['ships_to_launch']\n                    shipyard.next_action = ShipyardAction.launch_fleet_with_flight_plan(ships_to_launch, plan)\n        return me.next_actions\n\n    def __call__(self, obs, config):\n        board = Board(obs, config)\n        _, data = self.raw_outputs(board)\n        return self.raw_data_to_actions(board, data)\n    \n    \nPARAM = b'XXXXXXXXXX'\n\nstate_dict = pickle.loads(bz2.decompress(base64.b64decode(PARAM)))\n\nnet = KoreNet().eval()\nnet.load_state_dict(state_dict, strict=True)\nrl_agent = RLAgent(net, False)\n\ndef agent(obs, config):\n    return rl_agent(obs, config)","metadata":{"execution":{"iopub.status.busy":"2022-07-02T12:42:22.684374Z","iopub.execute_input":"2022-07-02T12:42:22.684949Z","iopub.status.idle":"2022-07-02T12:42:22.697664Z","shell.execute_reply.started":"2022-07-02T12:42:22.684917Z","shell.execute_reply":"2022-07-02T12:42:22.696456Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport base64\nimport bz2\nimport pickle\n\nweights = torch.load(\"../input/koreweights//kore.net\", map_location='cpu')\nPARAM = base64.b64encode(bz2.compress(pickle.dumps(weights)))\n\n# Read in the submission file\nwith open('submission.py',) as file:\n    filedata = file.read()\n\n# Replace the target string\nfiledata = filedata.replace('XXXXXXXXXX', PARAM.decode(\"utf-8\") )\n\n# Write the file out again\nwith open('submission.py','w') as file:\n    file.write(filedata)","metadata":{"execution":{"iopub.status.busy":"2022-07-02T12:42:35.251957Z","iopub.execute_input":"2022-07-02T12:42:35.252357Z","iopub.status.idle":"2022-07-02T12:42:35.388405Z","shell.execute_reply.started":"2022-07-02T12:42:35.252324Z","shell.execute_reply":"2022-07-02T12:42:35.387038Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from kaggle_environments import make, evaluate\n\nenv = make(\"kore_fleets\", debug=True)\n\nenv.run([\"submission.py\", \"balanced\"])\nenv.render(mode=\"ipython\", width=720, height=680)","metadata":{"execution":{"iopub.status.busy":"2022-07-02T12:42:36.555343Z","iopub.execute_input":"2022-07-02T12:42:36.556032Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}