From 19ba43666ccf6946a5c99bbd0ce7ab04477c1fc0 Mon Sep 17 00:00:00 2001 From: mperikov <145544228+mperikov@users.noreply.github.com> Date: Wed, 26 Nov 2025 23:40:07 +0300 Subject: [PATCH] hw8 without task 3 --- binding_affinity.ipynb | 725 +++++++++++++++++++++++++++++++++++++++++ 1 file changed, 725 insertions(+) create mode 100644 binding_affinity.ipynb diff --git a/binding_affinity.ipynb b/binding_affinity.ipynb new file mode 100644 index 0000000..69add9b --- /dev/null +++ b/binding_affinity.ipynb @@ -0,0 +1,725 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## Предсказание свободной энергии связывания" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "В этой практике вы реализуете собственную графовую архитектуру для предсказания свободной энергии связывания двух белков, которая будет более точно учитывать их геометрию, но сохранит инвариантность относительно движений в пространстве.\n", + "\n", + "На практике сделали всю подготовительную работу для проведения экспериментов, а также обнаружили, что в случае простой графовой модели лучший результат дал граф, построенный на атомной структуре интерфейса, но лишённый внутримолекулярных связей, т.е. рёбер, соединяющих атомы одной и той же молекулы.\n", + "\n", + "Однако, наша модель была крайне простой, и в своих экспериментах вы можете обнаружить, что другой представление входных данных в сочетании с более сложной архитектурой сработает ещё лучше. В качестве бонусного задания вы сможете провести любые эксперименты с архитектурой и способом представления данных." + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "#### Подготовка данных (с практики по GNN)" + ] + }, + { + "cell_type": "code", + "execution_count": 1, + "metadata": {}, + "outputs": [], + "source": [ + "import sys\n", + "from pathlib import Path\n", + "\n", + "import pandas as pd\n", + "import torch\n", + "import torch.nn.functional as F\n", + "from torch import Tensor, nn\n", + "from torch.optim import Adam\n", + "from torch_geometric.loader import DataLoader\n", + "from torch_geometric.nn.conv import (\n", + " GATConv,\n", + " GatedGraphConv,\n", + " GCNConv,\n", + " GraphConv,\n", + " MessagePassing,\n", + ")\n", + "from torch_geometric.nn.pool import global_mean_pool\n", + "\n", + "sys.path.append(str(Path.cwd().parent))\n", + "from assets.utils.affinity_dataset import (\n", + " ATOMS_INDICES,\n", + " AffinityDataset,\n", + " AtomicInterfaceGraphBuilder,\n", + " DataItem,\n", + " InterfaceGraph,\n", + " PlotlyVis,\n", + ")" + ] + }, + { + "cell_type": "code", + "execution_count": 2, + "metadata": {}, + "outputs": [ + { + "data": { + "text/plain": [ + "Data(edge_index=[2, 1532], y=-12.91, atoms=[333], residues=[333], coordinates=[333, 3], receptor_mask=[333], distances=[1532], num_nodes=333)" + ] + }, + "execution_count": 2, + "metadata": {}, + "output_type": "execute_result" + } + ], + "source": [ + "dataset_dir = Path(\"../assets/datasets/binding_affinity/\")\n", + "pdb_dir = dataset_dir / \"pdb\"\n", + "train_csv = pd.read_csv(dataset_dir / \"affinity_train.csv\")\n", + "record = train_csv.iloc[0]\n", + "item = DataItem(\n", + " uid=record[\"uid\"],\n", + " receptor_chains=record[\"receptor_chains\"],\n", + " ligand_chains=record[\"ligand_chains\"],\n", + " dG=record[\"dG\"],\n", + " pdb=pdb_dir / f'{record[\"uid\"]}.pdb',\n", + ")\n", + "graph_builder = AtomicInterfaceGraphBuilder(\n", + " interface_distance=5.0, radius=5.0, keep_inner_edges=False\n", + ")\n", + "graph = graph_builder.build_graph(item)\n", + "graph" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "#### Задание 1 (5 баллов). Реализация E(3)-инвариантной графовой сети" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "В нашей простой модели мы использовали межатомные расстояния, чтобы построить граф, но далее никакую информацию о геометрии интерфейса не использовали.\n", + "\n", + "Тем не менее, точное относительное положение атомов может существенно определять силу и характер физических взаимодействий.\n", + "\n", + "В этом задании вы реализуете архитектуру графовой сети, которая использует межатомные расстояния при создании сообщений, которыми обмениваются вершины графа. Тем самым результат не будет зависеть от положения и ориентации белкового комплекса в пространстве, но будет явным образом зависеть от геометрии атомных контактов.\n", + "\n", + "Благодаря `pytorch-geometric` реализация таких моделей сравнительно простая, но чтобы не возникло впечатления, что фреймворк делает совсем какую-то магию, перед выполнением задания ознакомьтесь с туториалом по реализации message-passing neural networks: https://pytorch-geometric.readthedocs.io/en/stable/tutorial/create_gnn.html" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "##### Задание 1.1 (2 балла). E(3)-инвариантный слой графовой сети" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "Наш слой будет обновлять эмбеддинги вершин в соответствии с уравнением\n", + "\n", + "$h_i^{(t+1)} = \\sum_{j \\in \\mathcal{N}(i)} \\text{MLP}^{(t)} \\left( \\text{concat} (h_i^{(t)}, h_j^{(t)}, e_{ij}) \\right)$\n", + "\n", + "т.е. сообщение между вершинами $i$ и $j$ будет формироваться перцептроном, который принимает на вход эмбеддинги вершин и эмбеддинг соединяющего их ребра\n", + "\n", + "Всю работу по распространению сообщений сделает метод `propagate`, вам нужно только реализовать метод `message`, который эти сообщения сформирует" + ] + }, + { + "cell_type": "code", + "execution_count": 3, + "metadata": {}, + "outputs": [], + "source": [ + "class InvariantLayer(MessagePassing):\n", + " def __init__(\n", + " self, edge_dim: int, node_dim: int, hidden_dim: int, aggr: str = \"sum\"\n", + " ) -> None:\n", + " super().__init__(aggr)\n", + " self.message_mlp = nn.Sequential(\n", + " nn.Linear(2 * node_dim + edge_dim, hidden_dim),\n", + " nn.SiLU(),\n", + " nn.Linear(hidden_dim, node_dim),\n", + " )\n", + "\n", + " def forward(self, h: Tensor, edge_index: Tensor, edge_attr: Tensor) -> Tensor:\n", + " return self.propagate(edge_index, h=h, edge_attr=edge_attr)\n", + "\n", + " def message(self, h_i, h_j, edge_attr) -> Tensor:\n", + " return self.message_mlp(torch.cat([h_i, h_j, edge_attr], dim=1))" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "Минимальный тест на работоспособность:" + ] + }, + { + "cell_type": "code", + "execution_count": 4, + "metadata": {}, + "outputs": [], + "source": [ + "h = torch.randn(4, 8)\n", + "edge_index = torch.tensor(\n", + " [\n", + " [0, 0, 1, 1, 2],\n", + " [1, 3, 2, 3, 3],\n", + " ]\n", + ")\n", + "edge_attr = torch.randn(5, 6)\n", + "\n", + "assert InvariantLayer(6, 8, 10).forward(h, edge_index, edge_attr).shape == torch.Size(\n", + " [4, 8]\n", + ")" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "##### Задание 1.2 (3 балла). E(3)-инвариантная графовая сеть" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "Постройте модель на основе реализованного вами слоя, которая принимает на вход `InterfaceGraph` и возвращает предсказанную свободную энергию связывания.\n", + "\n", + "Отличия от модели с практики небольшие:\n", + "1. Вместо слоя `GraphConv` в модели должен быть ваш `InvariantLayer`\n", + "2. Нужно преобразовать расстояния с помощью модуля `RadialBasisExpansion` и передавать их в каждый `InvariantLayer` вместе с очередными эмбеддингами вершин.\n", + "3. Для достижения нужной точности может потребоваться добавить нормализацию, например `nn.LayerNorm`\n", + "\n", + "Модуль `RadialBasisExpansion` преобразует значения межатомных расстояний в вектор со значениями в [0, 1] с помощью набора радиальных базисных функций. Подумайте, почему такой способ обработки количественных признаков может работать лучше?" + ] + }, + { + "cell_type": "code", + "execution_count": 5, + "metadata": {}, + "outputs": [ + { + "data": { + "text/plain": [ + "tensor([[0.2530, 0.0290, 0.0010, 0.0000, 0.0000],\n", + " [0.6580, 0.1600, 0.0140, 0.0000, 0.0000],\n", + " [0.9010, 0.3460, 0.0490, 0.0030, 0.0000],\n", + " [0.9600, 0.7750, 0.2300, 0.0250, 0.0010],\n", + " [0.7260, 0.9800, 0.4870, 0.0890, 0.0060]])" + ] + }, + "execution_count": 5, + "metadata": {}, + "output_type": "execute_result" + } + ], + "source": [ + "class RadialBasisExpansion(nn.Module):\n", + " offset: Tensor\n", + "\n", + " def __init__(\n", + " self,\n", + " start: float = 3.0,\n", + " stop: float = 10.0,\n", + " num_gaussians: int = 32,\n", + " ):\n", + " super().__init__()\n", + " offset = torch.linspace(start, stop, num_gaussians)\n", + " self.coeff = -0.5 / (offset[1] - offset[0]).item() ** 2\n", + " self.register_buffer(\"offset\", offset)\n", + "\n", + " def forward(self, dist: Tensor) -> Tensor:\n", + " dist = dist.view(-1, 1) - self.offset.view(1, -1)\n", + " return torch.exp(self.coeff * torch.pow(dist, 2))\n", + "\n", + "\n", + "# пример использования\n", + "dist = torch.tensor([0.1, 1.4, 2.2, 3.5, 4.4])\n", + "RadialBasisExpansion(num_gaussians=5).forward(dist).round(decimals=3)" + ] + }, + { + "cell_type": "code", + "execution_count": 6, + "metadata": {}, + "outputs": [], + "source": [ + "class InvariantGNN(nn.Module):\n", + " def __init__(\n", + " self,\n", + " node_vocab_size: int, # кол-во типов вершин, например атомов\n", + " node_dim: int, # размерность эмбеддинга вершины\n", + " edge_dim: int, # размерность эмбеддинга ребра\n", + " n_layers: int, # кол-во графовых слоёв\n", + " dropout: float = 0.0, # dropout rate\n", + " hidden_dim: int = 64,\n", + " ) -> None:\n", + " super().__init__()\n", + " # эмбеддинг для типов атомов\n", + " self.embed = nn.Embedding(node_vocab_size, node_dim)\n", + " # список графовых слоёв\n", + " self.conv = nn.ModuleList(\n", + " [InvariantLayer(edge_dim, node_dim, hidden_dim) for _ in range(n_layers)]\n", + " )\n", + "\n", + " # линейный слой для регрессии\n", + " self.fc = nn.Linear(node_dim, 1)\n", + " self.dropout = nn.Dropout(dropout, inplace=True)\n", + " self.norm = nn.LayerNorm(node_dim)\n", + " self.expansion = RadialBasisExpansion(num_gaussians=edge_dim)\n", + "\n", + " def forward(self, batch: InterfaceGraph) -> Tensor:\n", + " # 1. Эмбеддинги вершин\n", + " x = self.embed(batch.atoms)\n", + " dist = self.expansion(batch.distances)\n", + " for conv in self.conv:\n", + " x = (x + conv(x, batch.edge_index, dist)).relu()\n", + "\n", + " # 2. Эмбеддинг графа: усреднение по вершинам отдельных графов\n", + " x = global_mean_pool(x, batch.batch) # [batch_size, hidden_channels]\n", + "\n", + " # 3. Финальный регрессор поверх эмбеддинга графа\n", + " x = self.norm(x)\n", + " x = self.dropout(x)\n", + " x = self.fc(x)\n", + " return x" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "Минимальный тест:" + ] + }, + { + "cell_type": "code", + "execution_count": 7, + "metadata": {}, + "outputs": [], + "source": [ + "train_dataset = AffinityDataset(\n", + " datadir=pdb_dir,\n", + " subset_csv=dataset_dir / \"affinity_train.csv\",\n", + " graph_builder=graph_builder,\n", + ")\n", + "train_loader = DataLoader(train_dataset, batch_size=4, shuffle=True)\n", + "batch = next(iter(train_loader))\n", + "model = InvariantGNN(\n", + " node_vocab_size=len(ATOMS_INDICES) + 1,\n", + " node_dim=32,\n", + " edge_dim=16,\n", + " n_layers=2,\n", + " dropout=0.1,\n", + ")\n", + "assert model.forward(batch).shape == torch.Size([4, 1])" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "#### Задание 2 (4 балла). Обучение модели" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "Обучите реализованную модель, выведите в конце обучения метрики на тестовой выборке (MAE, корреляции Пирсона и Спирмена).\n", + "\n", + "Ваша задача: добиться MAE < 1.55\n", + "\n", + "Используйте `AtomicInterfaceGraphBuilder(interface_distance=5.0, radius=5.0, keep_inner_edges=False)`, эти параметры можно будет изменить в следующем задании.\n", + "\n", + "Но вы можете выбрать любой размер модели и способ регуляризации, а также любой оптимизатор." + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "*Без изменений архитектуры получилось добиться только MAE < 1.56, поэтому применил несколько модификаций из задания 3, а именно: \n", + "1)К эмбеддингам атомов добавил эмбеддинги аминокислот \n", + "2)Добавил линейный слой, преобразующий эмбеддинги рёбер*" + ] + }, + { + "cell_type": "code", + "execution_count": 8, + "metadata": {}, + "outputs": [], + "source": [ + "class InvariantGNN(nn.Module):\n", + " def __init__(\n", + " self,\n", + " node_vocab_size: int, # кол-во типов вершин, например атомов\n", + " node_dim: int, # размерность эмбеддинга вершины\n", + " edge_dim: int, # размерность эмбеддинга ребра\n", + " n_layers: int, # кол-во графовых слоёв\n", + " dropout: float = 0.0, # dropout rate\n", + " hidden_dim: int = 64,\n", + " ) -> None:\n", + " super().__init__()\n", + " # эмбеддинг для типов атомов\n", + " self.embed = nn.Embedding(node_vocab_size, node_dim)\n", + " self.embed_residues = nn.Embedding(node_vocab_size, node_dim)\n", + " # список графовых слоёв\n", + " self.conv = nn.ModuleList(\n", + " [InvariantLayer(edge_dim, 2 * node_dim, hidden_dim) for _ in range(n_layers)]\n", + " )\n", + "\n", + " # линейный слой для регрессии\n", + " self.fc = nn.Linear(2 * node_dim, 1)\n", + " self.dropout = nn.Dropout(dropout, inplace=True)\n", + " self.norm = nn.LayerNorm(2 * node_dim)\n", + " self.expansion = RadialBasisExpansion(num_gaussians=edge_dim)\n", + "\n", + " self.edge_linear = nn.Linear(edge_dim, edge_dim)\n", + "\n", + " def forward(self, batch: InterfaceGraph) -> Tensor:\n", + " # 1. Эмбеддинги вершин\n", + " x = torch.cat([self.embed(batch.atoms), self.embed_residues(batch.residues)], dim=1)\n", + " dist = self.edge_linear(self.expansion(batch.distances))\n", + " for conv in self.conv:\n", + " x = (x + conv(x, batch.edge_index, dist)).relu()\n", + "\n", + " # 2. Эмбеддинг графа: усреднение по вершинам отдельных графов\n", + " x = global_mean_pool(x, batch.batch) # [batch_size, hidden_channels]\n", + "\n", + " # 3. Финальный регрессор поверх эмбеддинга графа\n", + " x = self.norm(x)\n", + " x = self.dropout(x)\n", + " x = self.fc(x)\n", + " return x" + ] + }, + { + "cell_type": "code", + "execution_count": 9, + "metadata": {}, + "outputs": [], + "source": [ + "from scipy.stats import pearsonr, spearmanr\n", + "\n", + "device = torch.device('cuda')\n", + "\n", + "@torch.no_grad()\n", + "def validate(loader: DataLoader, model: nn.Module) -> tuple[list[float], list[float]]:\n", + " model.eval()\n", + " ys = []\n", + " yhats = []\n", + " loss = 0.0\n", + " for batch in loader:\n", + " batch = batch.to(device)\n", + " \n", + " yhat = model.forward(batch)\n", + " yhats.extend(yhat.flatten().tolist())\n", + " ys.extend(batch.y.tolist())\n", + " loss += F.mse_loss(yhat.flatten(), batch.y, reduction=\"sum\").item()\n", + "\n", + " print(f\"Loss: {loss / len(ys):.4f}, \", end=\"\")\n", + " print(f\"MAE: {(torch.tensor(ys).to(device) - torch.tensor(yhats).to(device)).abs().mean():.4f}, \", end=\"\")\n", + " print(f\"Pearson R: {pearsonr(ys, yhats).statistic:.4f}, \", end=\"\")\n", + " print(f\"Spearman R: {spearmanr(ys, yhats).statistic:.4f}\")\n", + " model.train()\n", + " return yhats, ys" + ] + }, + { + "cell_type": "code", + "execution_count": 10, + "metadata": {}, + "outputs": [], + "source": [ + "graph_builder = AtomicInterfaceGraphBuilder(\n", + " interface_distance=5.0, radius=5.0, keep_inner_edges=False\n", + ")\n", + "train_dataset = AffinityDataset(\n", + " datadir=pdb_dir,\n", + " subset_csv=dataset_dir / \"affinity_train.csv\",\n", + " graph_builder=graph_builder,\n", + ")\n", + "test_dataset = AffinityDataset(\n", + " datadir=pdb_dir,\n", + " subset_csv=dataset_dir / \"affinity_test.csv\",\n", + " graph_builder=graph_builder,\n", + ")\n", + "train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True)\n", + "test_loader = DataLoader(test_dataset, batch_size=64, shuffle=False)" + ] + }, + { + "cell_type": "code", + "execution_count": 11, + "metadata": {}, + "outputs": [], + "source": [ + "torch.manual_seed(42)\n", + "\n", + "model = InvariantGNN(\n", + " node_vocab_size=len(ATOMS_INDICES) + 1,\n", + " node_dim=128,\n", + " edge_dim=128,\n", + " n_layers=3,\n", + " dropout=0.5,\n", + " hidden_dim=64,\n", + ")\n", + "\n", + "model.to(device)\n", + "\n", + "optim = torch.optim.AdamW(model.parameters(), lr=0.0001, weight_decay=0.005)\n", + "scheduler = torch.optim.lr_scheduler.MultiStepLR(optim, milestones=[25], gamma=0.1)" + ] + }, + { + "cell_type": "code", + "execution_count": 12, + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "Loss: 6.5458, MAE: 2.0336, Pearson R: 0.0911, Spearman R: 0.1304\n", + "Loss: 3.9223, MAE: 1.5941, Pearson R: 0.1163, Spearman R: 0.1769\n", + "Loss: 3.7547, MAE: 1.5839, Pearson R: 0.2597, Spearman R: 0.2841\n", + "Loss: 3.7239, MAE: 1.5590, Pearson R: 0.2928, Spearman R: 0.2781\n", + "Loss: 3.6283, MAE: 1.5478, Pearson R: 0.2706, Spearman R: 0.2738\n", + "Loss: 3.6193, MAE: 1.5461, Pearson R: 0.2675, Spearman R: 0.2670\n", + "Loss: 3.6076, MAE: 1.5418, Pearson R: 0.2697, Spearman R: 0.2691\n", + "Loss: 3.5997, MAE: 1.5378, Pearson R: 0.2750, Spearman R: 0.2759\n", + "Loss: 3.5835, MAE: 1.5391, Pearson R: 0.2733, Spearman R: 0.2773\n", + "Loss: 3.6000, MAE: 1.5390, Pearson R: 0.2695, Spearman R: 0.2663\n" + ] + } + ], + "source": [ + "for i in range(50):\n", + " model.train()\n", + " for batch in train_loader:\n", + " batch = batch.to(device)\n", + " \n", + " yhat = model.forward(batch)\n", + " loss = F.mse_loss(yhat.flatten(), batch.y)\n", + " loss.backward()\n", + " optim.step()\n", + " optim.zero_grad()\n", + "\n", + " scheduler.step()\n", + "\n", + " if (i + 1) % 5 == 0:\n", + " validate(test_loader, model)" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "#### Задание 3 (необязательное). В погоне за точностью" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "Используйте любую графовую архитектуру, чтобы добиться MAE < 1.45.\n", + "\n", + "Баллы за задание:\n", + "- 3 балла — за MAE < 1.45\n", + "- +1 балл за каждые следующие 0.01\n", + "\n", + "Задание с полной свободой творчества, можно менять и архитектуру модели, и использовать любые модули из `pytorch-geometric`, и менять способ представления данных. Вот некоторые идеи, которые можно тестировать:\n", + "1. **Модификация модели с практики**: она является достаточно сильным бейзлайном, поэтому может иметь смысл поколдовать над ней: поменять гиперпараметры, функции активации, используемую функцию ошибки (например huber loss или log-cosh)\n", + "2. **Использование информации об аминокислотах**: В наших моделях мы кодируем только тип атома, но никак не используем информацию об аминокислотах, к которым эти атомы относятся. Можно к эмбеддингам атомов добавить эмбеддинги аминокислот, индексы которых находятся в атрибуте `residues`.\n", + "3. **Модификация реализованной модели**: тут много вариантов, например\n", + " - добавить линейный слой / перцептрон, который будет в каждом графовом слое преобразовывать эмбеддинг рёбер\n", + " - изменить метод `message`, чтобы иначе формировать сообщения\n", + " - изменить метод `update`, чтобы использовать более гибкий метод агрегации сообщений от соседей; например, реализовать механизм внимания, как в `torch_geometric.nn.conv.GATConv` \n", + "4. **Включение внутрибелковых рёбер**: возможно, модели не хватает обмена информацией с соседними вершинами того же белка, но в архитектуре модели сейчас нет ничего, что учитывает тип ребра: внутреннее (между атомами одного белка) и внешнее (между атомами рецептора и лиганда). Можно добавить в представление ребра его тип: как бинарную переменную или как эмбеддинг (`nn.Embedding(2, edge_dim)`). Получить тип ребра можно из тензоров `edge_index` и `receptor_mask`." + ] + }, + { + "cell_type": "code", + "execution_count": 66, + "metadata": {}, + "outputs": [], + "source": [ + "class InvariantGNN(nn.Module):\n", + " def __init__(\n", + " self,\n", + " node_vocab_size: int, # кол-во типов вершин, например атомов\n", + " node_dim: int, # размерность эмбеддинга вершины\n", + " edge_dim: int, # размерность эмбеддинга ребра\n", + " n_layers: int, # кол-во графовых слоёв\n", + " dropout: float = 0.0, # dropout rate\n", + " hidden_dim: int = 64,\n", + " ) -> None:\n", + " super().__init__()\n", + " # эмбеддинг для типов атомов\n", + " self.embed = nn.Embedding(node_vocab_size, node_dim)\n", + " self.embed_residues = nn.Embedding(node_vocab_size, node_dim)\n", + " self.embed_edge = nn.Embedding(2, edge_dim)\n", + " # список графовых слоёв\n", + " self.conv = nn.ModuleList(\n", + " [GATv2Conv(2 * node_dim, 2 * node_dim, edge_dim=2 * edge_dim) for _ in range(n_layers)]\n", + " )\n", + "\n", + " # линейный слой для регрессии\n", + " self.fc = nn.Linear(2 * node_dim, 1)\n", + " self.dropout = nn.Dropout(dropout, inplace=True)\n", + " self.norm = nn.LayerNorm(2 * node_dim)\n", + " self.expansion = RadialBasisExpansion(num_gaussians=edge_dim)\n", + "\n", + "\n", + "\n", + " def forward(self, batch: InterfaceGraph) -> Tensor:\n", + " # 1. Эмбеддинги вершин\n", + " x = torch.cat([self.embed(batch.atoms), self.embed_residues(batch.residues)], dim=1)\n", + "\n", + " edge_type = batch.receptor_mask[batch.edge_index[0]] == batch.receptor_mask[batch.edge_index[1]]\n", + "\n", + " edge_type = edge_type.to(torch.long)\n", + "\n", + " dist = torch.cat([self.expansion(batch.distances), self.embed_edge(edge_type)], dim=1)\n", + " \n", + " for conv in self.conv:\n", + " x = (x + conv(x, batch.edge_index, dist)).relu()\n", + "\n", + " # 2. Эмбеддинг графа: усреднение по вершинам отдельных графов\n", + " x = global_mean_pool(x, batch.batch) # [batch_size, hidden_channels]\n", + "\n", + " # 3. Финальный регрессор поверх эмбеддинга графа\n", + " x = self.norm(x)\n", + " x = self.dropout(x)\n", + " x = self.fc(x)\n", + " return x" + ] + }, + { + "cell_type": "code", + "execution_count": 51, + "metadata": {}, + "outputs": [], + "source": [ + "graph_builder = AtomicInterfaceGraphBuilder(\n", + " interface_distance=5.0, radius=5.0, keep_inner_edges=True\n", + ")\n", + "train_dataset = AffinityDataset(\n", + " datadir=pdb_dir,\n", + " subset_csv=dataset_dir / \"affinity_train.csv\",\n", + " graph_builder=graph_builder,\n", + ")\n", + "test_dataset = AffinityDataset(\n", + " datadir=pdb_dir,\n", + " subset_csv=dataset_dir / \"affinity_test.csv\",\n", + " graph_builder=graph_builder,\n", + ")\n", + "train_loader = DataLoader(train_dataset, batch_size=16, shuffle=True)\n", + "test_loader = DataLoader(test_dataset, batch_size=16, shuffle=False)" + ] + }, + { + "cell_type": "code", + "execution_count": 67, + "metadata": {}, + "outputs": [], + "source": [ + "torch.manual_seed(42)\n", + "\n", + "\n", + "model = InvariantGNN(\n", + " node_vocab_size=len(ATOMS_INDICES) + 1,\n", + " node_dim=64,\n", + " edge_dim=64,\n", + " n_layers=3,\n", + " dropout=0.5,\n", + " hidden_dim=64,\n", + ")\n", + "\n", + "model.to(device)\n", + "\n", + "optim = torch.optim.AdamW(model.parameters(), lr=0.0001, weight_decay=0.005)\n", + "scheduler = torch.optim.lr_scheduler.MultiStepLR(optim, milestones=[25, 50, 75], gamma=0.1)" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "Loss: 7.9297, MAE: 2.2743, Pearson R: 0.1544, Spearman R: 0.1460\n" + ] + } + ], + "source": [ + "for i in range(100):\n", + " model.train()\n", + " for batch in train_loader:\n", + " batch = batch.to(device)\n", + " \n", + " yhat = model.forward(batch)\n", + " loss = F.mse_loss(yhat.flatten(), batch.y)\n", + " loss.backward()\n", + " optim.step()\n", + " optim.zero_grad()\n", + "\n", + " scheduler.step()\n", + "\n", + " if (i + 1) % 5 == 0:\n", + " validate(test_loader, model)" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [] + } + ], + "metadata": { + "kernelspec": { + "display_name": "Python 3 (ipykernel)", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.13.7" + } + }, + "nbformat": 4, + "nbformat_minor": 4 +}