ChessResNet-30M / model.py
Joeyfully's picture
Upload 5 files
ba69de3 verified
Raw
History Blame Contribute Delete
6.02 kB
"""
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])")