Volver a la lista

Predicción de interacciones fármaco-diana mediante redes neuronales gráficas: aplicación práctica del cribado para el reposicionamiento de fármacos dirigidos a la Mpro del SARS-CoV-2.

Convierte las cadenas SMILES en gráficos moleculares mediante RDKit y entrena y predice las interacciones fármaco-diana (DTI) utilizando una red neuronal de grafos (GAT) de PyTorch Geometric. Incluye un escenario de reutilización de fármacos dirigidos a la Mpro del SARS-CoV-2, la reproducción de los conjuntos de referencia BindingDB, Davis y KIBA, la validación de la generalización en entornos de diana/fármaco desconocidos (cold-target/cold-drug) y un flujo de trabajo completo que abarca hasta la visualización de los átomos de unión mediante la atención GAT.

Intermedio
|
40min
|
Verificado (2026-07)
Progreso0/15 (0%)

Predicción de interacciones fármaco-diana mediante redes neuronales gráficas: caso práctico de cribado para el reposicionamiento de fármacos dirigidos a la Mpro del SARS-CoV-2

Predecir si un fármaco candidato de molécula pequeña se une a una proteína diana y la intensidad de esa unión es fundamental para el descubrimiento de nuevos fármacos. Mientras que el acoplamiento molecular (capítulo 08) se basa en coordenadas 3D, este capítulo se centra en cómo manejar las estructuras químicas como grafos. Las cadenas SMILES se representan como un grafo compuesto por nodos (átomos) y aristas (enlaces), y una red neuronal gráfica (GNN) aprende esta información topológica para predecir la unión. Es posible analizar decenas de miles de fármacos candidatos en segundos con una sola GPU pequeña, e incluso interpretar qué átomos contribuyen a la unión mediante la visualización de la atención. En este capítulo, se construye un flujo de trabajo completo utilizando como ejemplo el escenario de reposicionamiento de fármacos dirigido a la proteasa principal (Mpro) del SARS-CoV-2.

📚 Lecturas previas recomendadas (muy recomendable)

Este capítulo es una guía avanzada y técnica sobre IA y biotecnología. Antes de comenzar, se recomienda encarecidamente revisar primero los siguientes capítulos de DryBench:

Si se accede a este capítulo sin haber leído las lecturas previas, se comenzará directamente con el código práctico sin volver a explicar los principios de la convolución gráfica, el paso de mensajes, las cabeceras de atención y el autograd de PyTorch, por lo que será difícil seguir el ritmo.


Lo que ya aprendimos en DryBench

En DryBench ai-native #2, aprendimos que las redes neuronales toman patrones locales de la entrada como filtros y apilan capas para aprender representaciones jerárquicas; en #12, aprendimos que PyTorch proporciona los fundamentos de tensores, autograd y la transferencia a la GPU; y en #13, aprendimos que HuggingFace ofrece una API estándar para cargar modelos preentrenados.

Las redes neuronales gráficas (GNN) extienden estos principios a estructuras de grafos irregulares. Mientras que las imágenes son píxeles en una cuadrícula y el texto es una secuencia lineal, lo que hace que las CNN y las RNN sean naturales, las moléculas son grafos en los que cada nodo tiene un número diferente de vecinos. El paso de mensajes de la GNN consiste en que cada nodo intercambia información con sus vecinos para actualizar su propia representación, lo que lo convierte en el estándar para datos de tipo grafo, como moléculas, proteínas y redes sociales. En este capítulo, aplicamos estos principios de forma intensiva a la predicción de interacciones fármaco-diana (DTI).

Definición del problema

Escenario práctico: Reutilización de fármacos dirigidos a la Mpro de SARS-CoV-2

Durante el período pandémico de 2020 a 2023, varios equipos de investigación realizaron cribados de candidatos que se unen a la proteasa principal (Mpro, 3CLpro) de SARS-CoV-2 a partir de bibliotecas de fármacos ya aprobados (FDA · DrugBank · Repurposing Hub). Los experimentos reales cuestan decenas de miles de wones y requieren varios días por fármaco, por lo que un filtro in silico previo es esencial. El objetivo de nuestra canalización es:

  • Entrada: 3000 SMILES de fármacos aprobados + secuencia de Mpro (UniProt P0DTD1, 306 residuos) + estructura PDB (por ejemplo, 6LU7 · 7BQY).
  • Salida: Puntuaciones de probabilidad e intensidad de unión a Mpro para cada fármaco, con los 20 mejores recomendados.
  • Validación: Coherencia con ligandos de Mpro medidos experimentalmente en PDBbind y BindingDB, así como con inhibidores de proteasa conocidos (como nirmatrelvir).
  • Coste: Completado en menos de 30 minutos en una GPU pequeña local.

