File size: 918 Bytes
d7d79b9 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 | 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}
|