""" Modern Protein Language Model ============================= A <200M parameter encoder combining ModernBERT architecture + ELECTRA-style replaced token detection for protein sequence predictive tasks. Key innovations over ESM-2: 1. ModernBERT architecture: Pre-LN, RMSNorm, GeGLU, RoPE, FlashAttention 2. ELECTRA-style discriminative pre-training (not just MLM) 3. Deep & narrow design (24 layers, 512 hidden ~120M params) 4. 30% masking rate with curriculum decay 5. Span masking for structural motifs """ import math from dataclasses import dataclass from typing import Optional, Tuple import torch import torch.nn as nn import torch.nn.functional as F # --------------------------------------------------------------------------- # Config # --------------------------------------------------------------------------- @dataclass class ModernProteinConfig: vocab_size: int = 33 # 20 AA + special tokens (ESM-2 style) hidden_size: int = 512 num_hidden_layers: int = 24 num_attention_heads: int = 16 intermediate_size: int = 1536 # GeGLU: 2/3 * 4 * hidden for same params as GELU max_position_embeddings: int = 1024 layer_norm_eps: float = 1e-6 hidden_dropout_prob: float = 0.0 attention_probs_dropout_prob: float = 0.0 initializer_range: float = 0.02 rope_theta: float = 10000.0 use_rms_norm: bool = True use_geglu: bool = True use_flash_attn: bool = True tie_word_embeddings: bool = True # ELECTRA generator_size_multiplier: float = 0.25 # small generator discriminator_lambda: float = 50.0 # Masking mask_prob: float = 0.30 mask_prob_end: float = 0.05 span_masking: bool = True mean_span_length: float = 3.0 # --------------------------------------------------------------------------- # Normalization # --------------------------------------------------------------------------- class RMSNorm(nn.Module): def __init__(self, dim: int, eps: float = 1e-6): super().__init__() self.eps = eps self.weight = nn.Parameter(torch.ones(dim)) def forward(self, x): norm = x.norm(2, dim=-1, keepdim=True) * (x.size(-1) ** -0.5) return self.weight * (x / (norm + self.eps)) # --------------------------------------------------------------------------- # RoPE # --------------------------------------------------------------------------- def rotate_half(x): x1, x2 = x.chunk(2, dim=-1) return torch.cat([-x2, x1], dim=-1) def apply_rotary_pos_emb(q, k, cos, sin): q_embed = (q * cos) + (rotate_half(q) * sin) k_embed = (k * cos) + (rotate_half(k) * sin) return q_embed, k_embed class RotaryEmbedding(nn.Module): def __init__(self, dim: int, max_seq_len: int = 2048, base: float = 10000.0): super().__init__() inv_freq = 1.0 / (base ** (torch.arange(0, dim, 2).float() / dim)) self.register_buffer("inv_freq", inv_freq) self.max_seq_len = max_seq_len self.dim = dim t = torch.arange(max_seq_len, dtype=self.inv_freq.dtype) freqs = torch.einsum("i,j->ij", t, self.inv_freq) emb = torch.cat([freqs, freqs], dim=-1) self.register_buffer("cos_cached", emb.cos()[None, None, :, :]) self.register_buffer("sin_cached", emb.sin()[None, None, :, :]) def forward(self, seq_len: int): return ( self.cos_cached[:, :, :seq_len, :], self.sin_cached[:, :, :seq_len, :], ) # --------------------------------------------------------------------------- # Attention # --------------------------------------------------------------------------- class ModernProteinAttention(nn.Module): def __init__(self, config: ModernProteinConfig): super().__init__() self.num_heads = config.num_attention_heads self.head_dim = config.hidden_size // config.num_attention_heads self.scale = self.head_dim ** -0.5 self.qkv = nn.Linear(config.hidden_size, 3 * config.hidden_size, bias=False) self.out_proj = nn.Linear(config.hidden_size, config.hidden_size, bias=False) self.dropout = nn.Dropout(config.attention_probs_dropout_prob) self.rotary = RotaryEmbedding(self.head_dim, config.max_position_embeddings, config.rope_theta) def forward(self, x, attention_mask=None): bsz, seq_len, _ = x.shape qkv = self.qkv(x) q, k, v = qkv.chunk(3, dim=-1) q = q.view(bsz, seq_len, self.num_heads, self.head_dim).transpose(1, 2) k = k.view(bsz, seq_len, self.num_heads, self.head_dim).transpose(1, 2) v = v.view(bsz, seq_len, self.num_heads, self.head_dim).transpose(1, 2) cos, sin = self.rotary(seq_len) q, k = apply_rotary_pos_emb(q, k, cos, sin) # FlashAttention via scaled_dot_product_attention attn_output = F.scaled_dot_product_attention( q, k, v, attn_mask=attention_mask, dropout_p=self.dropout.p if self.training else 0.0, is_causal=False, ) attn_output = attn_output.transpose(1, 2).contiguous().view(bsz, seq_len, -1) return self.out_proj(attn_output) # --------------------------------------------------------------------------- # MLP # --------------------------------------------------------------------------- class GeGLU(nn.Module): def __init__(self, config: ModernProteinConfig): super().__init__() self.w1 = nn.Linear(config.hidden_size, config.intermediate_size, bias=False) self.w2 = nn.Linear(config.hidden_size, config.intermediate_size, bias=False) self.w3 = nn.Linear(config.intermediate_size, config.hidden_size, bias=False) def forward(self, x): return self.w3(F.gelu(self.w1(x)) * self.w2(x)) class ModernProteinMLP(nn.Module): def __init__(self, config: ModernProteinConfig): super().__init__() if config.use_geglu: self.mlp = GeGLU(config) else: self.mlp = nn.Sequential( nn.Linear(config.hidden_size, config.intermediate_size, bias=False), nn.GELU(), nn.Linear(config.intermediate_size, config.hidden_size, bias=False), ) def forward(self, x): return self.mlp(x) # --------------------------------------------------------------------------- # Transformer Layer # --------------------------------------------------------------------------- class ModernProteinLayer(nn.Module): def __init__(self, config: ModernProteinConfig): super().__init__() Norm = RMSNorm if config.use_rms_norm else nn.LayerNorm self.ln1 = Norm(config.hidden_size, eps=config.layer_norm_eps) self.attn = ModernProteinAttention(config) self.ln2 = Norm(config.hidden_size, eps=config.layer_norm_eps) self.mlp = ModernProteinMLP(config) def forward(self, x, attention_mask=None): x = x + self.attn(self.ln1(x), attention_mask) x = x + self.mlp(self.ln2(x)) return x # --------------------------------------------------------------------------- # Backbone # --------------------------------------------------------------------------- class ModernProteinEncoder(nn.Module): def __init__(self, config: ModernProteinConfig): super().__init__() self.config = config self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size) self.layers = nn.ModuleList([ModernProteinLayer(config) for _ in range(config.num_hidden_layers)]) Norm = RMSNorm if config.use_rms_norm else nn.LayerNorm self.ln_final = Norm(config.hidden_size, eps=config.layer_norm_eps) self._init_weights() def _init_weights(self): for module in self.modules(): if isinstance(module, nn.Linear): nn.init.normal_(module.weight, mean=0.0, std=self.config.initializer_range) if module.bias is not None: nn.init.zeros_(module.bias) elif isinstance(module, nn.Embedding): nn.init.normal_(module.weight, mean=0.0, std=self.config.initializer_range) def forward(self, input_ids, attention_mask=None): x = self.embed_tokens(input_ids) for layer in self.layers: x = layer(x, attention_mask) return self.ln_final(x) # --------------------------------------------------------------------------- # ELECTRA: Generator + Discriminator # --------------------------------------------------------------------------- class ModernProteinForELECTRA(nn.Module): """ ELECTRA-style pre-training for proteins. Small generator predicts masked tokens. Discriminator predicts whether each token is original or replaced. """ def __init__(self, config: ModernProteinConfig): super().__init__() self.config = config self.discriminator = ModernProteinEncoder(config) self.discriminator_head = nn.Linear(config.hidden_size, 1) # Smaller generator gen_config = ModernProteinConfig( vocab_size=config.vocab_size, hidden_size=int(config.hidden_size * config.generator_size_multiplier), num_hidden_layers=max(1, config.num_hidden_layers // 2), num_attention_heads=max(2, config.num_attention_heads // 2), intermediate_size=int(config.intermediate_size * config.generator_size_multiplier), max_position_embeddings=config.max_position_embeddings, layer_norm_eps=config.layer_norm_eps, use_rms_norm=config.use_rms_norm, use_geglu=config.use_geglu, tie_word_embeddings=False, ) self.generator = ModernProteinEncoder(gen_config) self.generator_head = nn.Linear(gen_config.hidden_size, config.vocab_size, bias=False) def forward(self, input_ids, attention_mask=None, labels=None, is_replaced=None): # Generator: predict masked tokens gen_hidden = self.generator(input_ids, attention_mask) gen_logits = self.generator_head(gen_hidden) # Sample replacements from generator with torch.no_grad(): sampled_tokens = torch.argmax(gen_logits, dim=-1) # Create corrupted input corrupted_input = input_ids.clone() mask = (input_ids == 32) # mask token id corrupted_input[mask] = sampled_tokens[mask] # Discriminator: detect replaced tokens disc_hidden = self.discriminator(corrupted_input, attention_mask) disc_logits = self.discriminator_head(disc_hidden).squeeze(-1) loss = None if labels is not None and is_replaced is not None: gen_loss = F.cross_entropy( gen_logits.view(-1, self.config.vocab_size), labels.view(-1), ignore_index=-100, ) disc_loss = F.binary_cross_entropy_with_logits( disc_logits.view(-1), is_replaced.view(-1).float(), ) loss = gen_loss + self.config.discriminator_lambda * disc_loss return { "loss": loss, "gen_logits": gen_logits, "disc_logits": disc_logits, } # --------------------------------------------------------------------------- # Fine-tuning heads # --------------------------------------------------------------------------- class ModernProteinForSequenceClassification(nn.Module): def __init__(self, config: ModernProteinConfig, num_labels: int): super().__init__() self.encoder = ModernProteinEncoder(config) self.classifier = nn.Linear(config.hidden_size, num_labels) def forward(self, input_ids, attention_mask=None, labels=None): hidden = self.encoder(input_ids, attention_mask) pooled = hidden[:, 0] # CLS token logits = self.classifier(pooled) loss = None if labels is not None: if self.classifier.out_features == 1: loss = F.mse_loss(logits.squeeze(), labels.float()) else: loss = F.cross_entropy(logits, labels) return {"loss": loss, "logits": logits} class ModernProteinForTokenClassification(nn.Module): def __init__(self, config: ModernProteinConfig, num_labels: int): super().__init__() self.encoder = ModernProteinEncoder(config) self.classifier = nn.Linear(config.hidden_size, num_labels) def forward(self, input_ids, attention_mask=None, labels=None): hidden = self.encoder(input_ids, attention_mask) logits = self.classifier(hidden) loss = None if labels is not None: loss = F.cross_entropy( logits.view(-1, self.classifier.out_features), labels.view(-1), ignore_index=-100, ) return {"loss": loss, "logits": logits} # --------------------------------------------------------------------------- # Masking utilities # --------------------------------------------------------------------------- def span_mask_tokens(input_ids, mask_token_id, vocab_size, mask_prob=0.30, mean_span_length=3.0, pad_token_id=1): """ Span masking for protein sequences. Masks contiguous spans (simulating structural motif masking). """ batch_size, seq_len = input_ids.shape masked_input = input_ids.clone() labels = input_ids.clone() labels.fill_(-100) is_replaced = torch.zeros_like(input_ids, dtype=torch.float) for b in range(batch_size): valid_len = (input_ids[b] != pad_token_id).sum().item() num_to_mask = int(valid_len * mask_prob) masked_count = 0 while masked_count < num_to_mask: span_len = max(1, int(torch.poisson(torch.tensor(mean_span_length)).item())) start = torch.randint(1, valid_len, (1,)).item() # avoid position 0 (CLS) if start + span_len > valid_len: span_len = valid_len - start end = start + span_len for pos in range(start, end): if masked_count >= num_to_mask: break rand = torch.rand(1).item() if rand < 0.8: masked_input[b, pos] = mask_token_id elif rand < 0.9: masked_input[b, pos] = torch.randint(0, vocab_size, (1,)).item() # else: keep original (10%) labels[b, pos] = input_ids[b, pos] is_replaced[b, pos] = 1.0 masked_count += 1 return masked_input, labels, is_replaced # --------------------------------------------------------------------------- # Count parameters # --------------------------------------------------------------------------- def count_parameters(model): return sum(p.numel() for p in model.parameters() if p.requires_grad) if __name__ == "__main__": config = ModernProteinConfig() model = ModernProteinForELECTRA(config) print(f"Discriminator params: {count_parameters(model.discriminator) / 1e6:.1f}M") print(f"Generator params: {count_parameters(model.generator) / 1e6:.1f}M") print(f"Total params: {count_parameters(model) / 1e6:.1f}M") # Test forward batch_size, seq_len = 2, 128 input_ids = torch.randint(0, 33, (batch_size, seq_len)) input_ids[:, 0] = 0 # CLS masked, labels, is_replaced = span_mask_tokens(input_ids, 32, 33) out = model(masked, labels=labels, is_replaced=is_replaced) print(f"Loss: {out['loss'].item():.4f}")