File size: 3,189 Bytes
65a273f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
import math
import torch
import torch.nn as nn


def fourier_encode(positions: torch.Tensor, max_pos: int, num_frequencies: int = 16) -> torch.Tensor:
    """
    Fourier positional encoding.
    positions: [B] integer positions (0..max_pos)
    returns: [B, 1+2*F]  (normalized scalar + sin/cos features)
    """
    # Normalize to [0, 1]  (or [-1,1] style)
    # Use float32
    pos_norm = positions.float() / float(max_pos)  # [B]
    # pos_norm unsqueezed for broadcasting
    feats = [pos_norm.unsqueeze(-1)]  # [B,1]
    for i in range(num_frequencies):
        freq = 2 ** i
        # scale frequency by pi to capture high frequency variation
        # sin(2^i * pi * pos_norm)
        feats.append(torch.sin(freq * math.pi * pos_norm).unsqueeze(-1))
        feats.append(torch.cos(freq * math.pi * pos_norm).unsqueeze(-1))
    return torch.cat(feats, dim=-1)  # [B, 1+2F]


class PiPredictor(nn.Module):
    def __init__(
        self,
        max_pos: int = 10000,
        num_frequencies: int = 16,
        hidden_dims: list = None,
        dropout: float = 0.1,
        encoding: str = "fourier",  # "fourier" | "scalar" | "embedding"
        embedding_dim: int = 64,
    ):
        super().__init__()
        self.max_pos = max_pos
        self.num_frequencies = num_frequencies
        self.encoding = encoding

        if hidden_dims is None:
            hidden_dims = [256, 256, 128]

        if encoding == "fourier":
            input_dim = 1 + 2 * num_frequencies
        elif encoding == "scalar":
            input_dim = 1
        elif encoding == "embedding":
            input_dim = embedding_dim
            self.embedding = nn.Embedding(max_pos + 5, embedding_dim)
        else:
            raise ValueError(f"unknown encoding {encoding}")

        layers = []
        prev_dim = input_dim
        for h in hidden_dims:
            layers.append(nn.Linear(prev_dim, h))
            layers.append(nn.LayerNorm(h))
            layers.append(nn.GELU())
            if dropout > 0:
                layers.append(nn.Dropout(dropout))
            prev_dim = h

        self.backbone = nn.Sequential(*layers)
        self.head = nn.Linear(prev_dim, 10)  # 10 digits

        # store for checkpoint reloading
        self.hidden_dims = hidden_dims
        self.dropout = dropout
        self.embedding_dim = embedding_dim

    def encode(self, positions: torch.Tensor) -> torch.Tensor:
        if self.encoding == "fourier":
            return fourier_encode(positions, self.max_pos, self.num_frequencies)
        elif self.encoding == "scalar":
            return (positions.float() / float(self.max_pos)).unsqueeze(-1)
        elif self.encoding == "embedding":
            return self.embedding(positions.long())
        else:
            raise ValueError

    def forward(self, positions: torch.Tensor) -> torch.Tensor:
        """
        positions: [B] long
        returns logits [B, 10]
        """
        x = self.encode(positions)
        x = self.backbone(x)
        logits = self.head(x)
        return logits

    def predict(self, positions: torch.Tensor) -> torch.Tensor:
        logits = self.forward(positions)
        return torch.argmax(logits, dim=-1)