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:
- DryBench ai-native #2: Fundamentos de redes neuronales
- DryBench ai-native #12: Fundamentos de PyTorch
- DryBench ai-native #13: HuggingFace y API comerciales
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
| Herramienta | Función | Licencia |
|---|---|---|
| RDKit | SMILES → grafo molecular · huella digital | BSD-3-Clause |
| PyTorch Geometric (PyG) | Framework de redes neuronales gráficas | MIT |
| DGL (opcional) | Alternativa a PyG | Apache 2.0 |
| GeNNius (GitHub) | Modelo base GNN ultraligero | MIT (previsto) |
| ChemBERTa (HuggingFace) | Incrustaciones moleculares preentrenadas (opcional) | Apache 2.0 |
| MolFormer (IBM) | Alternativa de modelo base molecular | Apache 2.0 |
| BindingDB · Davis · KIBA · BioSNAP | Conjuntos de datos de entrenamiento y validación | Gratuito para fines académicos |
| DrugBank · DrugRepurposing Hub | Bibliotecas de fármacos | Gratuito 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:
Paso 1. SMILES → gráfico molecular
from dataclasses import dataclassfrom typing import Any
import torchfrom torch_geometric.data import Datafrom 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 dataPaso 2. Codificación de la proteína objetivo (incrustaciones ESM2 o modelo base CNN)
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.
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 outPaso 4. DataLoader · Bucle de entrenamiento
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].
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, valPaso 6. Cribado para el reposicionamiento de fármacos contra la Mpro del SARS-CoV-2
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.
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
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)
| Modelo | Evaluación comparativa (Davis · KIBA · BindingDB) | AUC / Pearson r | Fuente |
|---|---|---|---|
| Random Forest + ECFP | Regresión Davis | r = 0.55 | Legacy |
| DeepDTA (CNN 1D) | Davis | r = 0.62 | Öztürk et al., Bioinformatics 2018 [1] |
| GraphDTA (GAT) | Davis | r = 0.68 | Nguyen et al., Bioinformatics 2020 [2] |
| MGraphDTA | Davis · KIBA | r = 0.72 | Yang et al., Chem Sci 2022 [3] |
| GeNNius | BindingDB | AUC 0.92 | ML4BM Lab 2024 [4] |
| DrugBAN | BioSNAP | AUC 0.89 | Bai et al., Nat Mach Intell 2023 |
| Configuración de "cold-target" | Davis cold-target | r = 0.35~0.50 (disminución notable) | Pahikkala 2015 [5] |
| Evaluación comparativa SARS-CoV-2 Mpro (COVID Moonshot) | Biblioteca de candidatos | Diversas evaluaciones comparativas de la literatura | Moonshot 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)
-
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].
-
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
canonicalizede RDKit, las notaciones comoCCOfrente aC(C)Ose procesan de forma diferente. Solución: ForzarChem.MolToSmiles(mol, canonical=True)durante el preprocesamiento de los datos de entrenamiento, eliminar los duplicados y aplicar la estandarización de tautómeros (RDKitMolStandardize). Fuente: Discusiones de RDKit [7]. -
Error en la función
collatedelDataLoaderde 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 objetosDatade PyG (Data) se apilan automáticamente, pero los tensores de proteínas requieren un manejo separado. Solución: Utilizar una funcióncollatepersonalizada para: (a) apilar las moléculas con el lote de PyG, (b) apilar las proteínas contorch.stack(torch.stack) y (c) apilar las etiquetas. O bien, adjuntar también las características de la proteína al objetoDatade PyG (Data). Fuente: Discusiones del repositorio GitHub de PyG [8]. -
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].
-
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
- Ö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 - 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 - 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 - 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 - Pahikkala T, Airola A, Pietilä S, et al. "Toward more realistic drug-target interaction predictions." Briefings in Bioinformatics 2015.
- COVID Moonshot consortium:
https://postera.ai/moonshot/ - RDKit Discussions:
https://github.com/rdkit/rdkit/discussions - PyG GitHub Discussions:
https://github.com/pyg-team/pytorch_geometric/discussions - Jain S, Wallace BC. "Attention is not Explanation." NAACL 2019.
https://arxiv.org/abs/1902.10186 - PyTorch Geometric documentation:
https://pytorch-geometric.readthedocs.io/ - DGL (Deep Graph Library):
https://www.dgl.ai/ - BindingDB:
https://www.bindingdb.org/ - Davis benchmark: Davis MI et al. Nat Biotechnol 2011.
- KIBA benchmark: Tang J et al. J Chem Inf Model 2014.
- BioSNAP dataset:
https://snap.stanford.edu/biodata/ - ChemBERTa:
https://huggingface.co/DeepChem/ChemBERTa-77M-MTR - MolFormer (IBM):
https://github.com/IBM/molformer - DrugBank:
https://go.drugbank.com/ - Drug Repurposing Hub (Broad Institute):
https://clue.io/repurposing - UniProt SARS-CoV-2 Mpro (P0DTD1):
https://www.uniprot.org/uniprotkb/P0DTD1/entry