import torch import torch.nn as nn from transformers import PreTrainedModel, PretrainedConfig from transformers.modeling_outputs import CausalLMOutputWithPast class ChessConfig(PretrainedConfig): model_type = "chess_lm" def __init__( self, vocab_size=354, n_positions=256, n_embd=96, n_layer=8, n_head=8, tie_word_embeddings=False, **kwargs ): super().__init__(**kwargs) self.vocab_size = vocab_size self.n_positions = n_positions self.n_embd = n_embd self.n_layer = n_layer self.n_head = n_head self.tie_word_embeddings = tie_word_embeddings class ChessForCausalLM(PreTrainedModel): config_class = ChessConfig def __init__(self, config): super().__init__(config) self.token_emb = nn.Embedding(config.vocab_size, config.n_embd) self.pos_emb = nn.Embedding(config.n_positions, config.n_embd) encoder_layer = nn.TransformerEncoderLayer( d_model=config.n_embd, nhead=config.n_head, dim_feedforward=4 * config.n_embd, batch_first=True, norm_first=True ) self.transformer = nn.TransformerEncoder(encoder_layer, 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_word_embeddings: self.lm_head.weight = self.token_emb.weight self.post_init() def forward(self, input_ids, labels=None): B, T = input_ids.shape pos = torch.arange(T, device=input_ids.device) x = self.token_emb(input_ids) + self.pos_emb(pos) mask = torch.triu(torch.ones(T, T, device=x.device), diagonal=1).bool() x = self.transformer(x, mask=mask) x = self.ln_f(x) logits = self.lm_head(x) loss = None if labels is not None: loss_fct = nn.CrossEntropyLoss() loss = loss_fct( logits[:, :-1].reshape(-1, logits.size(-1)), labels[:, 1:].reshape(-1) ) return CausalLMOutputWithPast(loss=loss, logits=logits)