Download models/simplefold/mlx/esm_network.py from OneScience-Group/SimpleFold: direct link, hf CLI and curl.
- Browser
- Download file 5.35 kB
-
https://huggingface.co/OneScience-Group/SimpleFold/resolve/main/models/simplefold/mlx/esm_network.py
- Command line
-
hf download hf://OneScience-Group/SimpleFold/models/simplefold/mlx/esm_network.py
-
curl -L -o esm_network.py https://huggingface.co/OneScience-Group/SimpleFold/resolve/main/models/simplefold/mlx/esm_network.py
5.35 kB
| # | |
| # For licensing see accompanying LICENSE file. | |
| # Copyright (c) 2025 Apple Inc. Licensed under MIT License. | |
| # | |
| # Started from https://github.com/facebookresearch/esm/tree/main, | |
| # licensed under MIT License, Copyright (c) Meta Platforms, Inc. and affiliates. | |
| from typing import Union | |
| import mlx.core as mx | |
| import mlx.nn as nn | |
| from .simplefold.mlx.esm_modules import ( | |
| ContactPredictionHead, | |
| ESM1bLayerNorm, | |
| RobertaLMHead, | |
| TransformerLayer, | |
| ) | |
| from onescience.datapipes.esm import Alphabet | |
| def masked_fill_mlx(x, mask, value): | |
| return mx.where(mask, value, x) | |
| class ESM2(nn.Module): | |
| def __init__( | |
| self, | |
| num_layers: int = 33, | |
| embed_dim: int = 1280, | |
| attention_heads: int = 20, | |
| alphabet: Union[Alphabet, str] = "ESM-1b", | |
| token_dropout: bool = True, | |
| ): | |
| super().__init__() | |
| self.num_layers = num_layers | |
| self.embed_dim = embed_dim | |
| self.attention_heads = attention_heads | |
| if not isinstance(alphabet, Alphabet): | |
| alphabet = Alphabet.from_architecture(alphabet) | |
| self.alphabet = alphabet | |
| self.alphabet_size = len(alphabet) | |
| self.padding_idx = alphabet.padding_idx | |
| self.mask_idx = alphabet.mask_idx | |
| self.cls_idx = alphabet.cls_idx | |
| self.eos_idx = alphabet.eos_idx | |
| self.prepend_bos = alphabet.prepend_bos | |
| self.append_eos = alphabet.append_eos | |
| self.token_dropout = token_dropout | |
| self._init_submodules() | |
| def _init_submodules(self): | |
| self.embed_scale = 1 | |
| self.embed_tokens = mx.zeros((self.alphabet_size, self.embed_dim)) | |
| self.layers = [ | |
| TransformerLayer( | |
| self.embed_dim, | |
| 4 * self.embed_dim, | |
| self.attention_heads, | |
| add_bias_kv=False, | |
| use_esm1b_layer_norm=True, | |
| use_rotary_embeddings=True, | |
| ) | |
| for _ in range(self.num_layers) | |
| ] | |
| self.contact_head = ContactPredictionHead( | |
| self.num_layers * self.attention_heads, | |
| self.prepend_bos, | |
| self.append_eos, | |
| eos_idx=self.eos_idx, | |
| ) | |
| self.emb_layer_norm_after = ESM1bLayerNorm(self.embed_dim) | |
| self.lm_head = RobertaLMHead( | |
| embed_dim=self.embed_dim, | |
| output_dim=self.alphabet_size, | |
| weight=self.embed_tokens, | |
| ) | |
| def __call__( | |
| self, tokens, repr_layers=[], need_head_weights=False, return_contacts=False | |
| ): | |
| if return_contacts: | |
| need_head_weights = True | |
| assert tokens.ndim == 2 | |
| padding_mask = mx.equal(tokens, self.padding_idx) # B, T | |
| x = self.embed_scale * self.embed_tokens[tokens, :] | |
| if self.token_dropout: | |
| x = masked_fill_mlx(x, (tokens == self.mask_idx)[..., None], 0.0) | |
| # x: B x T x C | |
| mask_ratio_train = 0.15 * 0.8 | |
| src_lengths = (~padding_mask).sum(axis=-1) | |
| mask_ratio_observed = (tokens == self.mask_idx).sum(axis=-1).astype( | |
| x.dtype | |
| ) / src_lengths | |
| x = x * (1 - mask_ratio_train) / (1 - mask_ratio_observed)[:, None, None] | |
| if padding_mask is not None: | |
| x = x * (1 - padding_mask[..., None].astype(x.dtype)) | |
| repr_layers = set(repr_layers) | |
| hidden_representations = {} | |
| if 0 in repr_layers: | |
| hidden_representations[0] = x | |
| if need_head_weights: | |
| attn_weights = [] | |
| # (B, T, E) => (T, B, E) | |
| x = mx.swapaxes(x, axis1=0, axis2=1) | |
| if not padding_mask.any(): | |
| padding_mask = None | |
| for layer_idx, layer in enumerate(self.layers): | |
| x, attn = layer( | |
| x, | |
| self_attn_padding_mask=padding_mask, | |
| need_head_weights=need_head_weights, | |
| ) | |
| if (layer_idx + 1) in repr_layers: | |
| hidden_representations[layer_idx + 1] = mx.swapaxes(x, axis1=0, axis2=1) | |
| if need_head_weights: | |
| # (H, B, T, T) => (B, H, T, T) | |
| attn_weights.append(mx.swapaxes(attn, axis1=1, axis2=0)) | |
| x = self.emb_layer_norm_after(x) | |
| x = mx.swapaxes(x, axis1=0, axis2=1) # (T, B, E) => (B, T, E) | |
| # last hidden representation should have layer norm applied | |
| if (layer_idx + 1) in repr_layers: | |
| hidden_representations[layer_idx + 1] = x | |
| x = self.lm_head(x) | |
| result = {"logits": x, "representations": hidden_representations} | |
| if need_head_weights: | |
| # attentions: B x L x H x T x T | |
| attentions = mx.stack(attn_weights, axis=1) | |
| if padding_mask is not None: | |
| attention_mask = 1 - padding_mask.astype(attentions.dtype) | |
| attention_mask = ( | |
| attention_mask[:, None, ...] * attention_mask[:, :, None, ...] | |
| ) | |
| attentions = attentions * attention_mask[:, None, None, :, :] | |
| result["attentions"] = attentions | |
| if return_contacts: | |
| contacts = self.contact_head(tokens, attentions) | |
| result["contacts"] = contacts | |
| return result | |
| def predict_contacts(self, tokens): | |
| return self(tokens, return_contacts=True)["contacts"] | |