| from __future__ import annotations |
| import math |
| import torch |
| import torch.nn as nn |
| import torch.nn.functional as F |
| from torch.utils.checkpoint import checkpoint |
|
|
| NUM_FROM_TO = 4096 |
| NUM_PROMO = 5 |
| MAX_PLIES = 96 |
|
|
|
|
| class RMSNorm(nn.Module): |
|
|
| def __init__(self, dim: int, eps: float = 1e-05): |
| super().__init__() |
| self.eps = eps |
| self.weight = nn.Parameter(torch.ones(dim)) |
|
|
| def forward(self, x: torch.Tensor) -> torch.Tensor: |
| norm = x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) |
| return norm * self.weight |
|
|
|
|
| def parallel_scan(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: |
| L = a.shape[1] |
| d = 1 |
| while d < L: |
| a_prev, b_prev = (a[:, :-d], b[:, :-d]) |
| a_cur, b_cur = (a[:, d:], b[:, d:]) |
| new_a = a_cur * a_prev |
| new_b = a_cur * b_prev + b_cur |
| a = torch.cat([a[:, :d], new_a], dim=1) |
| b = torch.cat([b[:, :d], new_b], dim=1) |
| d *= 2 |
| return b |
|
|
|
|
| class S6Block(nn.Module): |
|
|
| def __init__(self, dim: int, state_dim: int = 16, expand: int = 2): |
| super().__init__() |
| inner_dim = dim * expand |
| self.dim = dim |
| self.inner_dim = inner_dim |
| self.state_dim = state_dim |
| self.in_proj = nn.Linear(dim, inner_dim * 2, bias=False) |
| self.x_proj = nn.Linear(inner_dim, state_dim * 2 + inner_dim, bias=False) |
| self.dt_bias = nn.Parameter(torch.zeros(inner_dim)) |
| A = torch.arange(1, state_dim + 1, dtype=torch.float32).unsqueeze(0).repeat(inner_dim, 1) |
| self.A_log = nn.Parameter(torch.log(A)) |
| self.D = nn.Parameter(torch.ones(inner_dim)) |
| self.out_proj = nn.Linear(inner_dim, dim, bias=False) |
|
|
| def forward(self, x: torch.Tensor) -> torch.Tensor: |
| B_, L, _ = x.shape |
| xz = self.in_proj(x) |
| x_in, gate = xz.chunk(2, dim=-1) |
| x_in = F.silu(x_in) |
| x_dbl = self.x_proj(x_in) |
| Bmat, Cmat, delta_raw = torch.split( |
| x_dbl, [self.state_dim, self.state_dim, self.inner_dim], dim=-1 |
| ) |
| delta = F.softplus(delta_raw + self.dt_bias) |
| A = -torch.exp(self.A_log) |
| A_bar = torch.exp(delta.unsqueeze(-1) * A.view(1, 1, self.inner_dim, self.state_dim)) |
| Bx = (delta * x_in).unsqueeze(-1) * Bmat.unsqueeze(2) |
| h = parallel_scan(A_bar, Bx) |
| y = (h * Cmat.unsqueeze(2)).sum(-1) + self.D * x_in |
| y = y * F.silu(gate) |
| return self.out_proj(y) |
|
|
| def step( |
| self, x_t: torch.Tensor, h_prev: torch.Tensor | None |
| ) -> tuple[torch.Tensor, torch.Tensor]: |
| xz = self.in_proj(x_t) |
| x_in, gate = xz.chunk(2, dim=-1) |
| x_in = F.silu(x_in) |
| x_dbl = self.x_proj(x_in) |
| Bmat, Cmat, delta_raw = torch.split( |
| x_dbl, [self.state_dim, self.state_dim, self.inner_dim], dim=-1 |
| ) |
| delta = F.softplus(delta_raw + self.dt_bias) |
| A = -torch.exp(self.A_log) |
| A_bar = torch.exp(delta.unsqueeze(-1) * A.unsqueeze(0)) |
| Bx = (delta * x_in).unsqueeze(-1) * Bmat.unsqueeze(1) |
| if h_prev is None: |
| h_prev = x_t.new_zeros(x_t.shape[0], self.inner_dim, self.state_dim) |
| h_new = A_bar * h_prev + Bx |
| y = (h_new * Cmat.unsqueeze(1)).sum(-1) + self.D * x_in |
| y = y * F.silu(gate) |
| return (self.out_proj(y), h_new) |
|
|
|
|
| class MambaBlock(nn.Module): |
|
|
| def __init__(self, dim: int, state_dim: int = 16, expand: int = 2): |
| super().__init__() |
| self.norm = RMSNorm(dim) |
| self.ssm = S6Block(dim, state_dim=state_dim, expand=expand) |
|
|
| def forward(self, x: torch.Tensor) -> torch.Tensor: |
| return x + self.ssm(self.norm(x)) |
|
|
| def step( |
| self, x_t: torch.Tensor, h_prev: torch.Tensor | None |
| ) -> tuple[torch.Tensor, torch.Tensor]: |
| y, h_new = self.ssm.step(self.norm(x_t), h_prev) |
| return (x_t + y, h_new) |
|
|
|
|
| class ChessMamba(nn.Module): |
|
|
| def __init__( |
| self, |
| dim: int = 256, |
| depth: int = 8, |
| state_dim: int = 16, |
| expand: int = 2, |
| max_plies: int = MAX_PLIES, |
| use_checkpoint: bool = True, |
| ): |
| super().__init__() |
| self.dim = dim |
| self.max_plies = max_plies |
| self.use_checkpoint = use_checkpoint |
| self.from_embed = nn.Embedding(64, dim) |
| self.to_embed = nn.Embedding(64, dim) |
| self.promo_embed = nn.Embedding(NUM_PROMO, dim) |
| self.pos_embed = nn.Embedding(max_plies + 1, dim) |
| self.start_token = nn.Parameter(torch.zeros(1, 1, dim)) |
| self.blocks = nn.ModuleList([MambaBlock(dim, state_dim, expand) for _ in range(depth)]) |
| self.norm_f = RMSNorm(dim) |
| self.policy_head = nn.Linear(dim, NUM_FROM_TO) |
| self.promo_head = nn.Linear(dim, NUM_PROMO) |
| self.value_head = nn.Linear(dim, 1) |
| nn.init.normal_(self.from_embed.weight, std=0.02) |
| nn.init.normal_(self.to_embed.weight, std=0.02) |
| nn.init.normal_(self.promo_embed.weight, std=0.02) |
| nn.init.normal_(self.pos_embed.weight, std=0.02) |
|
|
| def embed_moves( |
| self, from_ids: torch.Tensor, to_ids: torch.Tensor, promo_ids: torch.Tensor |
| ) -> torch.Tensor: |
| B_ = from_ids.shape[0] |
| start = self.start_token.expand(B_, 1, -1) |
| if from_ids.shape[1] == 0: |
| tok = start |
| else: |
| mv = self.from_embed(from_ids) + self.to_embed(to_ids) + self.promo_embed(promo_ids) |
| tok = torch.cat([start, mv], dim=1) |
| positions = torch.arange(tok.shape[1], device=tok.device).unsqueeze(0) |
| return tok + self.pos_embed(positions) |
|
|
| def forward( |
| self, |
| from_ids: torch.Tensor, |
| to_ids: torch.Tensor, |
| promo_ids: torch.Tensor, |
| lengths: torch.Tensor | None = None, |
| ): |
| x = self.embed_moves(from_ids, to_ids, promo_ids) |
| for block in self.blocks: |
| if self.use_checkpoint and self.training: |
| x = checkpoint(block, x, use_reentrant=False) |
| else: |
| x = block(x) |
| x = self.norm_f(x) |
| if lengths is None: |
| pooled = x[:, -1] |
| else: |
| idx = lengths.view(-1, 1, 1).expand(-1, 1, x.shape[-1]) |
| pooled = x.gather(1, idx).squeeze(1) |
| policy_logits = self.policy_head(pooled) |
| promo_logits = self.promo_head(pooled) |
| value = torch.tanh(self.value_head(pooled)) |
| return (policy_logits, promo_logits, value) |
|
|
| @torch.no_grad() |
| def init_incremental(self, device: torch.device | str = "cpu"): |
| pos = torch.zeros(1, dtype=torch.long, device=device) |
| x = self.start_token.view(1, self.dim).to(device) + self.pos_embed(pos) |
| block_states = [] |
| for block in self.blocks: |
| x, s = block.step(x, None) |
| block_states.append(s) |
| x = self.norm_f(x) |
| outputs = self._heads(x) |
| return ((1, block_states), outputs) |
|
|
| @torch.no_grad() |
| def step_move(self, from_sq: int, to_sq: int, promo_id: int, state): |
| pos_idx, block_states = state |
| device = self.from_embed.weight.device |
| pos_idx_clamped = min(pos_idx, self.max_plies) |
| idx = lambda v: torch.tensor([v], device=device) |
| x = ( |
| self.from_embed(idx(from_sq)) |
| + self.to_embed(idx(to_sq)) |
| + self.promo_embed(idx(promo_id)) |
| + self.pos_embed(idx(pos_idx_clamped)) |
| ) |
| new_block_states = [] |
| for block, s in zip(self.blocks, block_states): |
| x, ns = block.step(x, s) |
| new_block_states.append(ns) |
| x = self.norm_f(x) |
| outputs = self._heads(x) |
| return ((pos_idx + 1, new_block_states), outputs) |
|
|
| @torch.no_grad() |
| def build_incremental_state( |
| self, |
| from_list: list[int], |
| to_list: list[int], |
| promo_list: list[int], |
| device: torch.device | str = "cpu", |
| ): |
| state, outputs = self.init_incremental(device) |
| for f, t, p in zip(from_list, to_list, promo_list): |
| state, outputs = self.step_move(f, t, p, state) |
| return (state, outputs) |
|
|
| def _heads(self, x: torch.Tensor): |
| policy_logits = self.policy_head(x) |
| promo_logits = self.promo_head(x) |
| value = torch.tanh(self.value_head(x)) |
| return (policy_logits, promo_logits, value) |
|
|
|
|
| def count_params(model: nn.Module) -> int: |
| return sum((p.numel() for p in model.parameters())) |
|
|
|
|
| if __name__ == "__main__": |
| m = ChessMamba(dim=256, depth=8) |
| print(f"params: {count_params(m):,}") |
| B_, L = (4, 20) |
| from_ids = torch.randint(0, 64, (B_, L)) |
| to_ids = torch.randint(0, 64, (B_, L)) |
| promo_ids = torch.zeros(B_, L, dtype=torch.long) |
| pl, pr, v = m(from_ids, to_ids, promo_ids) |
| print(pl.shape, pr.shape, v.shape) |
|
|