| |
| |
| |
| |
|
|
| import math |
| import mlx.nn as nn |
| import mlx.core as mx |
| from models.simplefold.mlx.layers import FinalLayer, ConditionEmbedder |
| from onescience.utils.simplefold.esm_utils import esm_model_dict |
|
|
|
|
| |
| def one_hot(indices, num_classes, dtype=None): |
| """ |
| MLX version of torch.one_hot. |
| |
| Args: |
| indices: integer MLX array of any shape, containing class indices in [0, num_classes). |
| num_classes: number of classes for the one-hot dimension. |
| dtype: output data type (defaults to float32). |
| |
| Returns: |
| MLX array of shape indices.shape + (num_classes,) and given dtype. |
| """ |
| |
| if dtype is None: |
| dtype = mx.float32 |
|
|
| classes = mx.arange(num_classes, dtype=indices.dtype) |
|
|
| |
| |
| mask = indices[..., mx.newaxis] == classes |
|
|
| |
| return mask.astype(dtype) |
|
|
|
|
| class FoldingDiT(nn.Module): |
| def __init__( |
| self, |
| trunk, |
| time_embedder, |
| aminoacid_pos_embedder, |
| pos_embedder, |
| atom_encoder_transformer, |
| atom_decoder_transformer, |
| hidden_size=1152, |
| num_heads=16, |
| atom_num_heads=4, |
| output_channels=3, |
| atom_hidden_size_enc=256, |
| atom_hidden_size_dec=256, |
| atom_n_queries_enc=32, |
| atom_n_keys_enc=128, |
| atom_n_queries_dec=32, |
| atom_n_keys_dec=128, |
| esm_model="esm2_3B", |
| esm_dropout_prob=0.0, |
| use_atom_mask=False, |
| use_length_condition=True, |
| ): |
| super().__init__() |
| self.pos_embedder = pos_embedder |
| pos_embed_channels = pos_embedder.embed_dim |
| self.aminoacid_pos_embedder = aminoacid_pos_embedder |
| aminoacid_pos_embed_channels = aminoacid_pos_embedder.embed_dim |
|
|
| self.time_embedder = time_embedder |
|
|
| self.atom_encoder_transformer = atom_encoder_transformer |
| self.atom_decoder_transformer = atom_decoder_transformer |
|
|
| self.trunk = trunk |
|
|
| self.hidden_size = hidden_size |
| self.output_channels = output_channels |
| self.num_heads = num_heads |
| self.atom_num_heads = atom_num_heads |
| self.use_atom_mask = use_atom_mask |
| self.esm_dropout_prob = esm_dropout_prob |
| self.use_length_condition = use_length_condition |
|
|
| esm_s_dim = esm_model_dict[esm_model]["esm_s_dim"] |
| esm_num_layers = esm_model_dict[esm_model]["esm_num_layers"] |
|
|
| self.atom_hidden_size_enc = atom_hidden_size_enc |
| self.atom_hidden_size_dec = atom_hidden_size_dec |
| self.atom_n_queries_enc = atom_n_queries_enc |
| self.atom_n_keys_enc = atom_n_keys_enc |
| self.atom_n_queries_dec = atom_n_queries_dec |
| self.atom_n_keys_dec = atom_n_keys_dec |
|
|
| atom_feat_dim = pos_embed_channels + aminoacid_pos_embed_channels + 427 |
| self.atom_feat_proj = nn.Sequential( |
| nn.Linear(atom_feat_dim, hidden_size), |
| nn.LayerNorm(hidden_size), |
| nn.SiLU(), |
| ) |
| self.atom_pos_proj = nn.Linear(pos_embed_channels, hidden_size, bias=False) |
|
|
| if self.use_length_condition: |
| self.length_embedder = nn.Sequential( |
| nn.Linear(1, hidden_size, bias=False), |
| nn.LayerNorm(hidden_size), |
| ) |
|
|
| self.atom_in_proj = nn.Linear(hidden_size * 2, hidden_size, bias=False) |
|
|
| self.esm_s_combine = mx.zeros(esm_num_layers) |
| self.esm_s_proj = ConditionEmbedder( |
| input_dim=esm_s_dim, |
| hidden_size=hidden_size, |
| dropout_prob=0, |
| ) |
|
|
| latent_cat_dim = hidden_size * 2 |
| self.esm_cat_proj = nn.Linear(latent_cat_dim, hidden_size) |
|
|
| self.context2atom_proj = nn.Sequential( |
| nn.Linear(hidden_size, self.atom_hidden_size_enc), |
| nn.LayerNorm(self.atom_hidden_size_enc), |
| ) |
|
|
| self.atom_enc_cond_proj = nn.Sequential( |
| nn.Linear(hidden_size, self.atom_hidden_size_enc), |
| nn.LayerNorm(self.atom_hidden_size_enc), |
| ) |
|
|
| self.atom2latent_proj = nn.Sequential( |
| nn.Linear(self.atom_hidden_size_enc, hidden_size), |
| nn.LayerNorm(hidden_size), |
| ) |
|
|
| self.atom_dec_cond_proj = nn.Sequential( |
| nn.Linear(hidden_size, self.atom_hidden_size_dec), |
| nn.LayerNorm(self.atom_hidden_size_dec), |
| ) |
|
|
| self.latent2atom_proj = nn.Sequential( |
| nn.Linear(hidden_size, hidden_size), |
| nn.SiLU(), |
| nn.LayerNorm(hidden_size), |
| nn.Linear(hidden_size, self.atom_hidden_size_dec), |
| ) |
|
|
| self.final_layer = FinalLayer( |
| self.atom_hidden_size_dec, output_channels, c_dim=hidden_size |
| ) |
|
|
| def create_local_attn_bias( |
| self, |
| n: int, |
| n_queries: int, |
| n_keys: int, |
| inf: float = 1e10, |
| ): |
| """Create local attention bias based on query window n_queries and kv window n_keys. |
| |
| Args: |
| n (int): the length of quiries |
| n_queries (int): window size of quiries |
| n_keys (int): window size of keys/values |
| inf (float, optional): the inf to mask attention. Defaults to 1e10. |
| device (torch.device, optional): cuda|cpu|None. Defaults to None. |
| |
| Returns: |
| torch.Tensor: the diagonal-like global attention bias |
| """ |
| n_trunks = int(math.ceil(n / n_queries)) |
| padded_n = n_trunks * n_queries |
| attn_mask = mx.zeros((padded_n, padded_n)) |
| for block_index in range(0, n_trunks): |
| i = block_index * n_queries |
| j1 = max(0, n_queries * block_index - (n_keys - n_queries) // 2) |
| j2 = n_queries * block_index + (n_queries + n_keys) // 2 |
| attn_mask[i : i + n_queries, j1:j2] = 1.0 |
| attn_bias = (1 - attn_mask) * -inf |
| return attn_bias[:n, :n] |
|
|
| def create_atom_attn_mask( |
| self, feats, natoms, atom_n_queries=None, atom_n_keys=None, inf: float = 1e10 |
| ): |
|
|
| if atom_n_queries is not None and atom_n_keys is not None: |
| atom_attn_mask = self.create_local_attn_bias( |
| n=natoms, n_queries=atom_n_queries, n_keys=atom_n_keys, inf=inf |
| ) |
| else: |
| atom_attn_mask = None |
|
|
| return atom_attn_mask |
|
|
| def __call__(self, noised_pos, t, feats, self_cond=None): |
|
|
| B, N, _ = feats["ref_pos"].shape |
| M = feats["mol_type"].shape[1] |
| atom_to_token = feats["atom_to_token"].astype(mx.float32) |
| atom_to_token_idx = feats["atom_to_token_idx"] |
| ref_space_uid = feats["ref_space_uid"] |
|
|
| |
| atom_attn_mask_enc = self.create_atom_attn_mask( |
| feats, |
| natoms=N, |
| atom_n_queries=self.atom_n_queries_enc, |
| atom_n_keys=self.atom_n_keys_enc, |
| ) |
| atom_attn_mask_dec = self.create_atom_attn_mask( |
| feats, |
| natoms=N, |
| atom_n_queries=self.atom_n_queries_dec, |
| atom_n_keys=self.atom_n_keys_dec, |
| ) |
|
|
| |
| c_emb = self.time_embedder(t) |
| if self.use_length_condition: |
| length = feats["max_num_tokens"].astype(mx.float32)[..., None] |
| c_emb = c_emb + self.length_embedder(mx.log(length)) |
|
|
| mol_type = feats["mol_type"] |
| mol_type = one_hot(mol_type, num_classes=4).astype(mx.float32) |
| res_type = feats["res_type"].astype(mx.float32) |
| pocket_feature = feats["pocket_feature"].astype(mx.float32) |
| res_feat = mx.concatenate( |
| [mol_type, res_type, pocket_feature], axis=-1 |
| ) |
| atom_feat_from_res = mx.matmul(atom_to_token, res_feat) |
| atom_res_pos = self.aminoacid_pos_embedder( |
| pos=atom_to_token_idx[..., None].astype(mx.float32) |
| ) |
| ref_pos_emb = self.pos_embedder(pos=feats["ref_pos"]) |
| atom_feat = mx.concatenate( |
| [ |
| ref_pos_emb, |
| atom_feat_from_res, |
| atom_res_pos, |
| feats["ref_charge"][..., None], |
| feats["atom_pad_mask"][..., None], |
| feats["ref_element"], |
| feats["ref_atom_name_chars"].reshape(B, N, 4 * 64), |
| ], |
| axis=-1, |
| ) |
| atom_feat = self.atom_feat_proj(atom_feat) |
|
|
| atom_coord = self.pos_embedder(pos=noised_pos) |
| atom_coord = self.atom_pos_proj(atom_coord) |
|
|
| atom_in = mx.concatenate([atom_feat, atom_coord], axis=-1) |
| atom_in = self.atom_in_proj(atom_in) |
|
|
| |
| atom_pe_pos = mx.concatenate( |
| [ |
| ref_space_uid[..., None].astype(mx.float32), |
| feats["ref_pos"], |
| ], |
| axis=-1, |
| ) |
|
|
| token_pe_pos = mx.concatenate( |
| [ |
| feats["residue_index"][..., None].astype(mx.float32), |
| feats["entity_id"][..., None].astype(mx.float32), |
| feats["asym_id"][..., None].astype(mx.float32), |
| feats["sym_id"][..., None].astype(mx.float32), |
| ], |
| axis=-1, |
| ) |
|
|
| atom_c_emb_enc = self.atom_enc_cond_proj(c_emb) |
| atom_latent = self.context2atom_proj(atom_in) |
| atom_latent = self.atom_encoder_transformer( |
| latents=atom_latent, |
| c=atom_c_emb_enc, |
| attention_mask=atom_attn_mask_enc, |
| pos=atom_pe_pos, |
| ) |
|
|
|
|
| atom_latent = self.atom2latent_proj(atom_latent) |
|
|
| |
| atom_to_token_mean = atom_to_token / ( |
| atom_to_token.sum(axis=1, keepdims=True) + 1e-6 |
| ) |
| latent = mx.matmul(atom_to_token_mean.swapaxes(axis1=1, axis2=2), atom_latent) |
| assert latent.shape[1] == M |
|
|
| esm_s = ( |
| mx.softmax(self.esm_s_combine, axis=0)[None, ...] @ feats["esm_s"] |
| ).squeeze(axis=2) |
|
|
| |
| esm_emb = self.esm_s_proj(esm_s, train=False) |
| assert esm_emb.shape[1] == latent.shape[1] |
|
|
| latent = self.esm_cat_proj(mx.concatenate([latent, esm_emb], axis=-1)) |
|
|
| |
| latent = self.trunk( |
| latents=latent, |
| c=c_emb, |
| attention_mask=None, |
| pos=token_pe_pos, |
| ) |
|
|
| |
| output = mx.matmul(atom_to_token, latent) |
| assert output.shape[1] == N |
|
|
| |
| output = output + atom_latent |
| output = self.latent2atom_proj(output) |
|
|
| |
| atom_c_emb_dec = self.atom_dec_cond_proj(c_emb) |
| output = self.atom_decoder_transformer( |
| latents=output, |
| c=atom_c_emb_dec, |
| attention_mask=atom_attn_mask_dec, |
| pos=atom_pe_pos, |
| ) |
|
|
| output = self.final_layer(output, c=c_emb) |
|
|
| return { |
| "predict_velocity": output, |
| "latent": latent, |
| } |
|
|