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}