MCPcopy Create free account
hub / github.com/AtlasAnalyticsLab/AdaFisher / MiniGPT1

Class MiniGPT1

Language_Model/GPT1.py:270–425  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

268
269
270class 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)`)

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected