Charles-the-Chess-Bot / modeling_chess_policy.py
Datdanboi25's picture
Upload folder using huggingface_hub
d2d7586 verified
Raw
History Blame Contribute Delete
13.1 kB
"""Self-contained Hugging Face implementation of the chess policy model."""
from __future__ import annotations
import math
from dataclasses import dataclass
import torch
import torch.nn.functional as F
from torch import nn
from torch.nn.attention import SDPBackend, sdpa_kernel
from torch.utils.checkpoint import checkpoint
from transformers import PreTrainedModel
from transformers.utils import ModelOutput
try:
from .configuration_chess_policy import ChessPolicyConfig
except ImportError: # Allows convert_checkpoint.py to run from this folder.
from configuration_chess_policy import ChessPolicyConfig
class SwiGLU(nn.Module):
def __init__(self, d_model: int, hidden: int, dropout: float) -> None:
super().__init__()
self.gate = nn.Linear(d_model, hidden, bias=False)
self.up = nn.Linear(d_model, hidden, bias=False)
self.down = nn.Linear(hidden, d_model, bias=False)
self.dropout = nn.Dropout(dropout)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.down(self.dropout(F.silu(self.gate(x)) * self.up(x)))
class SDPASelfAttention(nn.Module):
"""Bias-free bidirectional attention using PyTorch fused SDPA."""
def __init__(self, d_model: int, n_heads: int, dropout: float) -> None:
super().__init__()
self.d_model = d_model
self.n_heads = n_heads
self.head_dim = d_model // n_heads
self.dropout = dropout
self.in_proj_weight = nn.Parameter(torch.empty(3 * d_model, d_model))
self.out_proj = nn.Linear(d_model, d_model, bias=False)
def forward(
self,
x: torch.Tensor,
key_padding_mask: torch.Tensor | None = None,
) -> torch.Tensor:
batch, tokens, _ = x.shape
qkv = F.linear(x, self.in_proj_weight)
qkv = qkv.view(batch, tokens, 3, self.n_heads, self.head_dim)
query, key, value = qkv.unbind(dim=2)
query = query.transpose(1, 2)
key = key.transpose(1, 2)
value = value.transpose(1, 2)
attention_mask = None
if key_padding_mask is not None:
attention_mask = (~key_padding_mask)[:, None, None, :]
sdpa_args = {
"attn_mask": attention_mask,
"dropout_p": self.dropout if self.training else 0.0,
"is_causal": False,
}
if query.is_cuda:
with sdpa_kernel(SDPBackend.EFFICIENT_ATTENTION):
attended = F.scaled_dot_product_attention(
query, key, value, **sdpa_args
)
else:
attended = F.scaled_dot_product_attention(
query, key, value, **sdpa_args
)
attended = attended.transpose(1, 2).contiguous()
return self.out_proj(attended.view(batch, tokens, self.d_model))
class BidirectionalTransformerBlock(nn.Module):
def __init__(
self,
d_model: int,
n_heads: int,
swiglu_hidden: int,
dropout: float,
stack_depth: int,
) -> None:
super().__init__()
self.attention_norm = nn.LayerNorm(d_model)
self.attention = SDPASelfAttention(d_model, n_heads, dropout)
self.attention_dropout = nn.Dropout(dropout)
self.ffn_norm = nn.LayerNorm(d_model)
self.ffn = SwiGLU(d_model, swiglu_hidden, dropout)
self.ffn_dropout = nn.Dropout(dropout)
self.reset_parameters(stack_depth)
def reset_parameters(self, stack_depth: int) -> None:
residual_std = 0.02 / math.sqrt(2.0 * stack_depth)
nn.init.normal_(self.attention.in_proj_weight, mean=0.0, std=0.02)
nn.init.normal_(self.attention.out_proj.weight, mean=0.0, std=residual_std)
nn.init.normal_(self.ffn.gate.weight, mean=0.0, std=0.02)
nn.init.normal_(self.ffn.up.weight, mean=0.0, std=0.02)
nn.init.normal_(self.ffn.down.weight, mean=0.0, std=residual_std)
def forward(
self,
x: torch.Tensor,
key_padding_mask: torch.Tensor | None = None,
) -> torch.Tensor:
attended = self.attention(self.attention_norm(x), key_padding_mask)
x = x + self.attention_dropout(attended)
x = x + self.ffn_dropout(self.ffn(self.ffn_norm(x)))
return x
class BidirectionalTransformerStack(nn.Module):
def __init__(self, *, layers: int, config: ChessPolicyConfig) -> None:
super().__init__()
self.activation_checkpointing = config.activation_checkpointing
self.layers = nn.ModuleList(
[
BidirectionalTransformerBlock(
config.d_model,
config.n_heads,
config.swiglu_hidden,
config.dropout,
stack_depth=layers,
)
for _ in range(layers)
]
)
self.final_norm = nn.LayerNorm(config.d_model)
def forward(
self,
x: torch.Tensor,
key_padding_mask: torch.Tensor | None = None,
) -> torch.Tensor:
for layer in self.layers:
if self.activation_checkpointing and self.training:
x = checkpoint(layer, x, key_padding_mask, use_reentrant=False)
else:
x = layer(x, key_padding_mask)
return self.final_norm(x)
class BoardEncoder(nn.Module):
def __init__(self, config: ChessPolicyConfig) -> None:
super().__init__()
self.piece_embedding = nn.Embedding(config.piece_states, config.d_model)
self.square_embedding = nn.Embedding(config.board_squares, config.d_model)
self.summary_token = nn.Parameter(torch.empty(1, 1, config.d_model))
self.transformer = BidirectionalTransformerStack(
layers=config.board_layers, config=config
)
nn.init.normal_(self.piece_embedding.weight, mean=0.0, std=0.02)
nn.init.normal_(self.square_embedding.weight, mean=0.0, std=0.02)
nn.init.normal_(self.summary_token, mean=0.0, std=0.02)
def forward(self, boards: torch.Tensor) -> torch.Tensor:
if boards.ndim != 2 or boards.shape[1] != 64:
raise ValueError(f"Expected boards shaped [N, 64], got {boards.shape}")
boards = boards.to(torch.int64)
square_ids = torch.arange(64, device=boards.device)
squares = self.piece_embedding(boards) + self.square_embedding(square_ids)
summary = self.summary_token.expand(boards.shape[0], -1, -1)
encoded = self.transformer(torch.cat((summary, squares), dim=1))
return encoded[:, 0]
@dataclass
class ChessPolicyOutput(ModelOutput):
"""Hugging Face output with per-candidate move scores."""
loss: torch.Tensor | None = None
logits: torch.Tensor | None = None
candidate_mask: torch.Tensor | None = None
distill_loss: torch.Tensor | None = None
class ChessTransitionPolicy(PreTrainedModel):
"""Score legal chess moves as contextualized latent transitions."""
config_class = ChessPolicyConfig
base_model_prefix = "chess_policy"
main_input_name = "current_boards"
def _init_weights(self, module: nn.Module) -> None:
"""Keep the project's explicit initialization scheme unchanged."""
# The submodules are initialized explicitly below. This no-op lets
# Hugging Face's post_init register loader metadata without replacing
# the original initialization scheme.
del module
def __init__(self, config: ChessPolicyConfig) -> None:
super().__init__(config)
self.board_encoder = BoardEncoder(config)
self.candidate_type_embedding = nn.Embedding(2, config.d_model)
self.candidate_transformer = BidirectionalTransformerStack(
layers=config.candidate_layers, config=config
)
self.query = nn.Linear(config.d_model, config.d_model, bias=False)
self.key = nn.Linear(config.d_model, config.d_model, bias=False)
self.query_norm = nn.RMSNorm(config.d_model, elementwise_affine=False)
self.key_norm = nn.RMSNorm(config.d_model, elementwise_affine=False)
nn.init.normal_(self.candidate_type_embedding.weight, mean=0.0, std=0.02)
nn.init.normal_(self.query.weight, mean=0.0, std=0.02)
nn.init.normal_(self.key.weight, mean=0.0, std=0.02)
self.post_init()
def forward(
self,
*,
current_boards: torch.Tensor,
successor_boards: torch.Tensor,
candidate_owner: torch.Tensor,
candidate_offsets: torch.Tensor,
candidate_mask: torch.Tensor,
target_indices: torch.Tensor | None = None,
teacher_logits: torch.Tensor | None = None,
distill_temperature: float = 1.0,
teacher_temperature: float = 120.0,
return_dict: bool | None = None,
) -> ChessPolicyOutput | tuple[torch.Tensor, ...]:
batch_size = current_boards.shape[0]
all_boards = torch.cat((current_boards, successor_boards), dim=0)
all_states = self.board_encoder(all_boards)
current_states = all_states[:batch_size]
successor_states = all_states[batch_size:]
candidate_owner = candidate_owner.to(torch.int64)
transitions = successor_states - current_states[candidate_owner]
local_indices = (
torch.arange(transitions.shape[0], device=transitions.device)
- candidate_offsets[candidate_owner]
)
max_candidates = candidate_mask.shape[1]
candidate_sequence = transitions.new_zeros(
batch_size, max_candidates + 1, self.config.d_model
)
candidate_sequence[:, 0] = current_states
candidate_sequence[candidate_owner, local_indices + 1] = transitions
type_ids = torch.ones(
batch_size,
max_candidates + 1,
dtype=torch.int64,
device=transitions.device,
)
type_ids[:, 0] = 0
candidate_sequence = candidate_sequence + self.candidate_type_embedding(type_ids)
valid_mask = torch.cat(
(
torch.ones(
batch_size,
1,
dtype=torch.bool,
device=candidate_mask.device,
),
candidate_mask,
),
dim=1,
)
contextualized = self.candidate_transformer(
candidate_sequence, key_padding_mask=~valid_mask
)
move_states = contextualized[:, 1:]
query = self.query_norm(self.query(current_states))
keys = self.key_norm(self.key(move_states))
logits = torch.einsum("bd,bnd->bn", query, keys)
logits = logits / math.sqrt(self.config.d_model)
logits = logits.masked_fill(~candidate_mask, float("-inf"))
loss = None
distill_loss = None
if teacher_logits is not None:
if distill_temperature <= 0 or teacher_temperature <= 0:
raise ValueError("distillation temperatures must be positive")
teacher_logits = teacher_logits.to(logits.dtype).masked_fill(
~candidate_mask, float("-inf")
)
teacher_log_probs = F.log_softmax(
teacher_logits / teacher_temperature, dim=1
)
teacher_probs = teacher_log_probs.exp()
student_log_probs = F.log_softmax(logits / distill_temperature, dim=1)
safe_student_log_probs = torch.where(
candidate_mask, student_log_probs, torch.zeros_like(student_log_probs)
)
safe_teacher_log_probs = torch.where(
candidate_mask, teacher_log_probs, torch.zeros_like(teacher_log_probs)
)
distill_loss = (
teacher_probs * (safe_teacher_log_probs - safe_student_log_probs)
).sum(dim=1).mean() * (distill_temperature**2)
if target_indices is not None:
loss = F.cross_entropy(logits, target_indices)
elif distill_loss is not None:
loss = distill_loss
if return_dict is False:
return tuple(
value
for value in (loss, logits, candidate_mask, distill_loss)
if value is not None
)
return ChessPolicyOutput(
loss=loss,
logits=logits,
candidate_mask=candidate_mask,
distill_loss=distill_loss,
)
def parameter_breakdown(self) -> dict[str, int]:
board = sum(p.numel() for p in self.board_encoder.parameters())
candidates = sum(p.numel() for p in self.candidate_transformer.parameters())
candidate_types = sum(p.numel() for p in self.candidate_type_embedding.parameters())
scorer = sum(p.numel() for p in self.query.parameters()) + sum(
p.numel() for p in self.key.parameters()
)
return {
"board_encoder": board,
"candidate_transformer": candidates,
"candidate_type_embedding": candidate_types,
"scorer": scorer,
"total": sum(p.numel() for p in self.parameters()),
}