| 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} | |