| """
|
| 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
|
|
|
|
|
| self.stem = nn.Sequential(
|
| nn.Conv2d(18, channels, 3, padding=1, bias=False),
|
| nn.GroupNorm(norm_groups, channels),
|
| nn.GELU(),
|
| )
|
|
|
|
|
| tower = []
|
| for _ in range(blocks):
|
| tower.append(ResidualBlock(channels, norm_groups))
|
| self.tower = nn.Sequential(*tower)
|
|
|
|
|
|
|
| self.policy_head = nn.Conv2d(channels, 320, 1, bias=True)
|
|
|
|
|
| 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)
|
|
|
|
|
|
|
|
|
|
|
|
|
| pol = self.policy_head(x)
|
| B = pol.shape[0]
|
| pol = pol.permute(0, 2, 3, 1)
|
| policy_logits = pol.reshape(B, 8 * 8 * 320)
|
|
|
|
|
|
|
|
|
|
|
|
|
| value = self.value_head(x).squeeze(-1)
|
|
|
| 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")
|
|
|
|
|
| m = ChessResNet(channels=128, blocks=4)
|
| n_params = count_parameters(m)
|
| print(f"ChessResNet(channels=128, blocks=4): {n_params:,} parameters")
|
|
|
|
|
| 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])")
|
|
|