import torch import torch.nn as nn from transformers import PreTrainedModel, PretrainedConfig, GenerationMixin from transformers.modeling_outputs import CausalLMOutput import torch import torch.nn.functional as F import torch.nn as nn import math embedding = 256 heads = 4 layers = 4 dropout = 0.1 msl = 160 class PositionalEncoding(nn.Module): def __init__(self, d_model, max_len=5000): super(PositionalEncoding, self).__init__() pe = torch.zeros(max_len, d_model) position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1) div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model)) pe[:, 0::2] = torch.sin(position * div_term) pe[:, 1::2] = torch.cos(position * div_term) pe = pe.unsqueeze(0).transpose(0, 1) self.register_buffer('pe', pe) def forward(self, x): return x + self.pe[:x.size(1), :].transpose(0, 1) class CausalSelfAttention(nn.Module): def __init__(self, d_model, nhead, dropout=0.1): super().__init__() assert d_model % nhead == 0 self.nhead = nhead self.head_dim = d_model // nhead self.dropout = dropout self.qkv = nn.Linear(d_model, d_model * 3) self.out_proj = nn.Linear(d_model, d_model) def forward(self, x): B, T, C = x.shape # Create Q, K, V q, k, v = self.qkv(x).chunk(3, dim=-1) # [B, T, C] -> [B, heads, T, head_dim] q = q.view(B, T, self.nhead, self.head_dim).transpose(1, 2) k = k.view(B, T, self.nhead, self.head_dim).transpose(1, 2) v = v.view(B, T, self.nhead, self.head_dim).transpose(1, 2) y = F.scaled_dot_product_attention( q, k, v, attn_mask=None, dropout_p=self.dropout if self.training else 0.0, is_causal=True ) y = y.transpose(1, 2).contiguous().view(B, T, C) return self.out_proj(y) class TransformerBlock(nn.Module): def __init__(self, d_model, nhead, dropout=0.1): super().__init__() self.norm1 = nn.LayerNorm(d_model) self.attention = CausalSelfAttention( d_model, nhead, dropout ) self.norm2 = nn.LayerNorm(d_model) self.ffn = nn.Sequential( nn.Linear(d_model, d_model * 4), nn.GELU(), nn.Linear(d_model * 4, d_model), nn.Dropout(dropout) ) def forward(self, x): x = x + self.attention(self.norm1(x)) x = x + self.ffn(self.norm2(x)) return x class TransformerLanguageModel(nn.Module): def __init__( self, vocab_size, d_model=512, nhead=8, num_layers=8, dropout=0.1, max_seq_len=160 ): super().__init__() self.d_model = d_model self.max_seq_len = max_seq_len self.token_embedding = nn.Embedding( vocab_size, d_model ) self.positional_encoding = PositionalEncoding( d_model, max_seq_len ) self.transformer = nn.ModuleList([ TransformerBlock( d_model, nhead, dropout ) for _ in range(num_layers) ]) self.final_norm = nn.LayerNorm(d_model) self.output_layer = nn.Linear( d_model, vocab_size, bias=False ) def forward(self, src): x = self.token_embedding(src) x = self.positional_encoding(x) for layer in self.transformer: x = layer(x) x = self.final_norm(x) return self.output_layer(x) class LightningConfig(PretrainedConfig): model_type = "lightning" def __init__( self, vocab_size=50000, d_model=256, nhead=4, num_layers=4, dropout=0.1, max_seq_len=160, **kwargs ): super().__init__( tie_word_embeddings=False, **kwargs ) self.vocab_size = vocab_size self.d_model = d_model self.nhead = nhead self.num_layers = num_layers self.dropout = dropout self.num_hidden_layers = num_layers self.num_attention_heads = nhead self.hidden_size = d_model self.max_seq_len = max_seq_len class LightningForCausalLM(PreTrainedModel, GenerationMixin): config_class = LightningConfig base_model_prefix = "lightning" def __init__(self, config): super().__init__(config) self.lightning = TransformerLanguageModel( vocab_size=config.vocab_size, d_model=config.d_model, nhead=config.nhead, num_layers=config.num_layers, dropout=config.dropout, max_seq_len=config.max_seq_len ) self.post_init() def forward(self, input_ids=None, labels=None, **kwargs): logits = self.lightning(input_ids) loss = None if labels is not None: shift_logits = logits[..., :-1, :].contiguous() shift_labels = labels[..., 1:].contiguous() loss_fn = nn.CrossEntropyLoss() loss = loss_fn( shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1) ) return CausalLMOutput( loss=loss, logits=logits ) def get_input_embeddings(self): return self.lightning.token_embedding def set_input_embeddings(self, value): self.lightning.token_embedding = value def get_output_embeddings(self): return self.lightning.output_layer def set_output_embeddings(self, new_embeddings): self.lightning.output_layer = new_embeddings