Yossri23's picture
Chess Challenge submission by Yossri23
a0e7273 verified
Raw
History Blame Contribute Delete
3.77 kB
from __future__ import annotations
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from transformers import PretrainedConfig, PreTrainedModel
from transformers.modeling_outputs import CausalLMOutputWithPast
class ChessConfig(PretrainedConfig):
model_type = "chess_transformer"
def __init__(self, vocab_size=75, n_embd=96, n_layer=8, n_head=4, n_ctx=512, n_inner=None, dropout=0.1, tie_weights=True, **kwargs):
super().__init__(**kwargs)
self.vocab_size = vocab_size
self.n_embd = n_embd # 96 : Largeur confortable
self.n_layer = n_layer # 8 couches
self.n_head = n_head
self.n_ctx = n_ctx # 512 : Contexte doublé pour les coordonnées
self.n_inner = n_inner if n_inner is not None else 3 * n_embd
self.dropout = dropout
self.tie_weights = tie_weights
self.tie_word_embeddings = tie_weights
class MultiHeadAttention(nn.Module):
def __init__(self, config):
super().__init__()
self.n_head = config.n_head
self.n_embd = config.n_embd
self.head_dim = config.n_embd // config.n_head
self.c_attn = nn.Linear(config.n_embd, 3 * config.n_embd)
self.c_proj = nn.Linear(config.n_embd, config.n_embd)
self.dropout = nn.Dropout(config.dropout)
self.register_buffer("bias", torch.tril(torch.ones(config.n_ctx, config.n_ctx)).view(1, 1, config.n_ctx, config.n_ctx), persistent=False)
def forward(self, x, mask=None):
B, T, C = x.size()
q, k, v = self.c_attn(x).split(self.n_embd, dim=2)
k = k.view(B, T, self.n_head, self.head_dim).transpose(1, 2)
q = q.view(B, T, self.n_head, self.head_dim).transpose(1, 2)
v = v.view(B, T, self.n_head, self.head_dim).transpose(1, 2)
att = (q @ k.transpose(-2, -1)) * (1.0 / math.sqrt(k.size(-1)))
att = att.masked_fill(self.bias[:,:,:T,:T] == 0, float('-inf'))
if mask is not None: att = att.masked_fill(mask.unsqueeze(1).unsqueeze(2) == 0, float('-inf'))
y = self.dropout(F.softmax(att, dim=-1)) @ v
return self.c_proj(y.transpose(1, 2).contiguous().view(B, T, C))
class TransformerBlock(nn.Module):
def __init__(self, config):
super().__init__()
self.ln1 = nn.LayerNorm(config.n_embd); self.attn = MultiHeadAttention(config)
self.ln2 = nn.LayerNorm(config.n_embd); self.mlp = nn.Sequential(nn.Linear(config.n_embd, config.n_inner), nn.GELU(), nn.Linear(config.n_inner, config.n_embd), nn.Dropout(config.dropout))
def forward(self, x, mask=None): return x + self.mlp(self.ln2(x + self.attn(self.ln1(x), mask)))
class ChessForCausalLM(PreTrainedModel):
config_class = ChessConfig
keys_to_ignore_on_load_missing = ["lm_head.weight"]
def __init__(self, config):
super().__init__(config)
self.wte = nn.Embedding(config.vocab_size, config.n_embd)
self.wpe = nn.Embedding(config.n_ctx, config.n_embd)
self.h = nn.ModuleList([TransformerBlock(config) for _ in range(config.n_layer)])
self.ln_f = nn.LayerNorm(config.n_embd)
self.lm_head = nn.Linear(config.n_embd, config.vocab_size, bias=False)
if config.tie_weights: self.lm_head.weight = self.wte.weight
self.post_init()
def forward(self, input_ids, attention_mask=None, labels=None, **kwargs):
x = self.wte(input_ids) + self.wpe(torch.arange(input_ids.size(1), device=input_ids.device))
for b in self.h: x = b(x, attention_mask)
logits = self.lm_head(self.ln_f(x))
loss = None
if labels is not None: loss = F.cross_entropy(logits.view(-1, logits.size(-1)), labels.view(-1), ignore_index=-100)
return CausalLMOutputWithPast(loss=loss, logits=logits)