|
|
| import torch |
| import torch.nn as nn |
| from transformers import PretrainedConfig, PreTrainedModel |
|
|
| class SupernovaEncoderConfig(PretrainedConfig): |
| model_type = 'supernova_encoder' |
| def __init__( |
| self, |
| vocab_size=50257, |
| hidden_size=512, |
| num_hidden_layers=6, |
| num_attention_heads=8, |
| intermediate_size=2048, |
| max_position_embeddings=300, |
| output_dim=2304, |
| layer_norm_eps=1e-12, |
| **kwargs |
| ): |
| super().__init__(**kwargs) |
| self.vocab_size = vocab_size |
| self.hidden_size = hidden_size |
| self.num_hidden_layers = num_hidden_layers |
| self.num_attention_heads = num_attention_heads |
| self.intermediate_size = intermediate_size |
| self.max_position_embeddings = max_position_embeddings |
| self.output_dim = output_dim |
| self.layer_norm_eps = layer_norm_eps |
|
|
| class SupernovaNepaliEncoder(PreTrainedModel): |
| config_class = SupernovaEncoderConfig |
|
|
| def __init__(self, config): |
| super().__init__(config) |
| self.embeddings = nn.Embedding(config.vocab_size, config.hidden_size) |
| self.position_embeddings = nn.Embedding(config.max_position_embeddings, config.hidden_size) |
| |
| layer = nn.TransformerEncoderLayer( |
| d_model=config.hidden_size, |
| nhead=config.num_attention_heads, |
| dim_feedforward=config.intermediate_size, |
| batch_first=True, |
| norm_first=True |
| ) |
| self.encoder = nn.TransformerEncoder(layer, num_layers=config.num_hidden_layers) |
| self.projection = nn.Linear(config.hidden_size, config.output_dim) |
| self.ln_final = nn.LayerNorm(config.output_dim, eps=config.layer_norm_eps) |
| self.post_init() |
|
|
| def forward(self, input_ids, attention_mask=None): |
| seq_length = input_ids.size(1) |
| position_ids = torch.arange(seq_length, dtype=torch.long, device=input_ids.device).unsqueeze(0) |
| |
| x = self.embeddings(input_ids) + self.position_embeddings(position_ids) |
| |
| padding_mask = None |
| if attention_mask is not None: |
| padding_mask = ~(attention_mask.bool()) |
| |
| hidden_states = self.encoder(x, src_key_padding_mask=padding_mask) |
| projected = self.projection(hidden_states) |
| return projected |
|
|