Desafíos fundamentales del DTI

  • Tarea de clasificación: Predicción binaria de unión frente a no unión (BindingDB · DrugBank · BioSNAP).
  • Tarea de regresión: Cuantificación de la intensidad de unión mediante pIC50, Kd, Ki, etc. (Davis · KIBA · benchmark Metz).
  • Generalización cold-target: Capacidad de predicción para nuevas proteínas diana no presentes en el entrenamiento (Mpro no aparece en los datos previos a la pandemia de SARS-CoV-2).
  • Generalización cold-drug: Capacidad de predicción para nuevos candidatos a fármacos no presentes en el entrenamiento.
  • Diversidad de ensayos: Los datos de entrenamiento mezclan varios ensayos (bioquímicos, celulares, in vivo). Existen sesgos específicos por ensayo.

Espectro de enfoques previos y posición de este capítulo

  • Random Forest + características ECFP: Rápido y como referencia. R² 0,4~0,5.
  • CNN 1D + caracteres SMILES (DeepDTA 2018): Procesamiento como secuencia, con pérdida de información estructural. R² 0,5~0,6 [1].
  • GraphDTA (Nguyen 2020): Primer modelo de última generación (SOTA) basado en grafos para DTI. R² 0,6~0,65 [2].
  • MGraphDTA (Yang 2022): Grafos multi-escala y atención. R² 0,65~0,70 [3].
  • GeNNius (2024): Ultra ligero y ultra rápido. Entrenamiento e inferencia muy rápidos [4].
  • DrugBAN (2023): Atención bilineal y mejora de la interpretabilidad.
  • Modelo fundacional (ChemBERTa · MolFormer · DeepChem): Pre-entrenamiento a gran escala seguido de ajuste fino.
  • AlphaFold3 · Boltz-2 (Parte 11): Integración de acoplamiento. Precisión de primer nivel, pero con alta carga en la GPU.

En esta parte se aborda el modelo base basado en GAT (parte principal), la integración de incrustaciones de ChemBERTa (idea de extensión) y la visualización de la atención.

Indicadores de objetivo para esta parte

  • AUC de clasificación binaria de BindingDB superior a 0,90 (división aleatoria).
  • Coeficiente de correlación de Pearson en Davis superior a 0,80.
  • AUC del escenario de objetivo frío 2 superior a 0,75 (verificación de generalización).
  • Recall de los 20 principales inhibidores conocidos en el escenario Mpro superior a 0,5 (es decir, recuperar al menos 10 inhibidores conocidos entre los 20 primeros).
  • Tiempo de entrenamiento inferior a 30 minutos (GPU pequeña).
  • Rendimiento de inferencia superior a 1000 SMILES por segundo.

Conjunto de herramientas y requisitos de infraestructura

HerramientaFunciónLicencia
RDKitSMILES → grafo molecular · huella digitalBSD-3-Clause
PyTorch Geometric (PyG)Framework de redes neuronales gráficasMIT
DGL (opcional)Alternativa a PyGApache 2.0
GeNNius (GitHub)Modelo base GNN ultraligeroMIT (previsto)
ChemBERTa (HuggingFace)Incrustaciones moleculares preentrenadas (opcional)Apache 2.0
MolFormer (IBM)Alternativa de modelo base molecularApache 2.0
BindingDB · Davis · KIBA · BioSNAPConjuntos de datos de entrenamiento y validaciónGratuito para fines académicos
DrugBank · DrugRepurposing HubBibliotecas de fármacosGratuito para fines académicos · Registro

Requisitos de infraestructura:

  • GPU de consumo pequeña (se recomienda RTX 4060 con al menos 8 GB; la CPU también puede realizar el entrenamiento y la inferencia, pero es 5 a 10 veces más lenta).
  • RAM superior a 16 GB.
  • Disco: volcado de BindingDB de aproximadamente 500 MB, Davis y KIBA de menos de 100 MB cada uno, DrugBank de aproximadamente 200 MB.

Costo estimado para la reproducción por parte del estudiante: Costo de API 0 (completamente local). Tiempo de GPU inferior a 30 minutos para el entrenamiento, inferencia en segundos.

Implementación práctica del flujo de trabajo

Flujo completo:

mermaid

Paso 1. SMILES → gráfico molecular

