mac-cell-sair / model.py
suii0x's picture
Upload 10 files
17abf28 verified
Raw
History Blame Contribute Delete
3.76 kB
"""ModularMultiplicationModel implementation: fixed 2-phase schedule around
a trained ReductionCell. This module (and everything it imports from
`mac_cell`) is the inference path -- no hand-coded reduction against the
challenge prime `p` appears anywhere below or in `cell.py` / `schedule.py` /
`digits.py`. Only base-R / base-2 digit decomposition of individual
arguments (explicitly permitted inside a per-argument preprocessing hook),
embeddings, GRU, linear layers, and argmax.
At packaging time this file is copied verbatim into the submission
directory as `model.py` (top-level, matching `entry_class = "model.MacCellModel"`),
alongside a copied `mac_cell/` subpackage -- the submission is fully
self-contained and reads nothing outside its own directory.
"""
from __future__ import annotations
import json
from pathlib import Path
from typing import Any
import torch
from mac_cell.cell import P_BITS, ReductionCell
from mac_cell.digits import digits_needed, int_to_digits_truncating
from mac_cell.schedule import rollout
from modchallenge.interface.base_model import ModularMultiplicationModel
# Fallback defaults matching task 2's T1/T2-only domain (p < 2**8, operands
# up to 48 bits). `model_config.json` overrides both per submission -- task
# 3 ships p_bits=16 (T3 extends p up to 2**16) and operand_bits=64 (T3's
# operand_bits). Kept only so a task-2-era config without these keys still
# loads with its original behavior.
DEFAULT_P_BITS = P_BITS
DEFAULT_OPERAND_BITS = 48
class MacCellModel(ModularMultiplicationModel):
def load(self, model_dir: str) -> None:
model_dir_path = Path(model_dir)
config = json.loads((model_dir_path / "model_config.json").read_text())
self.radix = config["radix"]
self.p_bits = config.get("p_bits", DEFAULT_P_BITS)
operand_bits = config.get("operand_bits", DEFAULT_OPERAND_BITS)
self.state_digits = digits_needed(2**self.p_bits, self.radix)
self.operand_digits = digits_needed(2**operand_bits, self.radix)
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
self.cell = ReductionCell(
radix=self.radix,
state_digits=self.state_digits,
hidden_size=config["hidden_size"],
digit_embed_dim=config["digit_embed_dim"],
num_layers=config["num_layers"],
p_bits=self.p_bits,
).to(self.device)
state_dict = torch.load(
model_dir_path / "weights.pt", map_location=self.device, weights_only=True
)
self.cell.load_state_dict(state_dict)
self.cell.eval()
def preprocess_a(self, a: str) -> Any:
return int_to_digits_truncating(int(a), self.radix, self.operand_digits)
def preprocess_b(self, b: str) -> Any:
return int_to_digits_truncating(int(b), self.radix, self.operand_digits)
def preprocess_p(self, p: str) -> Any:
return int_to_digits_truncating(int(p), 2, self.p_bits)
def predict_digits(self, a_enc: Any, b_enc: Any, p_enc: Any) -> list[int]:
return self.predict_digits_batch([(a_enc, b_enc, p_enc)])[0]
@torch.no_grad()
def predict_digits_batch(
self, inputs: list[tuple[Any, Any, Any]]
) -> list[list[int]]:
a_digits = torch.tensor(
[a for a, _, _ in inputs], dtype=torch.long, device=self.device
)
b_digits = torch.tensor(
[b for _, b, _ in inputs], dtype=torch.long, device=self.device
)
p_bits = torch.tensor(
[p for _, _, p in inputs], dtype=torch.float32, device=self.device
)
state = rollout(self.cell, a_digits, b_digits, p_bits)
return state.tolist()
def max_batch_size(self) -> int:
return 256