MCPcopy Create free account
hub / github.com/OSU-NLP-Group/Loop-Think-Generalize / RecurrentGPT2Block

Class RecurrentGPT2Block

gpt_utils_extrapolation.py:322–602  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

320
321
322class RecurrentGPT2Block(nn.Module):
323 def __init__(self, config, num_iterations, positional_embedding_type='learned', input_injection=False,
324 c_scale=0.003):
325 super().__init__()
326 self.config = config
327 self.num_iterations = num_iterations
328 self.input_injection = input_injection
329 self.blocks = nn.ModuleList([GPT2Block(config) for _ in range(config.n_layer)])
330
331 self.token_embedding = nn.Embedding(config.vocab_size, config.n_embd)
332 self.positional_embedding_type = positional_embedding_type
333
334 if self.positional_embedding_type == 'learned':
335 self.position_embedding = nn.Embedding(config.n_positions, config.n_embd)
336 elif self.positional_embedding_type == 'sinusoidal':
337 self.position_embedding = None
338 pe = torch.zeros(config.n_positions, config.n_embd)
339 position = torch.arange(0, config.n_positions, dtype=torch.float).unsqueeze(1)
340 div_term = torch.exp(torch.arange(0, config.n_embd, 2).float() * (-math.log(10000.0) / config.n_embd))
341 pe[:, 0::2] = torch.sin(position * div_term)
342 pe[:, 1::2] = torch.cos(position * div_term)
343 self.register_buffer('pe', pe)
344 elif self.positional_embedding_type == 'none':
345 self.position_embedding = None
346 else:
347 raise ValueError(f"Unsupported positional embedding type: {self.positional_embedding_type}")
348
349 embd_pdrop = getattr(config, "embd_pdrop", 0.1)
350 self.dropout = nn.Dropout(embd_pdrop)
351
352 layer_norm_epsilon = getattr(config, "layer_norm_epsilon", 1e-5)
353 self.ln_f = nn.LayerNorm(config.n_embd, eps=layer_norm_epsilon)
354
355 self.lm_head = nn.Linear(config.n_embd, config.vocab_size, bias=False)
356 self.lm_head.weight = self.token_embedding.weight
357 scale = c_scale
358 for block in self.blocks:
359 block.attn.c_proj.weight.data.mul_(scale)
360 block.mlp.c_proj.weight.data.mul_(scale)
361 if block.attn.c_proj.bias is not None:
362 block.attn.c_proj.bias.data.zero_()
363 if block.mlp.c_proj.bias is not None:
364 block.mlp.c_proj.bias.data.zero_()
365
366 def forward(self, input_ids, attention_mask=None):
367 batch_size, seq_len = input_ids.size()
368 device = input_ids.device
369
370 position_ids = torch.arange(0, seq_len, dtype=torch.long, device=device)
371 position_ids = position_ids.unsqueeze(0).expand(batch_size, seq_len)
372
373 token_embeds = self.token_embedding(input_ids)
374
375 if self.positional_embedding_type == 'learned':
376 position_ids = torch.arange(0, seq_len, dtype=torch.long, device=device).unsqueeze(0)
377 pos_embeds = self.position_embedding(position_ids)
378 hidden_states = token_embeds + pos_embeds
379 elif self.positional_embedding_type == 'sinusoidal':

Calls

no outgoing calls

Tested by

no test coverage detected