File size: 3,265 Bytes
a0a9254
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
"""Decision network extracted from NanoJev; no trainer or game imports."""
import torch
from torch import nn
import torch.nn.functional as F

class DecisionModel(nn.Module):
    def __init__(self, backbone, set_head):
        super().__init__()
        self.backbone = backbone
        hidden = backbone.config.hidden_size
        self.norm = nn.LayerNorm(hidden)
        self.scalar = nn.Linear(hidden, 1)  # Nonzero random initialization avoids a dead first step.
        nn.init.normal_(self.scalar.weight, std=0.02)
        nn.init.zeros_(self.scalar.bias)
        self.set_head = set_head
        if set_head == 'attention':
            self.set_project = nn.Linear(hidden + 1, 128)
            self.set_attention = nn.MultiheadAttention(128, 4, dropout=0.0, batch_first=True)
            self.set_output = nn.Linear(128, 1)
            # Only the final residual projection starts at zero; its upstream layers are nonzero.
            nn.init.zeros_(self.set_output.weight)
            nn.init.zeros_(self.set_output.bias)

    def forward(self, examples, pad_token):
        paths = [ids for ex in examples for ids in ex['leaf_tokens']]
        device = self.scalar.weight.device
        lengths = torch.tensor([len(ids) for ids in paths], device=device)
        width = int(lengths.max())
        tokens = torch.full((len(paths), width), pad_token, dtype=torch.long, device=device)
        for i, ids in enumerate(paths):
            tokens[i, :len(ids)] = torch.tensor(ids, device=device)
        attention = torch.arange(width, device=device)[None, :] < lengths[:, None]
        hidden = self.backbone(input_ids=tokens, attention_mask=attention,
                               use_cache=False).last_hidden_state
        leaves = hidden[torch.arange(len(paths), device=device), lengths-1]
        kmax = max(len(ex['candidate_ids']) for ex in examples)
        h = leaves.new_zeros((len(examples), kmax, leaves.shape[-1]))
        valid = torch.zeros((len(examples), kmax), dtype=torch.bool, device=device)
        offset = 0
        for i, ex in enumerate(examples):
            n = len(ex['leaf_tokens'])
            h[i, :n] = leaves[offset:offset+n]
            valid[i, :len(ex['candidate_ids'])] = True
            offset += n
        h = self.norm(h)
        z = self.scalar(h).squeeze(-1).float()
        choice = torch.tensor([i for i, ex in enumerate(examples) if ex['type'] == 'choice'], device=device)
        if self.set_head == 'attention' and len(choice):
            log_k = valid[choice].sum(-1).float().log()[:, None, None].expand(-1, kmax, 1)
            u = self.set_project(torch.cat([h[choice], log_k.to(h.dtype)], dim=-1))
            mixed, _ = self.set_attention(u, u, u, key_padding_mask=~valid[choice], need_weights=False)
            delta = self.set_output(torch.tanh(u + mixed)).squeeze(-1).float()
            z = z.index_add(0, choice, delta)
        # Boolean has one semantic path and one scalar, representing logits [0,z].
        out = []
        for i, ex in enumerate(examples):
            if ex['type'] == 'boolean':
                out.append(F.pad(torch.stack([z[i, 0] * 0, z[i, 0]]), (0, kmax-2)))
            else:
                out.append(z[i])
        return torch.stack(out).masked_fill(~valid, -1e9), valid