import torch import torch.nn as nn from torch.nn import CrossEntropyLoss class MiniTransformer(nn.Module): def __init__(self, vocab_size, hidden_dim): super().__init__() self.embed = nn.Embedding(vocab_size, hidden_dim) self.transformer = nn.Transformer(hidden_dim, nhead=8, num_encoder_layers=6) self.out = nn.Linear(hidden_dim, vocab_size) def forward(self, input_ids, attention_mask=None, labels=None): x = self.embed(input_ids) x = self.transformer(x, x) logits = self.out(x) loss = None if labels is not None: shift_logits = logits[:, :-1, :].contiguous() shift_labels = labels[:, 1:].contiguous() loss_fct = CrossEntropyLoss() loss = loss_fct(shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1)) return {"loss": loss, "logits": logits}