pi-predicter / model.py
akaruineko's picture
Upload folder using huggingface_hub
65a273f verified
Raw
History Blame Contribute Delete
3.19 kB
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)