""" ChessResNet: a ResNet-style policy-value network for chess. Architecture: Stem: Conv3x3 18→channels, GroupNorm, GELU Tower: N residual blocks (Conv3x3→GN→GELU→Conv3x3→GN→+→GELU) Policy head: spatial Conv1x1 → 320 channels → reshape to [B, 20480] Value head: Conv1x1 → 32 → Flatten → Linear 256 → Linear 1 → tanh Default config (channels=256, blocks=24) yields ~29M parameters. """ import math from typing import Optional import torch from torch import nn class ResidualBlock(nn.Module): """Pre-activation residual block with GroupNorm.""" def __init__(self, channels: int, norm_groups: int = 32): super().__init__() self.conv1 = nn.Conv2d(channels, channels, 3, padding=1, bias=False) self.norm1 = nn.GroupNorm(norm_groups, channels) self.act1 = nn.GELU() self.conv2 = nn.Conv2d(channels, channels, 3, padding=1, bias=False) self.norm2 = nn.GroupNorm(norm_groups, channels) self.act2 = nn.GELU() def forward(self, x: torch.Tensor) -> torch.Tensor: residual = x out = self.conv1(x) out = self.norm1(out) out = self.act1(out) out = self.conv2(out) out = self.norm2(out) out = out + residual out = self.act2(out) return out class ChessResNet(nn.Module): """Policy-value ResNet for chess. Args: channels: Number of filters in the residual tower (default: 256). blocks: Number of residual blocks (default: 20). num_actions: Size of action space (default: 20480). norm_groups: Number of groups for GroupNorm (default: 32). """ def __init__( self, channels: int = 256, blocks: int = 24, num_actions: int = 20480, norm_groups: int = 32, ): super().__init__() self.channels = channels self.blocks = blocks self.num_actions = num_actions # ---- Stem ---- self.stem = nn.Sequential( nn.Conv2d(18, channels, 3, padding=1, bias=False), nn.GroupNorm(norm_groups, channels), nn.GELU(), ) # ---- Residual tower ---- tower = [] for _ in range(blocks): tower.append(ResidualBlock(channels, norm_groups)) self.tower = nn.Sequential(*tower) # ---- Policy head (spatial): linear logits, no activation ---- # 320 = 64 destination squares × 5 promotion types self.policy_head = nn.Conv2d(channels, 320, 1, bias=True) # ---- Value head ---- self.value_head = nn.Sequential( nn.Conv2d(channels, 32, 1, bias=False), nn.GroupNorm(8, 32), nn.GELU(), nn.Flatten(), nn.Linear(32 * 8 * 8, 256), nn.GELU(), nn.Linear(256, 1), nn.Tanh(), ) self._init_weights() def _init_weights(self): """Initialize weights with scaled normal for stability.""" for m in self.modules(): if isinstance(m, nn.Conv2d): nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu') elif isinstance(m, nn.Linear): nn.init.trunc_normal_(m.weight, std=0.02) if m.bias is not None: nn.init.zeros_(m.bias) def forward(self, boards: torch.Tensor): """ Args: boards: [B, 18, 8, 8] float tensor (values 0.0 or 1.0) Returns: policy_logits: [B, 20480] raw logits for all action IDs value: [B] tanh-squashed scalar [-1, 1] """ x = self.stem(boards) x = self.tower(x) # Policy head: [B, 320, 8, 8] (NCHW) -> [B, 20480] # Action ID encoding: # action_id = ((from_sq * 64) + to_sq) * 5 + promo_id # Spatial (h,w) = from_square = h*8 + w # Channel c = to_sq * 5 + promo_id (0-319) pol = self.policy_head(x) # [B, 320, 8, 8] NCHW B = pol.shape[0] pol = pol.permute(0, 2, 3, 1) # [B, 8, 8, 320] NHWC policy_logits = pol.reshape(B, 8 * 8 * 320) # [B, 20480] # After permute+reshape: # flat_idx = (h*8+w) * 320 + c # = from_sq * 320 + to_sq * 5 + promo_id # = ((from_sq * 64) + to_sq) * 5 + promo_id ✓ # Value head value = self.value_head(x).squeeze(-1) # [B] return policy_logits, value def count_parameters(model: nn.Module) -> int: """Return total number of trainable parameters.""" return sum(p.numel() for p in model.parameters() if p.requires_grad) def get_model_config(model: ChessResNet) -> dict: """Return model hyperparameters for checkpoint saving.""" return dict( channels=model.channels, blocks=model.blocks, num_actions=model.num_actions, ) def create_model_from_config(config: dict) -> ChessResNet: """Create a model from a config dict (as stored in checkpoints).""" return ChessResNet( channels=config.get("channels", 256), blocks=config.get("blocks", 24), num_actions=config.get("num_actions", 20480), ) if __name__ == "__main__": m = ChessResNet(channels=256, blocks=24) n_params = count_parameters(m) print(f"ChessResNet(channels=256, blocks=24): {n_params:,} parameters") # ~29.0M expected m = ChessResNet(channels=128, blocks=4) n_params = count_parameters(m) print(f"ChessResNet(channels=128, blocks=4): {n_params:,} parameters") # Test forward x = torch.randn(4, 18, 8, 8) pol, val = m(x) print(f"Policy logits shape: {pol.shape} (expected [4, 20480])") print(f"Value shape: {val.shape} (expected [4])")