python
from dataclasses import dataclass
from typing import Any
import torch
from torch_geometric.data import Data
from rdkit import Chem, RDLogger
RDLogger.DisableLog("rdApp.*")
ATOM_FEATURES = {
"atomic_num": list(range(1, 119)),
"degree": [0, 1, 2, 3, 4, 5, 6],
"formal_charge": [-3, -2, -1, 0, 1, 2, 3],
"hybridization": [
Chem.rdchem.HybridizationType.SP,
Chem.rdchem.HybridizationType.SP2,
Chem.rdchem.HybridizationType.SP3,
Chem.rdchem.HybridizationType.SP3D,
Chem.rdchem.HybridizationType.SP3D2,
],
"num_h": [0, 1, 2, 3, 4],
"is_aromatic": [False, True],
"is_in_ring": [False, True],
}
def one_hot(value: Any, choices: list) -> list[int]:
"""Si no coincide, establece 1 en la última ranura (tratamiento OOV)."""
if value in choices:
idx = choices.index(value)
else:
idx = len(choices) - 1
result = [0] * len(choices)
result[idx] = 1
return result
def atom_features(atom: Chem.Atom) -> list[float]:
features = []
features += one_hot(atom.GetAtomicNum(), ATOM_FEATURES["atomic_num"])
features += one_hot(atom.GetDegree(), ATOM_FEATURES["degree"])
features += one_hot(atom.GetFormalCharge(), ATOM_FEATURES["formal_charge"])
features += one_hot(atom.GetHybridization(), ATOM_FEATURES["hybridization"])
features += one_hot(atom.GetTotalNumHs(), ATOM_FEATURES["num_h"])
features += one_hot(atom.GetIsAromatic(), ATOM_FEATURES["is_aromatic"])
features += one_hot(atom.IsInRing(), ATOM_FEATURES["is_in_ring"])
return features
def bond_features(bond: Chem.Bond) -> list[float]:
bt = bond.GetBondType()
return [
int(bt == Chem.rdchem.BondType.SINGLE),
int(bt == Chem.rdchem.BondType.DOUBLE),
int(bt == Chem.rdchem.BondType.TRIPLE),
int(bt == Chem.rdchem.BondType.AROMATIC),
int(bond.GetIsConjugated()),
int(bond.IsInRing()),
]
def smiles_to_graph(smiles: str, drug_id: str = "") -> Data | None:
"""SMILES a Datos PyG. Devuelve None si falla el análisis."""
mol = Chem.MolFromSmiles(smiles)
if mol is None or mol.GetNumAtoms() == 0:
return None
canonical = Chem.MolToSmiles(mol, canonical=True)
node_features = torch.tensor(
[atom_features(atom) for atom in mol.GetAtoms()],
dtype=torch.float,
)
edge_indices = []
edge_attrs = []
for bond in mol.GetBonds():
i, j = bond.GetBeginAtomIdx(), bond.GetEndAtomIdx()
feat = bond_features(bond)
edge_indices += [[i, j], [j, i]] # bidireccional
edge_attrs += [feat, feat]
if not edge_indices:
edge_index = torch.empty((2, 0), dtype=torch.long)
edge_attr = torch.empty((0, 6), dtype=torch.float)
else:
edge_index = torch.tensor(edge_indices, dtype=torch.long).t().contiguous()
edge_attr = torch.tensor(edge_attrs, dtype=torch.float)
data = Data(x=node_features, edge_index=edge_index, edge_attr=edge_attr)
data.smiles = canonical
data.drug_id = drug_id
return data

Paso 2. Codificación de la proteína objetivo (incrustaciones ESM2 o modelo base CNN)

python
def protein_to_onehot(sequence: str, max_length: int = 1200) -> torch.Tensor:
"""Secuencia de proteína a one-hot (max_length, 21)."""
aa_to_idx = {aa: i for i, aa in enumerate("ACDEFGHIKLMNPQRSTVWY")}
seq = sequence[:max_length].upper()
encoded = torch.zeros(max_length, 21)
for i, aa in enumerate(seq):
encoded[i, aa_to_idx.get(aa, 20)] = 1.0
return encoded
class ProteinCNN(torch.nn.Module):
"""Secuencia de proteína → vector de incrustación (línea base)."""
def __init__(self, output_dim: int = 128):
super().__init__()
self.conv1 = torch.nn.Conv1d(21, 32, kernel_size=8, padding=3)
self.conv2 = torch.nn.Conv1d(32, 64, kernel_size=8, padding=3)
self.conv3 = torch.nn.Conv1d(64, output_dim, kernel_size=8, padding=3)
self.pool = torch.nn.AdaptiveAvgPool1d(1)
def forward(self, x: torch.Tensor) -> torch.Tensor:
# x: (B, max_length, 21)
x = x.transpose(1, 2) # (B, 21, L)
x = torch.relu(self.conv1(x))
x = torch.relu(self.conv2(x))
x = torch.relu(self.conv3(x))
return self.pool(x).squeeze(-1)
class ProteinESM2(torch.nn.Module):
"""Envoltorio de incrustación ESM2 (similar al capítulo 02)."""
MODEL_ID = "facebook/esm2_t6_8M_UR50D" # 8M pequeño · muy ligero
def __init__(self, device: str = "cuda"):
super().__init__()
from transformers import AutoTokenizer, AutoModel
self.tokenizer = AutoTokenizer.from_pretrained(self.MODEL_ID)
self.model = AutoModel.from_pretrained(self.MODEL_ID).to(device).eval()
self.device = device
self.output_dim = self.model.config.hidden_size # 320 for 8M
@torch.no_grad()
def forward(self, sequences: list[str]) -> torch.Tensor:
embs = []
for seq in sequences:
inputs = self.tokenizer(seq, return_tensors="pt", truncation=True, max_length=1024).to(self.device)
out = self.model(**inputs)
embs.append(out.last_hidden_state.mean(dim=1).squeeze(0))
return torch.stack(embs).to(self.device)

