File size: 1,794 Bytes
bdfb884
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Shared helpers (same as the neural-DDR project): bit<->int, MLP, verify, train."""
from __future__ import annotations
import torch
import torch.nn as nn

DEV = "cuda" if torch.cuda.is_available() else "cpu"


def bits_of(v: int, n: int) -> list[int]:
    return [(v >> k) & 1 for k in range(n)]


def int_of(bits) -> int:
    return sum((1 << k) for k, b in enumerate(bits) if b > 0)


def pm(bits) -> torch.Tensor:
    return torch.tensor([1.0 if b else -1.0 for b in bits], dtype=torch.float32)


def mlp(inp: int, out: int, h: int = 256, layers: int = 2) -> nn.Sequential:
    mods = [nn.Linear(inp, h), nn.GELU()]
    for _ in range(layers - 1):
        mods += [nn.Linear(h, h), nn.GELU()]
    mods += [nn.Linear(h, out)]
    return nn.Sequential(*mods)


@torch.no_grad()
def verify(net, X, Ybits) -> tuple[int, int]:
    net = net.to("cpu")
    pred = (net(X) > 0).int()
    ok = (pred == Ybits.int()).all(dim=1).sum().item()
    return ok, X.shape[0]


def train(net, X, Ybits, steps=8000, lr=2e-3, tag="", report=2000):
    net = net.to(DEV)
    Xd, Yd = X.to(DEV), Ybits.to(DEV)
    opt = torch.optim.Adam(net.parameters(), lr=lr)
    sch = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=steps)
    lossfn = nn.BCEWithLogitsLoss()
    for e in range(steps):
        opt.zero_grad()
        loss = lossfn(net(Xd), Yd)
        loss.backward(); opt.step(); sch.step()
        if e % report == 0 or e == steps - 1:
            ok, tot = verify(net, X, Ybits); net.to(DEV)
            print(f"    [{tag}] epoch {e:5d}  loss {loss.item():.2e}  verified {ok}/{tot}")
            if ok == tot:
                print(f"    [{tag}] -> N/N"); break
    return net.to("cpu")


def run8(net, v: int) -> int:
    return int_of((net(pm(bits_of(v, 8)).unsqueeze(0))[0] > 0).int().tolist())