File size: 4,243 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
"""
ESM-2 masked marginal scoring for mutation effect prediction.

For each mutation (from_aa -> to_aa at position i):
  score = log P(to_aa | context_masked_at_i) - log P(from_aa | context_masked_at_i)

Positive score: ESM-2 prefers the mutant AA given the surrounding context.
Correlates with experimental ddG (Meier et al. 2021, ESM-1v).

Phase D (Step 12): Upgraded to facebook/esm2_t12_35M_UR50D (35M params, 480-dim).
  ~3× more parameters than the 8M model; better cross-protein generalisation.
  Override via env var: ESM_MODEL_NAME=facebook/esm2_t6_8M_UR50D (for CPU-only)

Falls back gracefully to score=0.0 if torch/transformers not installed.
"""

import logging, os
from typing import Optional

logger = logging.getLogger(__name__)

# Phase D Step 12: 35M model; configurable via env var for low-memory machines
_DEFAULT_ESM_MODEL = 'facebook/esm2_t12_35M_UR50D'
ESM_MODEL_NAME     = os.environ.get('ESM_MODEL_NAME', _DEFAULT_ESM_MODEL)
BATCH_SIZE         = 16          # reduced from 32: 35M model uses more VRAM
AMINO_ACIDS    = list('ACDEFGHIKLMNPQRSTVWY')

_load_attempted = False
_ESM_AVAILABLE  = False
_model          = None
_tokenizer      = None
_aa_token_ids   = None


def _try_load() -> bool:
    global _ESM_AVAILABLE, _model, _tokenizer, _aa_token_ids, _load_attempted
    if _load_attempted:
        return _ESM_AVAILABLE
    _load_attempted = True
    try:
        import torch                                             # noqa: F401
        from transformers import EsmForMaskedLM, EsmTokenizer

        logger.info('Loading %s ...', ESM_MODEL_NAME)
        _tokenizer    = EsmTokenizer.from_pretrained(ESM_MODEL_NAME)
        _model        = EsmForMaskedLM.from_pretrained(ESM_MODEL_NAME)
        _model.eval()
        _aa_token_ids = {aa: _tokenizer.convert_tokens_to_ids(aa) for aa in AMINO_ACIDS}
        _ESM_AVAILABLE = True
        logger.info('ESM-2 loaded.')
        return True
    except Exception as exc:
        logger.warning('ESM-2 unavailable (%s) — esm_masked_marginal will be 0.0.', exc)
        return False


def is_available() -> bool:
    return _try_load()


def get_masked_marginals(sequence: str) -> Optional[dict]:
    """
    Run batched masked-LM inference: for every position i, mask it and record
    log P(AA | context) for all 20 canonical AAs.

    Returns  dict[(pos_0based: int, aa: str)] -> float
    or       None if ESM-2 is unavailable.

    All L masked sequences have identical length so no padding is added —
    the token at index (pos + 1) in each output is always the masked position
    (ESM tokenizer inserts <cls> at index 0).
    """
    if not _try_load():
        return None

    import torch
    import torch.nn.functional as F

    seq  = sequence.upper()
    L    = len(seq)
    mask = _tokenizer.mask_token

    masked_seqs = [seq[:i] + mask + seq[i + 1:] for i in range(L)]
    scores: dict = {}

    for batch_start in range(0, L, BATCH_SIZE):
        batch          = masked_seqs[batch_start: batch_start + BATCH_SIZE]
        batch_pos_orig = list(range(batch_start, batch_start + len(batch)))

        inputs = _tokenizer(batch, return_tensors='pt', padding=True,
                            add_special_tokens=True)

        with torch.no_grad():
            logits = _model(**inputs).logits      # (B, T, vocab)

        for j, pos in enumerate(batch_pos_orig):
            token_pos = pos + 1                   # +1 for <cls>
            log_probs = F.log_softmax(logits[j, token_pos, :], dim=-1)
            for aa in AMINO_ACIDS:
                scores[(pos, aa)] = float(log_probs[_aa_token_ids[aa]])

    return scores


def masked_marginal_score(sequence: str, pos_0based: int,
                           from_aa: str, to_aa: str,
                           precomputed: Optional[dict] = None) -> float:
    """
    Convenience wrapper: returns log P(to_aa|ctx) - log P(from_aa|ctx).
    Pass precomputed=get_masked_marginals(sequence) to reuse across mutations.
    """
    data = precomputed if precomputed is not None else get_masked_marginals(sequence)
    if data is None:
        return 0.0
    lp_to   = data.get((pos_0based, to_aa),   -20.0)
    lp_from = data.get((pos_0based, from_aa),  -20.0)
    return float(lp_to - lp_from)