Download model.py from Joeyfully/ChessResNet-30M: direct link, hf CLI and curl.
- Browser
- Download file 6.02 kB
-
https://huggingface.co/Joeyfully/ChessResNet-30M/resolve/main/model.py
- Command line
-
hf download hf://Joeyfully/ChessResNet-30M/model.py
-
curl -L -o model.py https://huggingface.co/Joeyfully/ChessResNet-30M/resolve/main/model.py
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])") | |