| 98 | # A toy transformer model, partly inspired by the nanoGPT model: |
| 99 | # https://github.com/karpathy/nanoGPT. |
| 100 | class Transformer(nn.Module): |
| 101 | def __init__(self, args: ModelArgs): |
| 102 | super().__init__() |
| 103 | assert args.vocab_size is not None |
| 104 | assert args.max_seq_len is not None |
| 105 | self.model_args = args |
| 106 | self.max_seq_len = args.max_seq_len |
| 107 | self.tok_embeddings = nn.Embedding(args.vocab_size, args.dim) |
| 108 | self.pos_embeddings = nn.Embedding(args.max_seq_len, args.dim) |
| 109 | self.dropout = nn.Dropout(args.dropout_p) |
| 110 | self.layers = nn.ModuleList() |
| 111 | for _ in range(args.n_layers): |
| 112 | self.layers.append(TransformerBlock(args)) |
| 113 | self.norm = nn.LayerNorm(args.dim) |
| 114 | self.output = nn.Linear(args.dim, args.vocab_size, bias=False) |
| 115 | |
| 116 | def forward(self, tokens): |
| 117 | _bsz, seq_len = tokens.size() |
| 118 | assert seq_len <= self.max_seq_len |
| 119 | h = self.tok_embeddings(tokens) |
| 120 | pos = torch.arange(0, seq_len, device=tokens.device) |
| 121 | p = self.pos_embeddings(pos) # positional embeddings of shape (seq_len, dim) |
| 122 | h = h + p |
| 123 | h = self.dropout(h) |
| 124 | for layer in self.layers: |
| 125 | h = layer(h) |
| 126 | h = self.norm(h) |
| 127 | output = self.output(h).float() |
| 128 | return output |
| 129 | |
| 130 | def reset_parameters(self): |
| 131 | self.tok_embeddings.reset_parameters() |
| 132 | self.pos_embeddings.reset_parameters() |
| 133 | self.norm.reset_parameters() |
| 134 | self.output.reset_parameters() |