fast-modular-nn / model.py
robertboy18's picture
Publish SAIR Modular Arithmetic Challenge submission
956b8cf verified
Raw
History Blame Contribute Delete
6.17 kB
from __future__ import annotations
from pathlib import Path
import torch
from torch import nn
from modchallenge.interface.base_model import ModularMultiplicationModel
def ints_to_bits(values: list[int], device: torch.device, width: int) -> torch.Tensor:
"""Convert nonnegative Python integers to fixed-width, MSB-first bits."""
byte_width = (width + 7) // 8
packed_bytes = bytearray().join(
int(value).to_bytes(byte_width, "big") for value in values
)
packed = torch.frombuffer(packed_bytes, dtype=torch.uint8)
packed = packed.reshape(len(values), byte_width).to(device=device)
shifts = torch.arange(7, -1, -1, device=device)
bits = ((packed[:, :, None] >> shifts) & 1).reshape(len(values), byte_width * 8)
return bits[:, byte_width * 8 - width :]
def ints_to_digits(
values: list[int],
radix: int,
width: int,
device: torch.device,
) -> torch.Tensor:
bits_per_digit = radix.bit_length() - 1
mask = radix - 1
rows = [
[
(value >> (bits_per_digit * position)) & mask
for position in range(width - 1, -1, -1)
]
for value in values
]
return torch.tensor(rows, dtype=torch.long, device=device)
class TransitionCell(nn.Module):
def __init__(
self,
radix: int = 2,
dmodel: int = 32,
hidden: int = 64,
layers: int = 2,
bidirectional: bool = True,
) -> None:
super().__init__()
self.input_projection = nn.Linear(3, dmodel)
self.digit_embedding = nn.Embedding(radix, dmodel)
self.recurrent = nn.GRU(
dmodel,
hidden,
num_layers=layers,
batch_first=True,
bidirectional=bidirectional,
)
directions = 2 if bidirectional else 1
self.output = nn.Linear(directions * hidden, 1)
def forward(self, features: torch.Tensor, digits: torch.Tensor) -> torch.Tensor:
embedded = self.input_projection(features)
embedded = embedded + self.digit_embedding(digits)[:, None, :]
hidden, _ = self.recurrent(embedded)
return self.output(hidden).squeeze(-1)
class FastModularModel(ModularMultiplicationModel):
def __init__(self) -> None:
self.model: TransitionCell | None = None
self.device: torch.device | None = None
self.radix = 2
self.max_width = 2048
self.bits_per_digit = 1
def load(self, model_dir: str) -> None:
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
checkpoint = torch.load(
Path(model_dir) / "weights.pt",
map_location=self.device,
weights_only=True,
)
config = checkpoint["config"]
self.radix = int(config["radix"])
self.bits_per_digit = self.radix.bit_length() - 1
self.max_width = int(checkpoint["max_width"])
self.model = TransitionCell(**config)
self.model.load_state_dict(checkpoint["state_dict"])
self.model.to(self.device)
self.model.recurrent.flatten_parameters()
self.model.eval()
if self.device.type == "cuda":
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True
def preprocess_a(self, a: str) -> int:
return int(a)
def preprocess_b(self, b: str) -> int:
return int(b)
def preprocess_p(self, p: str) -> int:
return int(p)
@torch.inference_mode()
def predict_digits(self, a_enc, b_enc, p_enc) -> list[int]:
return self.predict_digits_batch([(a_enc, b_enc, p_enc)])[0]
@torch.inference_mode()
def predict_digits_batch(self, inputs) -> list[list[int]]:
output: list[list[int]] = [[0] for _ in inputs]
indices: list[int] = []
a_values: list[int] = []
b_values: list[int] = []
moduli: list[int] = []
for index, (a_enc, b_enc, p_enc) in enumerate(inputs):
p = int(p_enc)
if p < 2 or p.bit_length() > self.max_width:
continue
indices.append(index)
a_values.append(int(a_enc) % p)
b_values.append(int(b_enc) % p)
moduli.append(p)
if not indices:
return output
assert self.device is not None
effective_width = max(p.bit_length() for p in moduli)
effective_width = min(self.max_width, max(8, ((effective_width + 7) // 8) * 8))
digit_width = max(
1,
(max(value.bit_length() for value in b_values) + self.bits_per_digit - 1)
// self.bits_per_digit,
)
p_bits = ints_to_bits(moduli, self.device, effective_width).float()
x_bits = ints_to_bits(a_values, self.device, effective_width).float()
control_digits = ints_to_digits(
b_values,
self.radix,
digit_width,
self.device,
)
state = torch.zeros(
(len(indices), effective_width),
dtype=torch.float32,
device=self.device,
)
for position in range(digit_width):
state = self._step(
state,
x_bits,
p_bits,
control_digits[:, position],
)
rows = state.to(dtype=torch.int64).tolist()
for row_index, output_index in enumerate(indices):
output[output_index] = [int(bit) for bit in rows[row_index]]
return output
def max_batch_size(self) -> int:
return 256
def _step(
self,
state: torch.Tensor,
multiplicand: torch.Tensor,
modulus: torch.Tensor,
digit: torch.Tensor,
) -> torch.Tensor:
assert self.model is not None
features = torch.stack((state, multiplicand, modulus), dim=-1)
if self.device is not None and self.device.type == "cuda":
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
logits = self.model(features, digit)
else:
logits = self.model(features, digit)
return (logits.float() > 0).float()