import math import re import torch import torch.nn as nn import torch.nn.functional as F from transformers import PreTrainedModel, PretrainedConfig from transformers.modeling_outputs import CausalLMOutput # ============================================================ # Sampling # ============================================================ def top_k_top_p_sample( logits, top_k=40, top_p=0.9 ): """ Sample one token from logits using top-k and/or top-p sampling. """ logits = logits.float() # -------------------------------------------------------- # Top-k # -------------------------------------------------------- if top_k is not None and top_k > 0: top_k = min( top_k, logits.size(-1) ) values, indices = torch.topk( logits, top_k ) filtered_logits = torch.full_like( logits, -float("inf") ) filtered_logits.scatter_( 0, indices, values ) logits = filtered_logits # -------------------------------------------------------- # Top-p # -------------------------------------------------------- if top_p is not None and 0.0 < top_p < 1.0: sorted_logits, sorted_indices = torch.sort( logits, descending=True ) probabilities = torch.softmax( sorted_logits, dim=-1 ) cumulative_probabilities = torch.cumsum( probabilities, dim=-1 ) remove_mask = ( cumulative_probabilities > top_p ) # Always keep the first token above the threshold. remove_mask[1:] = remove_mask[:-1].clone() remove_mask[0] = False sorted_logits[remove_mask] = -float("inf") logits = torch.full_like( logits, -float("inf") ) logits.scatter_( 0, sorted_indices, sorted_logits ) probabilities = torch.softmax( logits, dim=-1 ) next_token = torch.multinomial( probabilities, num_samples=1 ) return next_token.item() # ============================================================ # Positional Encoding # ============================================================ class PositionalEncoding(nn.Module): def __init__(self, d_model, max_len=5000): super().__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) # ============================================================ # Causal Self Attention # ============================================================ 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 q, k, v = self.qkv(x).chunk( 3, dim=-1 ) 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) # ============================================================ # Transformer Block # ============================================================ 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 # ============================================================ # Transformer Language Model # ============================================================ class TransformerLanguageModel(nn.Module): def __init__( self, vocab_size, d_model=256, nhead=4, num_layers=6, dropout=0.1, max_seq_len=200 ): 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) # ============================================================ # Lightning Config # ============================================================ class LightningConfig( PretrainedConfig ): model_type = "lightning" def __init__( self, vocab_size=75000, d_model=256, nhead=4, num_layers=6, dropout=0.1, max_seq_len=200, **kwargs ): kwargs.setdefault( "tie_word_embeddings", False ) super().__init__( **kwargs ) self.vocab_size = vocab_size self.d_model = d_model self.nhead = nhead self.num_layers = num_layers self.dropout = dropout self.max_seq_len = max_seq_len self.num_hidden_layers = ( num_layers ) self.num_attention_heads = ( nhead ) self.hidden_size = ( d_model ) # ============================================================ # Lightning Causal LM # ============================================================ class LightningForCausalLM( PreTrainedModel ): 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() # -------------------------------------------------------- # Forward # -------------------------------------------------------- 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 ) # -------------------------------------------------------- # Embeddings # -------------------------------------------------------- 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 ) # ============================================================ # Text Generation API # ============================================================ def generate_text( model, tokenizer, prompt, max_len, device, top_k=40, top_p=0.9, penalty=1.2, temperature=0.8, chat_history=None ): """ Generate text from Lightning. chat_history format: [ { "role": "user", "content": "Hello" }, { "role": "assistant", "content": "Hi!" } ] """ model.eval() # -------------------------------------------------------- # Maximum sequence length # -------------------------------------------------------- msl = getattr( model.config, "max_seq_len", 200 ) # -------------------------------------------------------- # Build conversation # -------------------------------------------------------- messages = [] if chat_history: for message in chat_history: role = message.get( "role", "" ).lower() content = message.get( "content", "" ).strip() if not content: continue if role == "user": messages.append( f"User: {content}" ) elif role == "assistant": messages.append( f"Assistant: {content}" ) messages.append( f"User: {prompt}" ) messages.append( "Assistant:" ) generation_prompt = "\n".join( messages ) # -------------------------------------------------------- # Tokenize # -------------------------------------------------------- encoding = tokenizer.encode( generation_prompt, add_special_tokens=False ) input_ids = torch.tensor( [encoding.ids], dtype=torch.long, device=device ) # -------------------------------------------------------- # Context window # -------------------------------------------------------- if input_ids.size(1) > msl: input_ids = input_ids[ :, -msl: ] prompt_len = input_ids.size(1) generated = ( input_ids[0].tolist() ) # -------------------------------------------------------- # Special tokens # -------------------------------------------------------- eos_id = tokenizer.token_to_id( "<|endoftext|>" ) eor_id = tokenizer.token_to_id( "<|eor|>" ) pad_id = getattr( tokenizer, "pad_id", None ) # -------------------------------------------------------- # Generation # -------------------------------------------------------- for _ in range(max_len): src = input_ids[ :, -msl: ] with torch.no_grad(): output = model( src ) logits = ( output.logits[:, -1, :] .squeeze(0) ) # ---------------------------------------------------- # Prevent PAD generation # ---------------------------------------------------- if pad_id is not None: logits[ pad_id ] = -float("inf") # ---------------------------------------------------- # Repetition penalty # ---------------------------------------------------- response_tokens = ( generated[prompt_len:] ) for idx in set( response_tokens[-32:] ): if logits[idx] > 0: logits[idx] /= penalty else: logits[idx] *= penalty # ---------------------------------------------------- # Temperature # ---------------------------------------------------- if temperature <= 0: raise ValueError( "temperature must be > 0" ) logits /= temperature # ---------------------------------------------------- # Sample # ---------------------------------------------------- next_token = top_k_top_p_sample( logits, top_k=top_k, top_p=top_p ) # ---------------------------------------------------- # Stop tokens # ---------------------------------------------------- if ( next_token == eor_id or next_token == eos_id ): break generated.append( next_token ) input_ids = torch.cat( [ input_ids, torch.tensor( [[next_token]], device=device ) ], dim=1 ) # -------------------------------------------------------- # Decode # -------------------------------------------------------- new_tokens = generated[ prompt_len: ] response = tokenizer.decode( new_tokens, skip_special_tokens=True ) response = ( response .replace("", "") .strip() ) # -------------------------------------------------------- # Cleanup # -------------------------------------------------------- response = re.sub( r"[{}\\/]", "", response ) if response.startswith( "Assistant:" ): response = ( response[ len("Assistant:"): ] .strip() ) return response