Download cdpreserve/common.py from NeuralVerified/neural-cd-preserve: direct link, hf CLI and curl.
- Browser
- Download file 1.79 kB
-
https://huggingface.co/NeuralVerified/neural-cd-preserve/resolve/de15a9f1e628d5f702dc74ff6fffd227d256be68/cdpreserve/common.py
- Command line
-
hf download hf://NeuralVerified/neural-cd-preserve@de15a9f1e628d5f702dc74ff6fffd227d256be68/cdpreserve/common.py
-
curl -L -o common.py https://huggingface.co/NeuralVerified/neural-cd-preserve/resolve/de15a9f1e628d5f702dc74ff6fffd227d256be68/cdpreserve/common.py
1.79 kB
| """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) | |
| 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()) | |