File size: 6,015 Bytes
ba69de3 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 | """
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])")
|