File size: 5,170 Bytes
b72d311
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
"""
ESM2-35M + LoRA regression model for protein ΔG prediction.

Reconstructs the architecture stored in best_model.pt (model_type=esm2_lora):
  - backbone : facebook/esm2_t12_35M_UR50D (EsmModel, 12 layers, hidden 480)
  - adapters : LoRA r=16, alpha=32 on attention query/key/value
  - head     : LayerNorm(480) -> Linear(480,256) -> GELU -> Dropout -> Linear(256,64) -> GELU -> Linear(64,1)
  - pooling  : attention-masked mean pool over token embeddings

The full ESM weights are bundled in the checkpoint, so no HuggingFace download
is needed — only the model config (built manually) and tokenizer (shared UR50D vocab).
"""

import os
os.environ.setdefault("HF_HUB_OFFLINE", "1")
os.environ.setdefault("TRANSFORMERS_OFFLINE", "1")

import torch
import torch.nn as nn
from transformers import EsmConfig, EsmModel, AutoTokenizer

# The model was trained with tokenizer max_length=82 (= 80 residues + BOS + EOS).
# Sequences longer than 80 aa are truncated to the first 80 — must match training/Colab
# exactly or predictions diverge (the gap grows with sequence length).
MAX_LEN = 80          # residue cap (what the user sees)
TOK_MAX_LEN = 82      # tokenizer max_length, mirrors Colab Config cell
BACKBONE = "facebook/esm2_t12_35M_UR50D"
_TOKENIZER_FALLBACK = "facebook/esm2_t6_8M_UR50D"  # identical UR50D vocab, cached locally


def _build_config() -> EsmConfig:
    # Config for esm2_t12_35M_UR50D (no download required)
    return EsmConfig(
        vocab_size=33,
        hidden_size=480,
        num_hidden_layers=12,
        num_attention_heads=20,
        intermediate_size=1920,
        max_position_embeddings=1026,
        position_embedding_type="rotary",
        token_dropout=True,
        emb_layer_norm_before=False,
        pad_token_id=1,
        mask_token_id=32,
    )


class ESM2LoRARegressor(nn.Module):
    def __init__(self, lora_r: int = 16, dropout: float = 0.3):
        super().__init__()
        from peft import LoraConfig, get_peft_model

        base = EsmModel(_build_config(), add_pooling_layer=True)
        lora_cfg = LoraConfig(
            r=lora_r,
            lora_alpha=lora_r * 2,
            target_modules=["query", "key", "value"],
            lora_dropout=0.0,
            bias="none",
        )
        self.esm = get_peft_model(base, lora_cfg)
        self.head = nn.Sequential(
            nn.LayerNorm(480),       # 0
            nn.Linear(480, 256),     # 1
            nn.GELU(),               # 2
            nn.Dropout(dropout),     # 3
            nn.Linear(256, 64),      # 4
            nn.GELU(),               # 5
            nn.Linear(64, 1),        # 6
        )

    def forward(self, input_ids, attention_mask):
        out = self.esm(input_ids=input_ids, attention_mask=attention_mask)
        hidden = out.last_hidden_state                      # (B, T, 480)
        mask = attention_mask.unsqueeze(-1).float()         # (B, T, 1)
        pooled = (hidden * mask).sum(1) / mask.sum(1).clamp(min=1e-9)
        return self.head(pooled).squeeze(-1)                # (B,)

    def count_parameters(self) -> int:
        return sum(p.numel() for p in self.parameters() if p.requires_grad)


_tokenizer = None


def _get_tokenizer():
    global _tokenizer
    if _tokenizer is None:
        try:
            _tokenizer = AutoTokenizer.from_pretrained(BACKBONE)
        except Exception:
            _tokenizer = AutoTokenizer.from_pretrained(_TOKENIZER_FALLBACK)
    return _tokenizer


def load_model(checkpoint_path: str, device: str = "cpu") -> ESM2LoRARegressor:
    ckpt = torch.load(checkpoint_path, map_location=device, weights_only=False)
    state = ckpt.get("state_dict", ckpt)
    model = ESM2LoRARegressor()
    missing, unexpected = model.load_state_dict(state, strict=False)
    # Report only non-trivial mismatches
    real_missing = [k for k in missing if "lora" not in k]
    if real_missing or unexpected:
        print(f"[esm2_lora] load: {len(missing)} missing, {len(unexpected)} unexpected")
        if real_missing[:5]:
            print("  e.g. missing:", real_missing[:5])
        if unexpected[:5]:
            print("  e.g. unexpected:", unexpected[:5])
    model.to(device)
    model.eval()
    return model


@torch.no_grad()
def predict_batch(seqs, model, device="cpu"):
    tok = _get_tokenizer()
    # Mirror Colab exactly: max_length=82, padding='max_length', truncation=True.
    # (Padding tokens are masked out of the mean pool, so padding mode does not affect output.)
    enc = tok([s.upper().strip() for s in seqs],
              return_tensors="pt", padding="max_length", truncation=True, max_length=TOK_MAX_LEN)
    enc = {k: v.to(device) for k, v in enc.items()}
    out = model(enc["input_ids"], enc["attention_mask"])
    vals = out.tolist()
    if isinstance(vals, float):
        vals = [vals]
    return [round(v, 4) for v in vals]


def predict_one(seq, model, device="cpu"):
    return predict_batch([seq], model, device)[0]


def stability_label(dg: float) -> str:
    if dg > 3.0:  return "highly stable"
    if dg > 0.5:  return "stable"
    if dg > -0.5: return "marginally stable"
    if dg > -3.0: return "unstable"
    return "highly unstable"