| """ | |
| Required for AutoModelForCausalLM(trust_remote_code=True) to know how to build your custom PyTorch architecture. | |
| (Note: Big companies like Meta/Mistral don't upload files like this because they merge their architecture code directly into the official `transformers` GitHub repository.) | |
| """ | |
| from __future__ import annotations | |
| import torch | |
| import torch.nn as nn | |
| from transformers import PretrainedConfig, PreTrainedModel | |
| from transformers.generation import GenerationMixin | |
| from transformers.modeling_outputs import CausalLMOutput | |
| class GPTCustomConfig(PretrainedConfig): | |
| model_type = "gpt-custom" | |
| attribute_map = { | |
| "num_hidden_layers": "number_of_transformer_block", | |
| "hidden_size": "d_model", | |
| "num_attention_heads": "num_heads", | |
| } | |
| def __init__( | |
| self, | |
| vocab_size: int = 32000, | |
| d_model: int = 768, | |
| num_heads: int = 8, | |
| number_of_transformer_block: int = 6, | |
| max_seq_len: int = 1024, | |
| dropout: float = 0.2, | |
| **kwargs, | |
| ) -> None: | |
| super().__init__(**kwargs) | |
| self.vocab_size = vocab_size | |
| self.d_model = d_model | |
| self.num_heads = num_heads | |
| self.number_of_transformer_block = number_of_transformer_block | |
| self.max_seq_len = max_seq_len | |
| self.dropout = dropout | |
| class _GPTBlock(nn.Module): | |
| def __init__(self, config: GPTCustomConfig) -> None: | |
| super().__init__() | |
| self.layer_norm_1 = nn.LayerNorm(config.d_model) | |
| self.layer_norm_2 = nn.LayerNorm(config.d_model) | |
| self.multihead_attention = nn.MultiheadAttention( | |
| embed_dim=config.d_model, | |
| num_heads=config.num_heads, | |
| batch_first=True, | |
| ) | |
| self.gelu = nn.GELU() | |
| self.ffn_1 = nn.Linear(config.d_model, config.d_model * 4) | |
| self.ffn_2 = nn.Linear(config.d_model * 4, config.d_model) | |
| self.mha_drop = nn.Dropout(config.dropout) | |
| self.ffn_drop = nn.Dropout(config.dropout) | |
| self.register_buffer( | |
| "causal_mask", | |
| torch.triu( | |
| torch.full((config.max_seq_len, config.max_seq_len), float("-inf")), | |
| diagonal=1, | |
| ), | |
| ) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| _, seq_len, _ = x.size() | |
| ln1 = self.layer_norm_1(x) | |
| attn_out, _ = self.multihead_attention( | |
| ln1, ln1, ln1, | |
| attn_mask=self.causal_mask[:seq_len, :seq_len], | |
| ) | |
| x = x + self.mha_drop(attn_out) | |
| ln2 = self.layer_norm_2(x) | |
| ff_out = self.ffn_2(self.gelu(self.ffn_1(ln2))) | |
| return x + self.ffn_drop(ff_out) | |
| class GPTCustomForCausalLM(PreTrainedModel, GenerationMixin): | |
| config_class = GPTCustomConfig | |
| def __init__(self, config: GPTCustomConfig) -> None: | |
| super().__init__(config) | |
| self.token_embedding = nn.Embedding(config.vocab_size, config.d_model) | |
| self.positional_encoding = nn.Embedding(config.max_seq_len, config.d_model) | |
| self.emb_dropout = nn.Dropout(config.dropout) | |
| self.transformer_blocks = nn.ModuleList( | |
| [_GPTBlock(config) for _ in range(config.number_of_transformer_block)] | |
| ) | |
| self.layer_norm_final = nn.LayerNorm(config.d_model) | |
| self.final_linear_layer = nn.Linear(config.d_model, config.vocab_size, bias=False) | |
| self.final_linear_layer.weight = self.token_embedding.weight | |
| self.config.is_decoder = True | |
| self.post_init() | |
| def forward( | |
| self, | |
| input_ids: torch.Tensor, | |
| attention_mask: torch.Tensor | None = None, | |
| labels: torch.Tensor | None = None, | |
| **kwargs, | |
| ) -> CausalLMOutput: | |
| batch_size, seq_len = input_ids.shape | |
| position_ids = ( | |
| torch.arange(seq_len, device=input_ids.device) | |
| .unsqueeze(0) | |
| .expand(batch_size, -1) | |
| ) | |
| x = self.token_embedding(input_ids) + self.positional_encoding(position_ids) | |
| x = self.emb_dropout(x) | |
| for block in self.transformer_blocks: | |
| x = block(x) | |
| logits = self.final_linear_layer(self.layer_norm_final(x)) | |
| loss = None | |
| if labels is not None: | |
| shift_logits = logits[..., :-1, :].contiguous() | |
| shift_labels = labels[..., 1:].contiguous() | |
| loss = nn.functional.cross_entropy( | |
| shift_logits.view(-1, self.config.vocab_size), | |
| shift_labels.view(-1), | |
| ) | |
| return CausalLMOutput(loss=loss, logits=logits) | |
| def get_input_embeddings(self) -> nn.Embedding: | |
| return self.token_embedding | |
| def set_input_embeddings(self, value: nn.Embedding) -> None: | |
| self.token_embedding = value | |
| def prepare_inputs_for_generation( | |
| self, | |
| input_ids: torch.Tensor, | |
| **kwargs, | |
| ) -> dict: | |
| return {"input_ids": input_ids} | |
| def tie_weights(self, **kwargs) -> None: | |
| self.final_linear_layer.weight = self.token_embedding.weight | |