| 268 | |
| 269 | |
| 270 | class MiniGPT1(nn.Module): |
| 271 | def __init__( |
| 272 | self, |
| 273 | vocabulary_size=40479, |
| 274 | embedding_size=768, |
| 275 | sequence_length=256, |
| 276 | num_heads=12, |
| 277 | num_layers=4, |
| 278 | learn_embeddings=False, |
| 279 | _tokens_embedding_weight=None, |
| 280 | _positional_embedding_weight=None, |
| 281 | ): |
| 282 | super(MiniGPT1, self).__init__() |
| 283 | self.vocabulary_size = vocabulary_size |
| 284 | self.embedding_size = embedding_size |
| 285 | self.sequence_length = sequence_length |
| 286 | self.num_heads = num_heads |
| 287 | self.num_layers = num_layers |
| 288 | self.learn_embeddings = learn_embeddings |
| 289 | self.head_size = embedding_size // num_heads |
| 290 | |
| 291 | self.embedding = GPT1Embedding( |
| 292 | vocabulary_size, |
| 293 | embedding_size, |
| 294 | sequence_length, |
| 295 | _tokens_embedding_weight=_tokens_embedding_weight, |
| 296 | _positional_embedding_weight=_positional_embedding_weight, |
| 297 | ) |
| 298 | self.layers = nn.ModuleList( |
| 299 | [ |
| 300 | Block(self.head_size, 4 * embedding_size, num_heads, sequence_length) |
| 301 | for _ in range(num_layers) |
| 302 | ] |
| 303 | ) |
| 304 | self.classifier = nn.Linear(embedding_size, vocabulary_size, bias=False) |
| 305 | |
| 306 | # Tying classifier and embedding weights |
| 307 | self.classifier.weight = self.embedding.tokens.weight |
| 308 | |
| 309 | # Freeze the embedding weights, depending on learn_embeddings |
| 310 | self.embedding.requires_grad_(learn_embeddings) |
| 311 | |
| 312 | def get_embeddings(self, inputs): |
| 313 | """Get the embeddings for some input sequence. |
| 314 | |
| 315 | This function computes the embedding vectors based on the input |
| 316 | sequence (and the positions of the tokens). See also the module |
| 317 | `GPT1Embedding` for details about the implementation of |
| 318 | `self.embedding`. |
| 319 | |
| 320 | Parameters |
| 321 | ---------- |
| 322 | inputs (`torch.LongTensor` of shape `(batch_size, sequence_length)`) |
| 323 | The input tensor containing the token sequences. |
| 324 | |
| 325 | Returns |
| 326 | ------- |
| 327 | embeddings (`torch.FloatTensor` of shape `(batch_size, sequence_length, embedding_size)`) |
nothing calls this directly
no outgoing calls
no test coverage detected