| import torch
|
| from torch import nn
|
| import torch.nn.functional as F
|
| from config import CFG
|
| import utils
|
| import math
|
| import numpy as np
|
| from cliplayers import QuickGELU, Transformer as MSTsfmEncoder
|
| from GNN import layers as gly
|
|
|
| loss_func_ms = nn.CrossEntropyLoss()
|
| loss_func = nn.CrossEntropyLoss()
|
|
|
|
|
| class MolGNNEncoder(nn.Module):
|
| def __init__(self,
|
| outdim,
|
| n_feats=74,
|
| n_filters_list=[256, 256, 256],
|
| n_head=4,
|
| mols=1,
|
| adj_chans=6,
|
| readout_layers=2,
|
| bias=True):
|
|
|
| super().__init__()
|
|
|
| n_filters_list = [i for i in n_filters_list if i is not None]
|
| lys = []
|
|
|
| for i, nf in enumerate(n_filters_list):
|
| if i == 0:
|
| nf1 = n_feats
|
| else:
|
| nf1 = prevnf
|
|
|
| prevnf = nf
|
|
|
| ly = gly.GConvBlockNoGF(nf1, nf, mols, adj_chans, bias)
|
| lys.append(ly)
|
|
|
| self.block_layers = nn.ModuleList(lys)
|
| self.attention_layer = gly.MultiHeadGlobalAttention(nf, n_head=n_head, concat=True, bias=bias)
|
| self.readout_layers = nn.ModuleList(
|
| [nn.Linear(nf * n_head, outdim, bias=bias)] + [nn.Linear(outdim, outdim) for _ in
|
| range(readout_layers - 1)])
|
| self.gelu = QuickGELU()
|
|
|
| def forward(self, batch):
|
| V = batch['V']
|
| A = batch['A']
|
| mol_size = batch['mol_size']
|
|
|
| for ly in self.block_layers:
|
| V = ly(V, A)
|
|
|
| X = self.attention_layer(V, mol_size)
|
|
|
| for ly in self.readout_layers:
|
| X = self.gelu(ly(X))
|
|
|
| return X
|
|
|
|
|
| class GATBlock(nn.Module):
|
| """GAT 卷积块: GraphAttentionLayer + BatchNorm + ELU + 残差连接"""
|
|
|
| def __init__(self, n_feats, n_filters, adj_chans=6, bias=True, dropout=0.1, alpha=0.2):
|
| super().__init__()
|
| self.gat_conv = gly.GraphAttentionLayer(n_feats, adj_chans, n_filters, bias, dropout, alpha)
|
| self.bn = nn.BatchNorm1d(n_filters)
|
|
|
| self.residual = nn.Linear(n_feats, n_filters, bias=False) if n_feats != n_filters else nn.Identity()
|
|
|
| def forward(self, V, A):
|
| V_res = self.residual(V)
|
| V_out = self.gat_conv(V, A)
|
| V_out = self.bn(V_out.transpose(1, 2).contiguous()).transpose(1, 2)
|
| V_out = F.elu(V_out + V_res)
|
| return V_out
|
|
|
|
|
| class MolGATEncoder(nn.Module):
|
| """
|
| 分子 GAT 编码器 — 与 MolGNNEncoder 同构, 将 GCN 替换为 GAT.
|
| 结构: GAT Block × N → MultiHead Global Attention → MLP Readout
|
| """
|
|
|
| def __init__(self,
|
| outdim,
|
| n_feats=74,
|
| n_filters_list=[256, 256, 256],
|
| n_head=4,
|
| mols=1,
|
| adj_chans=6,
|
| readout_layers=2,
|
| dropout=0.1,
|
| bias=True):
|
|
|
| super().__init__()
|
|
|
| n_filters_list = [i for i in n_filters_list if i is not None]
|
| lys = []
|
|
|
| for i, nf in enumerate(n_filters_list):
|
| if i == 0:
|
| nf1 = n_feats
|
| else:
|
| nf1 = prevnf
|
|
|
| prevnf = nf
|
|
|
| ly = GATBlock(nf1, nf, adj_chans, bias, dropout)
|
| lys.append(ly)
|
|
|
| self.block_layers = nn.ModuleList(lys)
|
| self.attention_layer = gly.MultiHeadGlobalAttention(nf, n_head=n_head, concat=True, bias=bias)
|
| self.readout_layers = nn.ModuleList(
|
| [nn.Linear(nf * n_head, outdim, bias=bias)] +
|
| [nn.Linear(outdim, outdim, bias=bias) for _ in range(readout_layers - 1)]
|
| )
|
| self.gelu = QuickGELU()
|
| self.dropout = nn.Dropout(dropout)
|
|
|
| def forward(self, batch):
|
| V = batch['V']
|
| A = batch['A']
|
| mol_size = batch['mol_size']
|
|
|
| for ly in self.block_layers:
|
| V = ly(V, A)
|
|
|
| X = self.attention_layer(V, mol_size)
|
|
|
| for ly in self.readout_layers:
|
| X = self.dropout(self.gelu(ly(X)))
|
|
|
| return X
|
|
|
|
|
| class ProjectionHead(nn.Module):
|
| def __init__(self,
|
| embedding_dim,
|
| projection_dim,
|
| cfg,
|
| transformer=True,
|
| lstm=False):
|
|
|
| super().__init__()
|
|
|
| self.projection = nn.Linear(embedding_dim, projection_dim)
|
| self.gelu = nn.GELU()
|
| self.transformer = None
|
| if transformer:
|
| self.transformer = MSTsfmEncoder(projection_dim, cfg.tsfm_layers, cfg.tsfm_heads)
|
| self.lstm = None
|
| if lstm:
|
| self.lstm = nn.LSTM(input_size=projection_dim, hidden_size=projection_dim, num_layers=cfg.lstm_layers,
|
| batch_first=True)
|
| self.dropout = nn.Dropout(cfg.dropout)
|
|
|
| def forward(self, x):
|
| projected = self.projection(x)
|
| if self.transformer is None:
|
| x = self.gelu(projected)
|
| else:
|
| x = self.transformer(projected)
|
| if not self.lstm is None:
|
| x, (_, _) = self.lstm(x)
|
| x = self.dropout(x)
|
|
|
| return x
|
|
|
|
|
|
|
| class FragSimiModel(nn.Module):
|
| def __init__(
|
| self,
|
| cfg
|
| ):
|
| super().__init__()
|
|
|
| self.cfg = cfg
|
| self.mol_gnn_encoder = None
|
| mol_embedding_dim = cfg.mol_embedding_dim
|
|
|
| if 'gnn' in self.cfg.mol_encoder:
|
| self.mol_gnn_encoder = MolGNNEncoder(outdim=cfg.mol_embedding_dim,
|
| n_filters_list=cfg.molgnn_n_filters_list,
|
| n_head=cfg.molgnn_nhead,
|
| readout_layers=cfg.molgnn_readout_layers)
|
| if 'fp' in self.cfg.mol_encoder:
|
| mol_embedding_dim = 2 * cfg.mol_embedding_dim
|
|
|
| if 'fm' in self.cfg.mol_encoder:
|
| mol_embedding_dim += 10
|
|
|
| self.ms_projection = ProjectionHead(cfg.ms_embedding_dim,
|
| cfg.projection_dim,
|
| cfg,
|
| cfg.tsfm_in_ms,
|
| cfg.lstm_in_ms)
|
|
|
| self.mol_projection = ProjectionHead(mol_embedding_dim,
|
| cfg.projection_dim,
|
| cfg,
|
| cfg.tsfm_in_mol,
|
| cfg.lstm_in_mol)
|
|
|
| def forward(self, batch):
|
| ms_features = batch["ms_bins"]
|
| mol_feat_list = []
|
| if 'gnn' in self.cfg.mol_encoder:
|
| mol_feat_list.append(self.mol_gnn_encoder(batch))
|
| if 'fp' in self.cfg.mol_encoder:
|
| mol_feat_list.append(batch["mol_fps"])
|
| if 'fm' in self.cfg.mol_encoder:
|
| mol_feat_list.append(batch["mol_fmvec"])
|
|
|
| if len(mol_feat_list) > 1:
|
| mol_features = torch.cat(mol_feat_list, dim=1)
|
| else:
|
| mol_features = mol_feat_list[0]
|
|
|
|
|
| ms_embeddings = self.ms_projection(ms_features)
|
| mol_embeddings = self.mol_projection(mol_features)
|
|
|
|
|
|
|
|
|
| logits = mol_embeddings @ ms_embeddings.t()
|
|
|
| ground_truth = torch.arange(ms_features.shape[0], dtype=torch.long, device=self.cfg.device)
|
|
|
| ms_loss = loss_func(logits, ground_truth)
|
| mol_loss = loss_func(logits.t(), ground_truth)
|
| loss = (ms_loss + mol_loss) / 2.0
|
|
|
| return loss.mean()
|
|
|
|
|
|
|
| class SinusoidalPositionEmbeddings(nn.Module):
|
| def __init__(self, dim):
|
| super().__init__()
|
| self.dim = dim
|
|
|
| def forward(self, time):
|
| device = time.device
|
| half_dim = self.dim // 2
|
| embeddings = math.log(10000) / (half_dim - 1)
|
| embeddings = torch.exp(torch.arange(half_dim, device=device) * -embeddings)
|
| embeddings = time[:, None] * embeddings[None, :]
|
| embeddings = torch.cat((embeddings.sin(), embeddings.cos()), dim=-1)
|
| return embeddings
|
|
|
|
|
|
|
| def batch_tanimoto_sim(x1, x2):
|
| '''
|
| 计算两个batch指纹之间的Tanimoto相似度矩阵
|
| x1, x2: shape (batch_size, fp_dim), assuming binary (0/1) floats
|
| '''
|
|
|
| dot_prod = torch.matmul(x1, x2.t())
|
| x1_sum = torch.sum(x1 ** 2, dim=1).view(-1, 1)
|
| x2_sum = torch.sum(x2 ** 2, dim=1).view(1, -1)
|
|
|
| denominator = x1_sum + x2_sum - dot_prod
|
|
|
| return dot_prod / (denominator + 1e-8)
|
|
|
|
|
|
|
| class HybridAlignLoss(nn.Module):
|
| def __init__(self, alpha=0.5, beta=1.0, temperature=0.07):
|
| super().__init__()
|
| self.alpha = alpha
|
| self.beta = beta
|
| self.temp = temperature
|
| self.cross_entropy = nn.CrossEntropyLoss()
|
| self.mse_loss = nn.MSELoss()
|
|
|
| def forward(self, ms_emb, mol_emb, mol_fps):
|
| """
|
| ms_emb: (batch, dim) normalized
|
| mol_emb: (batch, dim) normalized
|
| mol_fps: (batch, fp_dim) 原始指纹,用于计算Ground Truth相似度
|
| """
|
|
|
| logits = (ms_emb @ mol_emb.t()) / self.temp
|
| labels = torch.arange(ms_emb.shape[0], device=ms_emb.device)
|
| loss_i2t = self.cross_entropy(logits, labels)
|
| loss_t2i = self.cross_entropy(logits.t(), labels)
|
| loss_infonce = (loss_i2t + loss_t2i) / 2
|
|
|
|
|
|
|
|
|
| pred_sim_matrix = ms_emb @ mol_emb.t()
|
|
|
|
|
| with torch.no_grad():
|
| target_sim_matrix = batch_tanimoto_sim(mol_fps, mol_fps)
|
|
|
| loss_mse = self.mse_loss(pred_sim_matrix, target_sim_matrix)
|
|
|
|
|
| total_loss = self.beta * loss_infonce + self.alpha * loss_mse
|
|
|
| return total_loss, loss_infonce, loss_mse
|
|
|
|
|
| class MolFusionHead(nn.Module):
|
| def __init__(self, gnn_out_dim, fp_dim=1024, projection_dim=256, dropout=0.1):
|
| super().__init__()
|
|
|
|
|
|
|
| self.gnn_proj = nn.Linear(gnn_out_dim, 128)
|
|
|
|
|
| self.fp_proj = nn.Sequential(
|
| nn.Linear(fp_dim, 256),
|
| nn.LayerNorm(256),
|
| nn.GELU()
|
| )
|
|
|
|
|
| self.fusion_layer = nn.Sequential(
|
| nn.Linear(128 + 256, 256),
|
| nn.LayerNorm(256),
|
| nn.Dropout(dropout),
|
| nn.GELU()
|
| )
|
|
|
| def forward(self, gnn_feat, fps):
|
|
|
|
|
|
|
| x_graph = self.gnn_proj(gnn_feat)
|
| x_fp = self.fp_proj(fps)
|
|
|
|
|
| x_cat = torch.cat([x_graph, x_fp], dim=1)
|
|
|
|
|
| out = self.fusion_layer(x_cat)
|
| return out
|
|
|
|
|
| class SinusoidalMzEmbedding(nn.Module):
|
| def __init__(self, dim=256, max_mz=2000):
|
| super().__init__()
|
| self.dim = dim
|
| self.max_mz = max_mz
|
|
|
| div_term = torch.exp(torch.arange(0, dim, 2).float() * (-math.log(10000.0) / dim))
|
| self.register_buffer('div_term', div_term)
|
|
|
| def forward(self, mz_values):
|
| """
|
| mz_values: (Batch, Seq_Len) or (Seq_Len,)
|
| Returns: (..., 256)
|
| """
|
|
|
| pe = torch.zeros(*mz_values.shape, self.dim, device=mz_values.device)
|
| position = mz_values.unsqueeze(-1)
|
|
|
|
|
| pe[..., 0::2] = torch.sin(position * self.div_term)
|
| pe[..., 1::2] = torch.cos(position * self.div_term)
|
| return pe
|
|
|
|
|
| class FragSimiModelNew(nn.Module):
|
| def __init__(self, cfg):
|
| super().__init__()
|
| self.cfg = cfg
|
| self.experiment_name_type = cfg.experiment_name_type
|
|
|
|
|
|
|
| if 'gnn' in self.cfg.mol_encoder:
|
| self.mol_gnn_encoder = MolGNNEncoder(
|
| outdim=cfg.mol_embedding_dim,
|
| n_filters_list=cfg.molgnn_n_filters_list,
|
| n_head=cfg.molgnn_nhead,
|
| readout_layers=cfg.molgnn_readout_layers
|
| )
|
|
|
| self.mol_gat_encoder = MolGATEncoder(
|
| outdim=cfg.mol_embedding_dim,
|
| n_filters_list=cfg.molgnn_n_filters_list,
|
| n_head=cfg.molgnn_nhead,
|
| readout_layers=cfg.molgnn_readout_layers
|
| )
|
|
|
|
|
|
|
| fp_input_dim = getattr(cfg, 'fp_dim', 2048)
|
| self.mol_fusion = MolFusionHead(
|
| gnn_out_dim=cfg.mol_embedding_dim,
|
| fp_dim=fp_input_dim,
|
| projection_dim=256
|
| )
|
|
|
|
|
|
|
|
|
| ms_dim = 256
|
| self.ms_input_proj = nn.Linear(cfg.ms_embedding_dim, ms_dim)
|
| self.ms_transformer = MSTsfmEncoder(
|
| width=ms_dim,
|
| layers=6,
|
| heads=8
|
| )
|
|
|
| self.ms_input_proj3 = nn.Linear(cfg.ms_feature3_embedding_dim, ms_dim)
|
| self.ms_transformer3 = MSTsfmEncoder(
|
| width=ms_dim,
|
| layers=6,
|
| heads=8
|
| )
|
|
|
|
|
|
|
| self.ms_pos_embed = SinusoidalPositionEmbeddings(256)
|
|
|
|
|
| self.mz_embedder = SinusoidalMzEmbedding(dim=256)
|
|
|
|
|
| self.feature_proj = nn.Linear(29, 256)
|
|
|
|
|
|
|
|
|
|
|
|
|
| self.transformer = MSTsfmEncoder(
|
| width=256,
|
| layers=6,
|
| heads=8
|
| )
|
|
|
|
|
| self.final_proj = nn.Sequential(
|
| nn.Linear(256, 256),
|
| nn.LayerNorm(256),
|
| nn.Dropout(cfg.dropout)
|
| )
|
|
|
|
|
| self.ms_final_proj = nn.Sequential(
|
| nn.Linear(256, 256),
|
| nn.LayerNorm(256),
|
| nn.Dropout(cfg.dropout)
|
| )
|
|
|
| self.ms_final_proj3 = nn.Sequential(
|
| nn.Linear(256, 256),
|
| nn.LayerNorm(256),
|
| nn.Dropout(cfg.dropout)
|
| )
|
|
|
|
|
| self.all_final_proj1 = nn.Sequential(
|
| nn.Linear(256, 256),
|
| nn.LayerNorm(256),
|
| nn.Dropout(cfg.dropout)
|
| )
|
|
|
| self.all_final_proj = nn.Sequential(
|
| nn.Linear(512, 256),
|
| nn.LayerNorm(256),
|
| nn.Dropout(cfg.dropout)
|
| )
|
|
|
| self.all_final_proj3 = nn.Sequential(
|
| nn.Linear(256 * 3, 256),
|
| nn.LayerNorm(256),
|
| nn.Dropout(cfg.dropout)
|
| )
|
|
|
|
|
|
|
| alpha = self.cfg.alpha
|
| beta = self.cfg.beta
|
| self.loss_fn = HybridAlignLoss(alpha=alpha, beta=beta)
|
|
|
|
|
| self.logit_scale = nn.Parameter(torch.ones([]) * np.log(1 / 0.07))
|
|
|
| def forward(self, batch, is_predict=False):
|
|
|
|
|
|
|
| if "gcn" in self.experiment_name_type:
|
| mol_gnn_feat = self.mol_gnn_encoder(batch)
|
| elif 'gat' in self.experiment_name_type:
|
| mol_gnn_feat = self.mol_gat_encoder(batch)
|
|
|
| mol_fps = batch["mol_fps"]
|
|
|
|
|
| mol_embeddings = self.mol_fusion(mol_gnn_feat, mol_fps)
|
|
|
|
|
|
|
|
|
| ms_features = batch["ms_bins"]
|
|
|
|
|
| x_ms = self.ms_input_proj(ms_features)
|
|
|
|
|
| if len(x_ms.shape) == 3:
|
| pos_emb = self.ms_pos_embed(torch.arange(x_ms.shape[1], device=x_ms.device))
|
| x_ms = x_ms + pos_emb
|
| else:
|
|
|
| x_ms = x_ms.unsqueeze(1)
|
|
|
| x_ms = self.ms_transformer(x_ms)
|
|
|
|
|
| if len(x_ms.shape) == 3:
|
| x_ms = x_ms.mean(dim=1)
|
|
|
| ms_embeddings = self.ms_final_proj(x_ms)
|
|
|
| if 'ms3' in self.experiment_name_type:
|
|
|
| if 'ms_bins3' in batch:
|
| ms_features3 = batch["ms_bins3"].to(self.ms_input_proj3.weight.dtype)
|
| x_ms3 = self.ms_input_proj3(ms_features3)
|
| x_ms3 = self.ms_transformer3(x_ms3)
|
|
|
|
|
| if len(x_ms3.shape) == 3:
|
| x_ms3 = x_ms3.mean(dim=1)
|
|
|
| ms_embeddings3 = self.ms_final_proj3(x_ms3)
|
|
|
| if 'ms1' in self.experiment_name_type:
|
|
|
| ms_features, ms_mz = batch["ms_bins1"], batch["ms_bins2"]
|
|
|
| mz_emb = self.mz_embedder(ms_mz)
|
|
|
|
|
| feat_emb = self.feature_proj(ms_features)
|
|
|
|
|
| x = mz_emb + feat_emb
|
|
|
|
|
|
|
|
|
| x = self.transformer(x)
|
|
|
|
|
|
|
|
|
| x = x.mean(dim=1)
|
|
|
|
|
| out = self.final_proj(x)
|
|
|
| if 'ms3' in self.experiment_name_type and 'ms1' in self.experiment_name_type:
|
| out1 = torch.cat([ms_embeddings, ms_embeddings3, out], dim=-1)
|
| ms_embeddings = self.all_final_proj3(out1)
|
|
|
| elif 'ms3' in self.experiment_name_type:
|
| out1 = torch.cat([ms_embeddings, ms_embeddings3], dim=-1)
|
| ms_embeddings = self.all_final_proj(out1)
|
| elif 'ms1' in self.experiment_name_type:
|
| out1 = torch.cat([ms_embeddings, out], dim=-1)
|
| ms_embeddings = self.all_final_proj(out1)
|
| else:
|
| ms_embeddings = self.all_final_proj1(ms_embeddings)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| ms_embeddings = F.normalize(ms_embeddings, dim=-1, p=2)
|
| mol_embeddings = F.normalize(mol_embeddings, dim=-1, p=2)
|
|
|
|
|
| loss, loss_infonce, loss_mse = self.loss_fn(ms_embeddings, mol_embeddings, mol_fps)
|
|
|
| if 'loss' not in self.experiment_name_type:
|
| loss = loss_infonce
|
|
|
| if is_predict:
|
| if 'loss' not in self.experiment_name_type:
|
| loss = loss_infonce
|
|
|
| return loss, loss_infonce, loss_mse, ms_embeddings, mol_embeddings
|
|
|
| return loss
|
|
|