Spaces:
Sleeping
Sleeping
File size: 7,705 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 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 | """
ESM2-150M + LoRA + environmental-gating regressor for protein ΔG prediction.
Architecture supplied by the model author, verified layer-by-layer (strict
`load_state_dict`, zero missing/unexpected keys) against the trained checkpoint:
- backbone : facebook/esm2_t30_150M_UR50D (EsmModel, 30 layers, hidden 640)
- adapters : LoRA r=32, alpha=64, dropout=0.1 on every submodule literally
named query/key/value/dense (attention q/k/v, attention output
projection, and both FFN dense layers)
- env_gate : Linear(2,64) -> SiLU -> Linear(64,64) -> Sigmoid — takes
[temperature, pH] and produces a per-feature multiplicative gate
- projector: Linear(640,64) — the checkpoint has ONLY this single Linear
(no LayerNorm/second Linear — the author's reference script had
a fuller Sequential here, but those extra layers have no
matching weights in this specific checkpoint, so they're
omitted to match what was actually trained)
- fusion : protein_feats * (1 + gate) — residual gating, never fully
zeroes the sequence signal
- head : Linear(64,32) -> SiLU -> Dropout(0.2) -> Linear(32,1)
- pooling : attention-masked mean pool over token embeddings (not the
ESM pooler — the pooler's LoRA weights exist in the checkpoint
because target_modules=["...", "dense"] matches pooler.dense
too, but forward() here never calls it)
UNCONFIRMED — do not treat predictions from this model as trustworthy until
verified with the model author:
- ENV_FEATURE normalization: temperature/pH are currently passed through
RAW (see `_env_tensor`). If the author normalized them during training
(min-max, z-score, etc.), predictions will be systematically wrong until
this is corrected to match.
- TOK_MAX_LEN: not recoverable from the checkpoint. Placeholder below;
confirm the exact tokenizer max_length/truncation used in training.
- This checkpoint has no `model_type`/`model_name` tag, so it is detected
at load time via the presence of "env_gate." keys — see protstab_predict.py.
"""
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
BACKBONE = "facebook/esm2_t30_150M_UR50D"
_TOKENIZER_FALLBACK = "facebook/esm2_t12_35M_UR50D" # shared UR50D vocab
# TODO(confirm with model author): exact training truncation length.
MAX_LEN = 512
TOK_MAX_LEN = MAX_LEN + 2 # + BOS/EOS
LORA_R = 32
LORA_ALPHA = 64
# TODO(confirm with model author): were these raw values, or normalized
# (min-max / z-score) before training? Using raw °C / pH units for now.
DEFAULT_TEMPERATURE_C = 37.0
DEFAULT_PH = 7.0
def _build_config() -> EsmConfig:
# facebook/esm2_t30_150M_UR50D — no network download required.
return EsmConfig(
vocab_size=33,
hidden_size=640,
num_hidden_layers=30,
num_attention_heads=20,
intermediate_size=2560,
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 ESM2GatedStabilityModel(nn.Module):
def __init__(self):
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_ALPHA,
target_modules=["query", "key", "value", "dense"],
lora_dropout=0.1,
bias="none",
)
self.esm2 = get_peft_model(base, lora_cfg)
self.env_gate = nn.Sequential(
nn.Linear(2, 64),
nn.SiLU(),
nn.Linear(64, 64),
nn.Sigmoid(),
)
# Matches the checkpoint exactly: single Linear, no LayerNorm/extra
# Linear (see module docstring).
self.protein_projector = nn.Sequential(
nn.Linear(640, 64),
)
self.regression_head = nn.Sequential(
nn.Linear(64, 32),
nn.SiLU(),
nn.Dropout(0.2),
nn.Linear(32, 1),
)
def forward(self, input_ids, attention_mask, env_features):
out = self.esm2(input_ids=input_ids, attention_mask=attention_mask)
hidden = out.last_hidden_state # (B, T, 640)
mask = attention_mask.unsqueeze(-1).float()
pooled = (hidden * mask).sum(1) / mask.sum(1).clamp(min=1e-9)
protein_feats = self.protein_projector(pooled) # (B, 64)
gate = self.env_gate(env_features) # (B, 64)
modulated = protein_feats * (1.0 + gate)
return self.regression_head(modulated).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") -> ESM2GatedStabilityModel:
ckpt = torch.load(checkpoint_path, map_location=device, weights_only=False)
state = ckpt.get("state_dict", ckpt)
model = ESM2GatedStabilityModel()
missing, unexpected = model.load_state_dict(state, strict=False)
if missing or unexpected:
print(f"[esm2_gated] load: {len(missing)} missing, {len(unexpected)} unexpected")
if missing[:5]:
print(" e.g. missing:", missing[:5])
if unexpected[:5]:
print(" e.g. unexpected:", unexpected[:5])
model.to(device)
model.eval()
return model
def _env_tensor(conditions_list, device) -> torch.Tensor:
"""conditions_list: list of (temperature_c, ph) tuples, one per sequence.
RAW pass-through — see UNCONFIRMED note in the module docstring."""
return torch.tensor(conditions_list, dtype=torch.float32, device=device)
@torch.no_grad()
def predict_batch(seqs, model, device="cpu", conditions=None):
"""
conditions: optional list of {"temperature": float, "ph": float} dicts,
same length as seqs. Missing entries fall back to DEFAULT_TEMPERATURE_C /
DEFAULT_PH (physiological defaults — NOT confirmed to match the training
distribution; see module docstring).
"""
tok = _get_tokenizer()
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()}
conditions = conditions or [{}] * len(seqs)
env_list = [
(c.get("temperature", DEFAULT_TEMPERATURE_C) or DEFAULT_TEMPERATURE_C,
c.get("ph", DEFAULT_PH) or DEFAULT_PH)
for c in conditions
]
env = _env_tensor(env_list, device)
out = model(enc["input_ids"], enc["attention_mask"], env)
vals = out.tolist()
if isinstance(vals, float):
vals = [vals]
return [round(v, 4) for v in vals]
def predict_one(seq, model, device="cpu", conditions=None):
return predict_batch([seq], model, device, conditions=[conditions or {}])[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"
|