GrimSqueaker commited on
Commit
d8552bf
·
verified ·
1 Parent(s): 7ab905c

Upload modeling_modern_protein.py

Browse files
Files changed (1) hide show
  1. modeling_modern_protein.py +400 -0
modeling_modern_protein.py ADDED
@@ -0,0 +1,400 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Modern Protein Language Model
3
+ =============================
4
+ A <200M parameter encoder combining ModernBERT architecture + ELECTRA-style
5
+ replaced token detection for protein sequence predictive tasks.
6
+
7
+ Key innovations over ESM-2:
8
+ 1. ModernBERT architecture: Pre-LN, RMSNorm, GeGLU, RoPE, FlashAttention
9
+ 2. ELECTRA-style discriminative pre-training (not just MLM)
10
+ 3. Deep & narrow design (24 layers, 512 hidden ~120M params)
11
+ 4. 30% masking rate with curriculum decay
12
+ 5. Span masking for structural motifs
13
+ """
14
+
15
+ import math
16
+ from dataclasses import dataclass
17
+ from typing import Optional, Tuple
18
+
19
+ import torch
20
+ import torch.nn as nn
21
+ import torch.nn.functional as F
22
+
23
+
24
+ # ---------------------------------------------------------------------------
25
+ # Config
26
+ # ---------------------------------------------------------------------------
27
+
28
+ @dataclass
29
+ class ModernProteinConfig:
30
+ vocab_size: int = 33 # 20 AA + special tokens (ESM-2 style)
31
+ hidden_size: int = 512
32
+ num_hidden_layers: int = 24
33
+ num_attention_heads: int = 16
34
+ intermediate_size: int = 1536 # GeGLU: 2/3 * 4 * hidden for same params as GELU
35
+ max_position_embeddings: int = 1024
36
+ layer_norm_eps: float = 1e-6
37
+ hidden_dropout_prob: float = 0.0
38
+ attention_probs_dropout_prob: float = 0.0
39
+ initializer_range: float = 0.02
40
+ rope_theta: float = 10000.0
41
+ use_rms_norm: bool = True
42
+ use_geglu: bool = True
43
+ use_flash_attn: bool = True
44
+ tie_word_embeddings: bool = True
45
+ # ELECTRA
46
+ generator_size_multiplier: float = 0.25 # small generator
47
+ discriminator_lambda: float = 50.0
48
+ # Masking
49
+ mask_prob: float = 0.30
50
+ mask_prob_end: float = 0.05
51
+ span_masking: bool = True
52
+ mean_span_length: float = 3.0
53
+
54
+
55
+ # ---------------------------------------------------------------------------
56
+ # Normalization
57
+ # ---------------------------------------------------------------------------
58
+
59
+ class RMSNorm(nn.Module):
60
+ def __init__(self, dim: int, eps: float = 1e-6):
61
+ super().__init__()
62
+ self.eps = eps
63
+ self.weight = nn.Parameter(torch.ones(dim))
64
+
65
+ def forward(self, x):
66
+ norm = x.norm(2, dim=-1, keepdim=True) * (x.size(-1) ** -0.5)
67
+ return self.weight * (x / (norm + self.eps))
68
+
69
+
70
+ # ---------------------------------------------------------------------------
71
+ # RoPE
72
+ # ---------------------------------------------------------------------------
73
+
74
+ def rotate_half(x):
75
+ x1, x2 = x.chunk(2, dim=-1)
76
+ return torch.cat([-x2, x1], dim=-1)
77
+
78
+
79
+ def apply_rotary_pos_emb(q, k, cos, sin):
80
+ q_embed = (q * cos) + (rotate_half(q) * sin)
81
+ k_embed = (k * cos) + (rotate_half(k) * sin)
82
+ return q_embed, k_embed
83
+
84
+
85
+ class RotaryEmbedding(nn.Module):
86
+ def __init__(self, dim: int, max_seq_len: int = 2048, base: float = 10000.0):
87
+ super().__init__()
88
+ inv_freq = 1.0 / (base ** (torch.arange(0, dim, 2).float() / dim))
89
+ self.register_buffer("inv_freq", inv_freq)
90
+ self.max_seq_len = max_seq_len
91
+ self.dim = dim
92
+ t = torch.arange(max_seq_len, dtype=self.inv_freq.dtype)
93
+ freqs = torch.einsum("i,j->ij", t, self.inv_freq)
94
+ emb = torch.cat([freqs, freqs], dim=-1)
95
+ self.register_buffer("cos_cached", emb.cos()[None, None, :, :])
96
+ self.register_buffer("sin_cached", emb.sin()[None, None, :, :])
97
+
98
+ def forward(self, seq_len: int):
99
+ return (
100
+ self.cos_cached[:, :, :seq_len, :],
101
+ self.sin_cached[:, :, :seq_len, :],
102
+ )
103
+
104
+
105
+ # ---------------------------------------------------------------------------
106
+ # Attention
107
+ # ---------------------------------------------------------------------------
108
+
109
+ class ModernProteinAttention(nn.Module):
110
+ def __init__(self, config: ModernProteinConfig):
111
+ super().__init__()
112
+ self.num_heads = config.num_attention_heads
113
+ self.head_dim = config.hidden_size // config.num_attention_heads
114
+ self.scale = self.head_dim ** -0.5
115
+
116
+ self.qkv = nn.Linear(config.hidden_size, 3 * config.hidden_size, bias=False)
117
+ self.out_proj = nn.Linear(config.hidden_size, config.hidden_size, bias=False)
118
+ self.dropout = nn.Dropout(config.attention_probs_dropout_prob)
119
+ self.rotary = RotaryEmbedding(self.head_dim, config.max_position_embeddings, config.rope_theta)
120
+
121
+ def forward(self, x, attention_mask=None):
122
+ bsz, seq_len, _ = x.shape
123
+ qkv = self.qkv(x)
124
+ q, k, v = qkv.chunk(3, dim=-1)
125
+ q = q.view(bsz, seq_len, self.num_heads, self.head_dim).transpose(1, 2)
126
+ k = k.view(bsz, seq_len, self.num_heads, self.head_dim).transpose(1, 2)
127
+ v = v.view(bsz, seq_len, self.num_heads, self.head_dim).transpose(1, 2)
128
+
129
+ cos, sin = self.rotary(seq_len)
130
+ q, k = apply_rotary_pos_emb(q, k, cos, sin)
131
+
132
+ # FlashAttention via scaled_dot_product_attention
133
+ attn_output = F.scaled_dot_product_attention(
134
+ q, k, v,
135
+ attn_mask=attention_mask,
136
+ dropout_p=self.dropout.p if self.training else 0.0,
137
+ is_causal=False,
138
+ )
139
+ attn_output = attn_output.transpose(1, 2).contiguous().view(bsz, seq_len, -1)
140
+ return self.out_proj(attn_output)
141
+
142
+
143
+ # ---------------------------------------------------------------------------
144
+ # MLP
145
+ # ---------------------------------------------------------------------------
146
+
147
+ class GeGLU(nn.Module):
148
+ def __init__(self, config: ModernProteinConfig):
149
+ super().__init__()
150
+ self.w1 = nn.Linear(config.hidden_size, config.intermediate_size, bias=False)
151
+ self.w2 = nn.Linear(config.hidden_size, config.intermediate_size, bias=False)
152
+ self.w3 = nn.Linear(config.intermediate_size, config.hidden_size, bias=False)
153
+
154
+ def forward(self, x):
155
+ return self.w3(F.gelu(self.w1(x)) * self.w2(x))
156
+
157
+
158
+ class ModernProteinMLP(nn.Module):
159
+ def __init__(self, config: ModernProteinConfig):
160
+ super().__init__()
161
+ if config.use_geglu:
162
+ self.mlp = GeGLU(config)
163
+ else:
164
+ self.mlp = nn.Sequential(
165
+ nn.Linear(config.hidden_size, config.intermediate_size, bias=False),
166
+ nn.GELU(),
167
+ nn.Linear(config.intermediate_size, config.hidden_size, bias=False),
168
+ )
169
+
170
+ def forward(self, x):
171
+ return self.mlp(x)
172
+
173
+
174
+ # ---------------------------------------------------------------------------
175
+ # Transformer Layer
176
+ # ---------------------------------------------------------------------------
177
+
178
+ class ModernProteinLayer(nn.Module):
179
+ def __init__(self, config: ModernProteinConfig):
180
+ super().__init__()
181
+ Norm = RMSNorm if config.use_rms_norm else nn.LayerNorm
182
+ self.ln1 = Norm(config.hidden_size, eps=config.layer_norm_eps)
183
+ self.attn = ModernProteinAttention(config)
184
+ self.ln2 = Norm(config.hidden_size, eps=config.layer_norm_eps)
185
+ self.mlp = ModernProteinMLP(config)
186
+
187
+ def forward(self, x, attention_mask=None):
188
+ x = x + self.attn(self.ln1(x), attention_mask)
189
+ x = x + self.mlp(self.ln2(x))
190
+ return x
191
+
192
+
193
+ # ---------------------------------------------------------------------------
194
+ # Backbone
195
+ # ---------------------------------------------------------------------------
196
+
197
+ class ModernProteinEncoder(nn.Module):
198
+ def __init__(self, config: ModernProteinConfig):
199
+ super().__init__()
200
+ self.config = config
201
+ self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size)
202
+ self.layers = nn.ModuleList([ModernProteinLayer(config) for _ in range(config.num_hidden_layers)])
203
+ Norm = RMSNorm if config.use_rms_norm else nn.LayerNorm
204
+ self.ln_final = Norm(config.hidden_size, eps=config.layer_norm_eps)
205
+ self._init_weights()
206
+
207
+ def _init_weights(self):
208
+ for module in self.modules():
209
+ if isinstance(module, nn.Linear):
210
+ nn.init.normal_(module.weight, mean=0.0, std=self.config.initializer_range)
211
+ if module.bias is not None:
212
+ nn.init.zeros_(module.bias)
213
+ elif isinstance(module, nn.Embedding):
214
+ nn.init.normal_(module.weight, mean=0.0, std=self.config.initializer_range)
215
+
216
+ def forward(self, input_ids, attention_mask=None):
217
+ x = self.embed_tokens(input_ids)
218
+ for layer in self.layers:
219
+ x = layer(x, attention_mask)
220
+ return self.ln_final(x)
221
+
222
+
223
+ # ---------------------------------------------------------------------------
224
+ # ELECTRA: Generator + Discriminator
225
+ # ---------------------------------------------------------------------------
226
+
227
+ class ModernProteinForELECTRA(nn.Module):
228
+ """
229
+ ELECTRA-style pre-training for proteins.
230
+ Small generator predicts masked tokens.
231
+ Discriminator predicts whether each token is original or replaced.
232
+ """
233
+ def __init__(self, config: ModernProteinConfig):
234
+ super().__init__()
235
+ self.config = config
236
+ self.discriminator = ModernProteinEncoder(config)
237
+ self.discriminator_head = nn.Linear(config.hidden_size, 1)
238
+
239
+ # Smaller generator
240
+ gen_config = ModernProteinConfig(
241
+ vocab_size=config.vocab_size,
242
+ hidden_size=int(config.hidden_size * config.generator_size_multiplier),
243
+ num_hidden_layers=max(1, config.num_hidden_layers // 2),
244
+ num_attention_heads=max(2, config.num_attention_heads // 2),
245
+ intermediate_size=int(config.intermediate_size * config.generator_size_multiplier),
246
+ max_position_embeddings=config.max_position_embeddings,
247
+ layer_norm_eps=config.layer_norm_eps,
248
+ use_rms_norm=config.use_rms_norm,
249
+ use_geglu=config.use_geglu,
250
+ tie_word_embeddings=False,
251
+ )
252
+ self.generator = ModernProteinEncoder(gen_config)
253
+ self.generator_head = nn.Linear(gen_config.hidden_size, config.vocab_size, bias=False)
254
+
255
+ def forward(self, input_ids, attention_mask=None, labels=None, is_replaced=None):
256
+ # Generator: predict masked tokens
257
+ gen_hidden = self.generator(input_ids, attention_mask)
258
+ gen_logits = self.generator_head(gen_hidden)
259
+
260
+ # Sample replacements from generator
261
+ with torch.no_grad():
262
+ sampled_tokens = torch.argmax(gen_logits, dim=-1)
263
+
264
+ # Create corrupted input
265
+ corrupted_input = input_ids.clone()
266
+ mask = (input_ids == 32) # mask token id
267
+ corrupted_input[mask] = sampled_tokens[mask]
268
+
269
+ # Discriminator: detect replaced tokens
270
+ disc_hidden = self.discriminator(corrupted_input, attention_mask)
271
+ disc_logits = self.discriminator_head(disc_hidden).squeeze(-1)
272
+
273
+ loss = None
274
+ if labels is not None and is_replaced is not None:
275
+ gen_loss = F.cross_entropy(
276
+ gen_logits.view(-1, self.config.vocab_size),
277
+ labels.view(-1),
278
+ ignore_index=-100,
279
+ )
280
+ disc_loss = F.binary_cross_entropy_with_logits(
281
+ disc_logits.view(-1),
282
+ is_replaced.view(-1).float(),
283
+ )
284
+ loss = gen_loss + self.config.discriminator_lambda * disc_loss
285
+
286
+ return {
287
+ "loss": loss,
288
+ "gen_logits": gen_logits,
289
+ "disc_logits": disc_logits,
290
+ }
291
+
292
+
293
+ # ---------------------------------------------------------------------------
294
+ # Fine-tuning heads
295
+ # ---------------------------------------------------------------------------
296
+
297
+ class ModernProteinForSequenceClassification(nn.Module):
298
+ def __init__(self, config: ModernProteinConfig, num_labels: int):
299
+ super().__init__()
300
+ self.encoder = ModernProteinEncoder(config)
301
+ self.classifier = nn.Linear(config.hidden_size, num_labels)
302
+
303
+ def forward(self, input_ids, attention_mask=None, labels=None):
304
+ hidden = self.encoder(input_ids, attention_mask)
305
+ pooled = hidden[:, 0] # CLS token
306
+ logits = self.classifier(pooled)
307
+ loss = None
308
+ if labels is not None:
309
+ if self.classifier.out_features == 1:
310
+ loss = F.mse_loss(logits.squeeze(), labels.float())
311
+ else:
312
+ loss = F.cross_entropy(logits, labels)
313
+ return {"loss": loss, "logits": logits}
314
+
315
+
316
+ class ModernProteinForTokenClassification(nn.Module):
317
+ def __init__(self, config: ModernProteinConfig, num_labels: int):
318
+ super().__init__()
319
+ self.encoder = ModernProteinEncoder(config)
320
+ self.classifier = nn.Linear(config.hidden_size, num_labels)
321
+
322
+ def forward(self, input_ids, attention_mask=None, labels=None):
323
+ hidden = self.encoder(input_ids, attention_mask)
324
+ logits = self.classifier(hidden)
325
+ loss = None
326
+ if labels is not None:
327
+ loss = F.cross_entropy(
328
+ logits.view(-1, self.classifier.out_features),
329
+ labels.view(-1),
330
+ ignore_index=-100,
331
+ )
332
+ return {"loss": loss, "logits": logits}
333
+
334
+
335
+ # ---------------------------------------------------------------------------
336
+ # Masking utilities
337
+ # ---------------------------------------------------------------------------
338
+
339
+ def span_mask_tokens(input_ids, mask_token_id, vocab_size, mask_prob=0.30,
340
+ mean_span_length=3.0, pad_token_id=1):
341
+ """
342
+ Span masking for protein sequences.
343
+ Masks contiguous spans (simulating structural motif masking).
344
+ """
345
+ batch_size, seq_len = input_ids.shape
346
+ masked_input = input_ids.clone()
347
+ labels = input_ids.clone()
348
+ labels.fill_(-100)
349
+ is_replaced = torch.zeros_like(input_ids, dtype=torch.float)
350
+
351
+ for b in range(batch_size):
352
+ valid_len = (input_ids[b] != pad_token_id).sum().item()
353
+ num_to_mask = int(valid_len * mask_prob)
354
+ masked_count = 0
355
+
356
+ while masked_count < num_to_mask:
357
+ span_len = max(1, int(torch.poisson(torch.tensor(mean_span_length)).item()))
358
+ start = torch.randint(1, valid_len, (1,)).item() # avoid position 0 (CLS)
359
+ if start + span_len > valid_len:
360
+ span_len = valid_len - start
361
+ end = start + span_len
362
+
363
+ for pos in range(start, end):
364
+ if masked_count >= num_to_mask:
365
+ break
366
+ rand = torch.rand(1).item()
367
+ if rand < 0.8:
368
+ masked_input[b, pos] = mask_token_id
369
+ elif rand < 0.9:
370
+ masked_input[b, pos] = torch.randint(0, vocab_size, (1,)).item()
371
+ # else: keep original (10%)
372
+ labels[b, pos] = input_ids[b, pos]
373
+ is_replaced[b, pos] = 1.0
374
+ masked_count += 1
375
+
376
+ return masked_input, labels, is_replaced
377
+
378
+
379
+ # ---------------------------------------------------------------------------
380
+ # Count parameters
381
+ # ---------------------------------------------------------------------------
382
+
383
+ def count_parameters(model):
384
+ return sum(p.numel() for p in model.parameters() if p.requires_grad)
385
+
386
+
387
+ if __name__ == "__main__":
388
+ config = ModernProteinConfig()
389
+ model = ModernProteinForELECTRA(config)
390
+ print(f"Discriminator params: {count_parameters(model.discriminator) / 1e6:.1f}M")
391
+ print(f"Generator params: {count_parameters(model.generator) / 1e6:.1f}M")
392
+ print(f"Total params: {count_parameters(model) / 1e6:.1f}M")
393
+
394
+ # Test forward
395
+ batch_size, seq_len = 2, 128
396
+ input_ids = torch.randint(0, 33, (batch_size, seq_len))
397
+ input_ids[:, 0] = 0 # CLS
398
+ masked, labels, is_replaced = span_mask_tokens(input_ids, 32, 33)
399
+ out = model(masked, labels=labels, is_replaced=is_replaced)
400
+ print(f"Loss: {out['loss'].item():.4f}")