| 320 | |
| 321 | |
| 322 | class 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': |
no outgoing calls
no test coverage detected