File size: 1,846 Bytes
4b77beb
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Tier-1/2 specialist: trained classifier over the finite small-prime domain.

Weights are trained (training/train_t2_enum.py) from random init on the
complete enumeration of (a mod p, b mod p, p) for all primes < 256 and
verified exact on that full domain. At inference the network's argmax IS the
answer digit; there is no arithmetic here.
"""

from __future__ import annotations

import torch
import torch.nn as nn


class T2Net(nn.Module):
    def __init__(self, d: int = 256, hidden: int = 2048):
        super().__init__()
        self.emb_a = nn.Embedding(256, d)
        self.emb_b = nn.Embedding(256, d)
        self.emb_p = nn.Embedding(256, d)
        self.net = nn.Sequential(
            nn.Linear(3 * d, hidden),
            nn.GELU(),
            nn.Linear(hidden, hidden),
            nn.GELU(),
            nn.Linear(hidden, hidden),
            nn.GELU(),
            nn.Linear(hidden, 256),
        )

    def forward(self, ra, rb, p):
        h = torch.cat([self.emb_a(ra), self.emb_b(rb), self.emb_p(p)], dim=-1)
        return self.net(h)


class T2EnumSpecialist:
    def __init__(self, weights_path, device):
        blob = torch.load(weights_path, map_location=device, weights_only=True)
        self.model = T2Net(**blob["config"]).to(device)
        self.model.load_state_dict(blob["state_dict"])
        self.model.eval()
        self.device = device

    @torch.no_grad()
    def predict_batch(self, batch) -> list[list[int]]:
        ra = torch.tensor([r_a for r_a, _, _ in batch], dtype=torch.long, device=self.device)
        rb = torch.tensor([r_b for _, r_b, _ in batch], dtype=torch.long, device=self.device)
        p = torch.tensor([p_enc["p"] for _, _, p_enc in batch], dtype=torch.long, device=self.device)
        preds = self.model(ra, rb, p).argmax(-1).tolist()
        return [[int(v)] for v in preds]