Paso 3. Modelo de interacción fármaco-objetivo (DTI) basado en GAT.

python
from torch_geometric.nn import GATConv, global_mean_pool
class GAT_DTI(torch.nn.Module):
"""GAT molecular + CNN/ESM de proteína + MLP concat → predicción DTI."""
def __init__(
self,
atom_feat_dim: int = 133, # dimensión de salida de atom_features
gat_hidden: int = 128,
gat_heads: int = 4,
protein_dim: int = 128,
fusion_dim: int = 256,
task: str = "classification", # "regression" or "classification"
):
super().__init__()
self.gat1 = GATConv(atom_feat_dim, gat_hidden, heads=gat_heads, dropout=0.1)
self.gat2 = GATConv(gat_hidden * gat_heads, gat_hidden, heads=1, dropout=0.1)
self.protein_encoder = ProteinCNN(output_dim=protein_dim)
self.fusion = torch.nn.Sequential(
torch.nn.Linear(gat_hidden + protein_dim, fusion_dim),
torch.nn.ReLU(),
torch.nn.Dropout(0.2),
torch.nn.Linear(fusion_dim, fusion_dim // 2),
torch.nn.ReLU(),
)
self.head = torch.nn.Linear(fusion_dim // 2, 1)
self.task = task
def forward(self, mol_data: Data, protein_onehot: torch.Tensor, return_attention: bool = False) -> torch.Tensor:
x, edge_index = mol_data.x, mol_data.edge_index
if return_attention:
x, (_, alpha1) = self.gat1(x, edge_index, return_attention_weights=True)
else:
x = self.gat1(x, edge_index)
x = torch.relu(x)
x = self.gat2(x, edge_index)
mol_emb = global_mean_pool(x, mol_data.batch) # (B, gat_hidden)
prot_emb = self.protein_encoder(protein_onehot) # (B, protein_dim)
combined = torch.cat([mol_emb, prot_emb], dim=-1)
hidden = self.fusion(combined)
out = self.head(hidden).squeeze(-1)
if self.task == "classification":
out = torch.sigmoid(out)
if return_attention:
return out, alpha1
return out

Paso 4. DataLoader · Bucle de entrenamiento

python
from torch_geometric.loader import DataLoader as PyGLoader
class DTIDataset(torch.utils.data.Dataset):
"""Conjunto de datos de entrenamiento DTI."""
def __init__(self, records: list[dict]):
self.records = []
for r in records:
g = smiles_to_graph(r["smiles"], drug_id=r.get("drug_id", ""))
if g is None:
continue
g.protein = protein_to_onehot(r["protein_sequence"])
g.y = torch.tensor(r["label"], dtype=torch.float)
g.target_id = r.get("target_id", "")
self.records.append(g)
def __len__(self):
return len(self.records)
def __getitem__(self, idx):
return self.records[idx]
def train_dti(
train_records: list[dict],
val_records: list[dict],
epochs: int = 50,
batch_size: int = 64,
lr: float = 1e-3,
device: str = "cuda",
task: str = "classification",
) -> torch.nn.Module:
train_ds = DTIDataset(train_records)
val_ds = DTIDataset(val_records)
train_loader = PyGLoader(train_ds, batch_size=batch_size, shuffle=True)
val_loader = PyGLoader(val_ds, batch_size=batch_size)
model = GAT_DTI(task=task).to(device)
optimizer = torch.optim.Adam(model.parameters(), lr=lr, weight_decay=1e-5)
criterion = torch.nn.BCELoss() if task == "classification" else torch.nn.MSELoss()
best_val_metric = 0.0 if task == "classification" else float("inf")
for epoch in range(epochs):
# Train
model.train()
total_loss = 0.0
for batch in train_loader:
batch = batch.to(device)
optimizer.zero_grad()
pred = model(batch, batch.protein.view(batch.num_graphs, -1, 21))
loss = criterion(pred, batch.y)
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
total_loss += loss.item()
# Validate
val_metric = evaluate_model(model, val_loader, device, task)
print(f"Epoch {epoch+1}: train_loss={total_loss/len(train_loader):.4f}, val_metric={val_metric:.4f}")
# Guardar el mejor
if (task == "classification" and val_metric > best_val_metric) or \
(task == "regression" and val_metric < best_val_metric):
best_val_metric = val_metric
torch.save(model.state_dict(), "best_dti.pt")
return model
def evaluate_model(model, loader, device: str, task: str) -> float:
from sklearn.metrics import roc_auc_score, mean_squared_error
model.eval()
all_preds, all_labels = [], []
with torch.no_grad():
for batch in loader:
batch = batch.to(device)
pred = model(batch, batch.protein.view(batch.num_graphs, -1, 21))
all_preds.extend(pred.cpu().tolist())
all_labels.extend(batch.y.cpu().tolist())
if task == "classification":
return roc_auc_score(all_labels, all_preds)
return mean_squared_error(all_labels, all_preds, squared=False)

Paso 5. División de objetivos fríos y fármacos fríos (validación generalizada)

El verdadero desafío de los bancos de datos de interacciones fármaco-objetivo (DTI). La división "fría" se asemeja más a los escenarios reales de descubrimiento de fármacos que la división aleatoria [5].

python
import random
def cold_target_split(
records: list[dict], val_target_ratio: float = 0.2, seed: int = 42,
) -> tuple[list, list]:
"""División cold-target: los objetivos de validación no aparecen en el entrenamiento."""
random.seed(seed)
all_targets = list({r["target_id"] for r in records})
random.shuffle(all_targets)
n_val = int(len(all_targets) * val_target_ratio)
val_targets = set(all_targets[:n_val])
train = [r for r in records if r["target_id"] not in val_targets]
val = [r for r in records if r["target_id"] in val_targets]
return train, val
def cold_drug_split(
records: list[dict], val_drug_ratio: float = 0.2, seed: int = 42,
) -> tuple[list, list]:
"""División cold-drug: los fármacos de validación no aparecen en el entrenamiento."""
random.seed(seed)
all_drugs = list({r["drug_id"] for r in records})
random.shuffle(all_drugs)
n_val = int(len(all_drugs) * val_drug_ratio)
val_drugs = set(all_drugs[:n_val])
train = [r for r in records if r["drug_id"] not in val_drugs]
val = [r for r in records if r["drug_id"] in val_drugs]
return train, val

Paso 6. Cribado para el reposicionamiento de fármacos contra la Mpro del SARS-CoV-2

python
import pandas as pd
MPRO_SEQUENCE = ( # SARS-CoV-2 Mpro (UniProt P0DTD1, 306 residues)
"SGFRKMAFPSGKVEGCMVQVTCGTTTLNGLWLDDVVYCPRHVICTSEDMLNPNYEDLLIRKSNHNFLVQAGNVQLRVIGHSMQNCVLKLKVDTANPKTPKYKFVR"
"IQPGQTFSVLACYNGSPSGVYQCAMRPNFTIKGSFLNGSCGSVGFNIDYDCVSFCYMHHMELPTGVHAGTDLEGNFYGPFVDRQTAQAAGTDTTITVNVLAWLYA"
"AVINGDRWFLNRFTTTLNDFNLVAMKYNYEPLTQDHVDILGPLSAQTGIAVLDMCASLKELLQNGMNGRTILGSALLEDEFTPFDVVRQCSGVTFQ"
)
KNOWN_MPRO_INHIBITORS = [
# nirmatrelvir (componente de Paxlovid)
"CC1(C)C2CC1C(NC(=O)C1CCCN1C(=O)C(NC(=O)OC(C)(C)C)C(C)(C)C)C(=O)NC(C#N)CC2=O",
# ensitrelvir (Xocova) y otros ejemplos
# En producción, obtener de ChEMBL · DrugBank, etc.
]
def screen_mpro_repurposing(
model: torch.nn.Module,
drug_smiles_list: list[tuple[str, str]], # (drug_id, smiles)
protein_sequence: str = MPRO_SEQUENCE,
top_k: int = 20,
device: str = "cuda",
) -> pd.DataFrame:
"""Devuelve los K mejores candidatos de unión a Mpro desde la biblioteca de fármacos aprobados."""
model.eval()
prot_encoded = protein_to_onehot(protein_sequence).unsqueeze(0).to(device)
results = []
for drug_id, smiles in drug_smiles_list:
g = smiles_to_graph(smiles, drug_id)
if g is None:
continue
g.batch = torch.zeros(g.x.size(0), dtype=torch.long)
g = g.to(device)
with torch.no_grad():
pred = model(g, prot_encoded).item()
results.append({
"drug_id": drug_id,
"smiles": smiles,
"mpro_score": pred,
})
df = pd.DataFrame(results).sort_values("mpro_score", ascending=False)
return df.head(top_k)
def validate_against_known_inhibitors(
model: torch.nn.Module,
all_screened: pd.DataFrame,
known_inhibitors_smiles: list[str],
top_k: int = 20,
) -> dict:
"""¿Cuántos inhibidores conocidos se incluyen en el top K?"""
known_set = set(known_inhibitors_smiles)
top_smiles = set(all_screened.head(top_k)["smiles"].tolist())
recall = len(top_smiles & known_set) / max(len(known_set), 1)
return {
"top_k": top_k,
"recall_of_known": recall,
"known_in_top_k": list(top_smiles & known_set),
}

Paso 7. Visualización de la atención (interpretabilidad)

Visualizar qué átomos contribuyen a la predicción mediante los pesos de atención del GAT.

python
def visualize_attention_on_molecule(
model: GAT_DTI,
smiles: str,
protein_sequence: str,
output_path: str = "attention.png",
device: str = "cuda",
) -> None:
"""Superposición de color por átomo con atención GAT."""
from rdkit.Chem import Draw
from rdkit.Chem.Draw import rdMolDraw2D
from matplotlib.cm import get_cmap
g = smiles_to_graph(smiles)
g.batch = torch.zeros(g.x.size(0), dtype=torch.long)
g = g.to(device)
prot = protein_to_onehot(protein_sequence).unsqueeze(0).to(device)
model.eval()
with torch.no_grad():
pred, alpha = model(g, prot, return_attention=True)
# alpha: (num_edges, num_heads). Agregado de atención por átomo
edge_index = g.edge_index.cpu().numpy()
alpha_np = alpha.mean(dim=1).cpu().numpy() # (num_edges,)
n_atoms = g.x.size(0)
atom_attn = np.zeros(n_atoms)
for e_idx, dest in enumerate(edge_index[1]):
atom_attn[dest] += float(alpha_np[e_idx])
atom_attn = atom_attn / max(atom_attn.max(), 1e-8)
# Resaltado de color con RDKit
mol = Chem.MolFromSmiles(smiles)
cmap = get_cmap("Reds")
highlight_colors = {i: cmap(float(atom_attn[i]))[:3] for i in range(n_atoms)}
drawer = rdMolDraw2D.MolDraw2DCairo(500, 500)
drawer.drawOptions().addAtomIndices = False
rdMolDraw2D.PrepareAndDrawMolecule(
drawer, mol,
highlightAtoms=list(range(n_atoms)),
highlightAtomColors=highlight_colors,
)
drawer.FinishDrawing()
with open(output_path, "wb") as f:
f.write(drawer.GetDrawingText())
print(f"Visualización de atención guardada en: {output_path} (puntuación predicha={pred.item():.3f})")

Flujo de trabajo integrado

python
def full_pipeline(
training_records: list[dict],
mpro_repurposing_library: list[tuple[str, str]],
output_dir: Path,
device: str = "cuda",
) -> None:
"""Entrenamiento del modelo DTI → cribado de Mpro → visualización de atención."""
output_dir.mkdir(parents=True, exist_ok=True)
print("[1/4] Cold-target split")
train, val = cold_target_split(training_records, val_target_ratio=0.2)
print(f" train {len(train)} / val {len(val)}")
print("[2/4] Entrenamiento GAT_DTI")
model = train_dti(train, val, epochs=50, task="classification", device=device)
print("[3/4] Cribado de reposicionamiento de Mpro")
top20 = screen_mpro_repurposing(model, mpro_repurposing_library, top_k=20, device=device)
top20.to_csv(output_dir / "mpro_top20.csv", index=False)
print("[4/4] Visualización de atención de los 3 mejores")
for _, row in top20.head(3).iterrows():
visualize_attention_on_molecule(
model, row["smiles"], MPRO_SEQUENCE,
output_path=str(output_dir / f"attn_{row['drug_id']}.png"),
device=device,
)
# Validación de reproducción de inhibidores conocidos
validation = validate_against_known_inhibitors(
model, top20, KNOWN_MPRO_INHIBITORS,
)
print(f"Recall de inhibidores conocidos@20: {validation['recall_of_known']:.2f}")

Rendimiento, costo y casos de fallo conocidos

Referencias de rendimiento (citación de evaluaciones comparativas públicas)

ModeloEvaluación comparativa (Davis · KIBA · BindingDB)AUC / Pearson rFuente
Random Forest + ECFPRegresión Davisr = 0.55Legacy
DeepDTA (CNN 1D)Davisr = 0.62Öztürk et al., Bioinformatics 2018 [1]
GraphDTA (GAT)Davisr = 0.68Nguyen et al., Bioinformatics 2020 [2]
MGraphDTADavis · KIBAr = 0.72Yang et al., Chem Sci 2022 [3]
GeNNiusBindingDBAUC 0.92ML4BM Lab 2024 [4]
DrugBANBioSNAPAUC 0.89Bai et al., Nat Mach Intell 2023
Configuración de "cold-target"Davis cold-targetr = 0.35~0.50 (disminución notable)Pahikkala 2015 [5]
Evaluación comparativa SARS-CoV-2 Mpro (COVID Moonshot)Biblioteca de candidatosDiversas evaluaciones comparativas de la literaturaMoonshot consortium [6]

Costo estimado de reproducción para el estudiante

  • Costo de la API: 0 (totalmente local).
  • Entrenamiento en Davis con una GPU de consumo pequeño (~30 000 pares): aprox. 30~60 minutos.
  • Inferencia: más de 1000 SMILES por segundo.
  • Cribado de Mpro para 3000 fármacos: completado en segundos.

5 casos de fallo conocidos (recopilación comunitaria y bibliográfica)

  1. Disminución notable del rendimiento en la configuración de "cold-target" (fallo de generalización) Síntoma: AUC de 0.90 con una división aleatoria, pero cae a 0.55 en la configuración de "cold-target". Causa: Sobreajuste a las dianas comunes en los datos de entrenamiento (p. ej., quinasas). Las nuevas dianas tienen representaciones de secuencia diferentes. Prevención: (a) Mejorar el codificador de dianas con ESM2/ESM3 preentrenado, (b) Regularizar el espacio de incrustación de la diana mediante el aprendizaje contrastivo, (c) Forzar la evaluación comparativa con una división de "cold-target" (Paso 5 de este capítulo), (d) Utilizar valores iniciales de incrustación molecular de MolFormer y ChemBERTa. Fuente: Pahikkala et al. "Toward more realistic drug-target interaction predictions." Brief Bioinform 2015 [5].

  2. Ausencia de la canónica SMILES que provoca un aprendizaje redundante Síntoma: Múltiples representaciones SMILES de la misma molécula se aprenden como muestras distintas, lo que provoca una fuga de datos. Causa: Si no se utiliza canonicalize de RDKit, las notaciones como CCO frente a C(C)O se procesan de forma diferente. Solución: Forzar Chem.MolToSmiles(mol, canonical=True) durante el preprocesamiento de los datos de entrenamiento, eliminar los duplicados y aplicar la estandarización de tautómeros (RDKit MolStandardize). Fuente: Discusiones de RDKit [7].

  3. Error en la función collate del DataLoader de PyG (desajuste entre los lotes de moléculas y proteínas) Síntoma: Durante el entrenamiento, las moléculas y las proteínas dentro del lote no coinciden o hay un error de forma. Causa: Los objetos Data de PyG (Data) se apilan automáticamente, pero los tensores de proteínas requieren un manejo separado. Solución: Utilizar una función collate personalizada para: (a) apilar las moléculas con el lote de PyG, (b) apilar las proteínas con torch.stack (torch.stack) y (c) apilar las etiquetas. O bien, adjuntar también las características de la proteína al objeto Data de PyG (Data). Fuente: Discusiones del repositorio GitHub de PyG [8].

  4. Sesgo en las etiquetas debido a la diversidad de los ensayos Síntoma: Las etiquetas de los datos de entrenamiento proceden de múltiples ensayos (Kd bioquímico, IC50 celular, in vivo). Las escalas de los valores difieren. Causa: Rango dinámico y límites de detección distintos según el tipo de ensayo. Solución: (a) Añadir el tipo de ensayo como característica, (b) entrenar únicamente dentro del mismo tipo de ensayo, (c) aplicar una transformación logarítmica seguida de una normalización z-score, (d) aprendizaje multitarea (separando la cabeza por ensayo). Fuente: Blog de química de RDKit de Landrum et al. [7].

  5. Trampas en la interpretación de la visualización de la atención Síntoma: Que un átomo tenga una alta atención en GAT no implica necesariamente que contribuya significativamente a la unión. Esto lleva a interpretaciones erróneas. Causa: Los pesos de atención son un subproducto del resultado del aprendizaje y no garantizan la causalidad ni el significado biológico. Solución: (a) Combinar en conjunto varios métodos de explicabilidad, como Integrated Gradients y GNNExplainer, (b) utilizar la atención únicamente para generar hipótesis, verificándolas con datos experimentales, (c) confiar solo en los átomos que aparecen de forma estable tras entrenar con múltiples semillas. Fuente: Jain & Wallace "Attention is not Explanation". NAACL 2019 [9].

Ideas de expansión

  • Combinación de incrustaciones base: Se utilizan las incrustaciones de ChemBERTa, MolFormer y SELFormer para inicializar las características de los nodos GNN.
  • Aprendizaje multi-tarea: Se realiza un aprendizaje simultáneo de DTI, solubilidad (acuosa), toxicidad hERG y permeabilidad BBB.
  • Mejora de la interpretabilidad: Se emplea un conjunto de Attention, GNNExplainer e IntegratedGradients.
  • Aprendizaje activo: Se prioriza la experimentación con nuevos pares que presenten alta incertidumbre, se obtienen las etiquetas y se vuelve a entrenar el modelo.
  • GNN con conciencia 3D: Se utilizan coordenadas 3D, como las de SchNet, DimeNet y Equiformer, y se combinan con el acoplamiento (episodio 08).
  • Reordenamiento basado en la estructura: Se revalida la afinidad de unión 3D de los candidatos principales de GNN mediante Boltz-2 (episodio 11).

Próximos episodios

  • Episodio 08 docking-hybrid-diffusion: Se verifica la pose de los candidatos principales predichos por GNN mediante el acoplamiento.
  • Episodio 11 structure-affinity-boltz: Se realiza la clasificación final de afinidad con Boltz-2.
  • Episodio 12 single-cell-perturbation: Se conectan los resultados de DTI con la predicción de la respuesta celular.
  • Episodio 14 bio-mcp-agent: Se exponen las predicciones de DTI a través de la herramienta MCP para la detección autónoma por parte del agente.

Referencias

  1. Öztürk H, Özgür A, Ozkirimli E. "DeepDTA: deep drug-target binding affinity prediction." Bioinformatics 2018. https://academic.oup.com/bioinformatics/article/34/17/i821/5093245
  2. Nguyen T, Le H, Quinn TP, et al. "GraphDTA: Predicting drug-target binding affinity with graph neural networks." Bioinformatics 2020. https://academic.oup.com/bioinformatics/article/37/8/1140/5942970
  3. Yang Z, Zhong W, Zhao L, Chen CY-C. "MGraphDTA: deep multiscale graph neural network for explainable drug-target binding affinity prediction." Chemical Science 2022. https://pubs.rsc.org/en/content/articlelanding/2022/sc/d1sc05180f
  4. Muñoz-Gil G, et al. "GeNNius: An ultrafast drug-target interaction inference method based on graph neural networks." 2024. https://github.com/ML4BM-Lab/GeNNius
  5. Pahikkala T, Airola A, Pietilä S, et al. "Toward more realistic drug-target interaction predictions." Briefings in Bioinformatics 2015.
  6. COVID Moonshot consortium: https://postera.ai/moonshot/
  7. RDKit Discussions: https://github.com/rdkit/rdkit/discussions
  8. PyG GitHub Discussions: https://github.com/pyg-team/pytorch_geometric/discussions
  9. Jain S, Wallace BC. "Attention is not Explanation." NAACL 2019. https://arxiv.org/abs/1902.10186
  10. PyTorch Geometric documentation: https://pytorch-geometric.readthedocs.io/
  11. DGL (Deep Graph Library): https://www.dgl.ai/
  12. BindingDB: https://www.bindingdb.org/
  13. Davis benchmark: Davis MI et al. Nat Biotechnol 2011.
  14. KIBA benchmark: Tang J et al. J Chem Inf Model 2014.
  15. BioSNAP dataset: https://snap.stanford.edu/biodata/
  16. ChemBERTa: https://huggingface.co/DeepChem/ChemBERTa-77M-MTR
  17. MolFormer (IBM): https://github.com/IBM/molformer
  18. DrugBank: https://go.drugbank.com/
  19. Drug Repurposing Hub (Broad Institute): https://clue.io/repurposing
  20. UniProt SARS-CoV-2 Mpro (P0DTD1): https://www.uniprot.org/uniprotkb/P0DTD1/entry

💬 Preguntas y comentarios

0 comentarios

Puedes publicar sin iniciar sesión. Los comentarios de invitados no pueden editarse ni eliminarse después.

0/2000

Cargando...