@pikaa-ai/pikaa 0.3.23 → 0.3.25
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- package/assets/brand/orbit-logo-option4-whale.jpg +0 -0
- package/assets/brand/orbit-logo.jpg +0 -0
- package/assets/brand/orbit-logo.png +0 -0
- package/assets/brand/orbit-logo.svg +3 -0
- package/dist/cli.js +407 -219
- package/dist/index.js +7 -2
- package/package.json +1 -2
- package/skills/adaptyv/SKILL.md +0 -240
- package/skills/aeon/SKILL.md +0 -402
- package/skills/analytical-method-validation/SKILL.md +0 -299
- package/skills/anndata/SKILL.md +0 -431
- package/skills/arbor/SKILL.md +0 -152
- package/skills/arboreto/SKILL.md +0 -267
- package/skills/astropy/SKILL.md +0 -353
- package/skills/autoskill/SKILL.md +0 -233
- package/skills/benchling-integration/SKILL.md +0 -229
- package/skills/bgpt-paper-search/SKILL.md +0 -75
- package/skills/bids/SKILL.md +0 -237
- package/skills/biopython/SKILL.md +0 -472
- package/skills/bioservices/SKILL.md +0 -399
- package/skills/bulk-rnaseq/SKILL.md +0 -198
- package/skills/cellxgene-census/SKILL.md +0 -283
- package/skills/cirq/SKILL.md +0 -370
- package/skills/citation-management/SKILL.md +0 -329
- package/skills/clinical-decision-support/SKILL.md +0 -238
- package/skills/clinical-decision-support/references/README.md +0 -62
- package/skills/clinical-reports/SKILL.md +0 -248
- package/skills/clinical-reports/references/README.md +0 -34
- package/skills/cobrapy/SKILL.md +0 -496
- package/skills/consciousness-council/SKILL.md +0 -151
- package/skills/dask/SKILL.md +0 -482
- package/skills/database-lookup/SKILL.md +0 -386
- package/skills/datamol/SKILL.md +0 -200
- package/skills/deepchem/SKILL.md +0 -244
- package/skills/deepspot-m/SKILL.md +0 -175
- package/skills/deeptools/SKILL.md +0 -412
- package/skills/depmap/SKILL.md +0 -301
- package/skills/dhdna-profiler/SKILL.md +0 -184
- package/skills/diffdock/SKILL.md +0 -488
- package/skills/dnanexus-integration/SKILL.md +0 -325
- package/skills/docx/SKILL.md +0 -99
- package/skills/esm/SKILL.md +0 -334
- package/skills/etetoolkit/SKILL.md +0 -327
- package/skills/exa-search/SKILL.md +0 -102
- package/skills/executing-plans/SKILL.md +0 -14
- package/skills/experimental-design/SKILL.md +0 -234
- package/skills/exploratory-data-analysis/SKILL.md +0 -280
- package/skills/flowio/SKILL.md +0 -310
- package/skills/fluidsim/SKILL.md +0 -279
- package/skills/frontend-design/SKILL.md +0 -100
- package/skills/generate-image/SKILL.md +0 -304
- package/skills/geniml/SKILL.md +0 -310
- package/skills/genomic-coordinates/SKILL.md +0 -189
- package/skills/genomic-intelligence/SKILL.md +0 -243
- package/skills/geomaster/README.md +0 -105
- package/skills/geomaster/SKILL.md +0 -366
- package/skills/geopandas/SKILL.md +0 -250
- package/skills/get-available-resources/SKILL.md +0 -260
- package/skills/gget/SKILL.md +0 -153
- package/skills/ginkgo-cloud-lab/SKILL.md +0 -106
- package/skills/glycoengineering/SKILL.md +0 -339
- package/skills/gtars/SKILL.md +0 -282
- package/skills/guardian-rails/SKILL.md +0 -54
- package/skills/histolab/SKILL.md +0 -243
- package/skills/hugging-science/SKILL.md +0 -132
- package/skills/hypogenic/SKILL.md +0 -290
- package/skills/hypothesis-generation/SKILL.md +0 -264
- package/skills/imaging-data-commons/SKILL.md +0 -496
- package/skills/infographics/SKILL.md +0 -315
- package/skills/iso-standards-readiness/SKILL.md +0 -352
- package/skills/lab-hardware-cad/SKILL.md +0 -372
- package/skills/labarchive-integration/SKILL.md +0 -216
- package/skills/lamindb/SKILL.md +0 -408
- package/skills/latchbio-integration/SKILL.md +0 -227
- package/skills/latex-posters/SKILL.md +0 -369
- package/skills/latex-posters/references/README.md +0 -439
- package/skills/liteparse/SKILL.md +0 -295
- package/skills/literature-review/SKILL.md +0 -263
- package/skills/markdown-mermaid-writing/SKILL.md +0 -322
- package/skills/market-research-reports/SKILL.md +0 -337
- package/skills/markitdown/SKILL.md +0 -264
- package/skills/matchms/SKILL.md +0 -276
- package/skills/matlab/SKILL.md +0 -274
- package/skills/matplotlib/SKILL.md +0 -378
- package/skills/medchem/SKILL.md +0 -321
- package/skills/modal/SKILL.md +0 -468
- package/skills/molecular-dynamics/SKILL.md +0 -458
- package/skills/molfeat/SKILL.md +0 -348
- package/skills/ncats-arax/SKILL.md +0 -178
- package/skills/networkx/SKILL.md +0 -440
- package/skills/neurokit2/SKILL.md +0 -323
- package/skills/neuropixels-analysis/SKILL.md +0 -412
- package/skills/nextflow/SKILL.md +0 -195
- package/skills/omero-integration/SKILL.md +0 -222
- package/skills/onekgpd/SKILL.md +0 -371
- package/skills/ontology-term-resolution/SKILL.md +0 -147
- package/skills/open-notebook/SKILL.md +0 -297
- package/skills/openpiv/SKILL.md +0 -469
- package/skills/opentrons-integration/SKILL.md +0 -322
- package/skills/optimize-for-gpu/SKILL.md +0 -176
- package/skills/owasp-top10/SKILL.md +0 -48
- package/skills/pacsomatic/LICENSE +0 -21
- package/skills/pacsomatic/SKILL.md +0 -150
- package/skills/paper-lookup/SKILL.md +0 -263
- package/skills/paperclip/SKILL.md +0 -413
- package/skills/paperzilla/SKILL.md +0 -159
- package/skills/parallel-web/SKILL.md +0 -128
- package/skills/pathml/SKILL.md +0 -222
- package/skills/pathogen-variant-surveillance/SKILL.md +0 -208
- package/skills/pathway-enrichment/SKILL.md +0 -194
- package/skills/pdf/SKILL.md +0 -322
- package/skills/peer-review/SKILL.md +0 -288
- package/skills/penetration-testing/SKILL.md +0 -31
- package/skills/pennylane/SKILL.md +0 -240
- package/skills/phylogenetics/SKILL.md +0 -409
- package/skills/pi-agent/SKILL.md +0 -83
- package/skills/pkpd-modeling/SKILL.md +0 -381
- package/skills/polars/SKILL.md +0 -393
- package/skills/polars-bio/SKILL.md +0 -379
- package/skills/ponytail/SKILL.md +0 -31
- package/skills/ponytail-audit/SKILL.md +0 -18
- package/skills/pptx/SKILL.md +0 -246
- package/skills/pptx-posters/SKILL.md +0 -258
- package/skills/primekg/SKILL.md +0 -99
- package/skills/protocolsio-integration/SKILL.md +0 -236
- package/skills/pufferlib/SKILL.md +0 -328
- package/skills/pydeseq2/SKILL.md +0 -369
- package/skills/pydicom/SKILL.md +0 -381
- package/skills/pyhealth/SKILL.md +0 -124
- package/skills/pylabrobot/SKILL.md +0 -216
- package/skills/pymatgen/SKILL.md +0 -404
- package/skills/pymc/SKILL.md +0 -310
- package/skills/pymoo/SKILL.md +0 -276
- package/skills/pyopenms/SKILL.md +0 -179
- package/skills/pysam/SKILL.md +0 -330
- package/skills/pytdc/SKILL.md +0 -297
- package/skills/pytorch-lightning/SKILL.md +0 -191
- package/skills/pyzotero/SKILL.md +0 -137
- package/skills/qiskit/SKILL.md +0 -259
- package/skills/qutip/SKILL.md +0 -317
- package/skills/rdkit/SKILL.md +0 -94
- package/skills/relsa-severity-assessment/SKILL.md +0 -354
- package/skills/research-grants/SKILL.md +0 -296
- package/skills/research-grants/references/README.md +0 -287
- package/skills/research-lookup/README.md +0 -106
- package/skills/research-lookup/SKILL.md +0 -338
- package/skills/rowan/SKILL.md +0 -398
- package/skills/scanpy/SKILL.md +0 -303
- package/skills/scholar-evaluation/SKILL.md +0 -296
- package/skills/scientific-brainstorming/SKILL.md +0 -282
- package/skills/scientific-critical-thinking/SKILL.md +0 -180
- package/skills/scientific-schematics/SKILL.md +0 -370
- package/skills/scientific-slides/SKILL.md +0 -379
- package/skills/scientific-visualization/SKILL.md +0 -285
- package/skills/scientific-writing/SKILL.md +0 -356
- package/skills/scikit-bio/SKILL.md +0 -470
- package/skills/scikit-learn/SKILL.md +0 -324
- package/skills/scikit-survival/SKILL.md +0 -313
- package/skills/scvelo/SKILL.md +0 -328
- package/skills/scvi-tools/SKILL.md +0 -201
- package/skills/seaborn/SKILL.md +0 -254
- package/skills/security-auditor/SKILL.md +0 -37
- package/skills/shap/SKILL.md +0 -282
- package/skills/simpy/SKILL.md +0 -283
- package/skills/stable-baselines3/SKILL.md +0 -325
- package/skills/statistical-analysis/SKILL.md +0 -446
- package/skills/statistical-power/SKILL.md +0 -200
- package/skills/statsmodels/SKILL.md +0 -238
- package/skills/sympy/SKILL.md +0 -354
- package/skills/systematic-debugging/SKILL.md +0 -35
- package/skills/tamarind/SKILL.md +0 -285
- package/skills/tdd/SKILL.md +0 -26
- package/skills/tiledbvcf/SKILL.md +0 -456
- package/skills/timesfm-forecasting/SKILL.md +0 -408
- package/skills/timesfm-forecasting/examples/global-temperature/README.md +0 -178
- package/skills/torch-geometric/SKILL.md +0 -458
- package/skills/torchdrug/SKILL.md +0 -241
- package/skills/transformers/SKILL.md +0 -195
- package/skills/treatment-plans/SKILL.md +0 -174
- package/skills/treatment-plans/references/README.md +0 -19
- package/skills/umap-learn/SKILL.md +0 -488
- package/skills/uncertainty-and-units/SKILL.md +0 -384
- package/skills/usfiscaldata/SKILL.md +0 -171
- package/skills/vaex/SKILL.md +0 -204
- package/skills/venue-templates/SKILL.md +0 -269
- package/skills/verification-before-completion/SKILL.md +0 -22
- package/skills/waypoint-bio/SKILL.md +0 -273
- package/skills/what-if-oracle/SKILL.md +0 -184
- package/skills/writing-plans/SKILL.md +0 -15
- package/skills/xlsx/SKILL.md +0 -110
- package/skills/zarr-python/SKILL.md +0 -241
|
@@ -1,458 +0,0 @@
|
|
|
1
|
-
---
|
|
2
|
-
name: torch-geometric
|
|
3
|
-
description: PyTorch Geometric (PyG) for graph neural networks — node/link/graph classification, message passing (GCN, GAT, GraphSAGE, GIN), heterogeneous graphs, neighbor sampling, and custom datasets. Use when working with torch_geometric, not for general NetworkX analytics or non-graph PyTorch models.
|
|
4
|
-
license: MIT license
|
|
5
|
-
compatibility: Requires Python 3.10+, PyTorch 2.6+, and torch-geometric 2.7.x. Optional extension wheels (pyg-lib, torch-scatter, torch-sparse, torch-cluster) must match your PyTorch/CUDA build from https://data.pyg.org/whl.
|
|
6
|
-
metadata:
|
|
7
|
-
version: "1.1"
|
|
8
|
-
skill-author: K-Dense Inc.
|
|
9
|
-
---
|
|
10
|
-
|
|
11
|
-
# PyTorch Geometric (PyG)
|
|
12
|
-
|
|
13
|
-
PyG is the standard library for Graph Neural Networks built on PyTorch. It provides data structures for graphs, 60+ GNN layer implementations, scalable mini-batch training, and support for heterogeneous graphs.
|
|
14
|
-
|
|
15
|
-
## Installation
|
|
16
|
-
|
|
17
|
-
Tested against **torch-geometric 2.7.x** (Oct 2025). Requires **Python 3.10+** and **PyTorch 2.6+**.
|
|
18
|
-
|
|
19
|
-
```bash
|
|
20
|
-
# 1. Install PyTorch first (match your CUDA/CPU setup — see https://pytorch.org/get-started/locally/)
|
|
21
|
-
uv pip install torch
|
|
22
|
-
|
|
23
|
-
# 2. Core PyG (no extension wheels required for basic usage)
|
|
24
|
-
uv pip install torch_geometric
|
|
25
|
-
```
|
|
26
|
-
|
|
27
|
-
Optional accelerated ops (`pyg-lib`, `torch-scatter`, `torch-sparse`, `torch-cluster`) are **not required** for basic PyG usage (since PyG 2.3). Install version-matched wheels from the [PyG wheel index](https://data.pyg.org/whl) after checking your PyTorch and CUDA versions:
|
|
28
|
-
|
|
29
|
-
```bash
|
|
30
|
-
python -c "import torch; print(torch.__version__, torch.version.cuda)"
|
|
31
|
-
# Then install wheels for your torch+CUDA combo, e.g.:
|
|
32
|
-
uv pip install pyg-lib torch-scatter torch-sparse torch-cluster \
|
|
33
|
-
-f https://data.pyg.org/whl/torch-2.8.0+cu128.html
|
|
34
|
-
```
|
|
35
|
-
|
|
36
|
-
Check your version:
|
|
37
|
-
|
|
38
|
-
```python
|
|
39
|
-
import torch_geometric
|
|
40
|
-
print(torch_geometric.__version__)
|
|
41
|
-
```
|
|
42
|
-
|
|
43
|
-
**Conda:** the `pyg` conda channel is no longer maintained for PyTorch >2.5 — use `uv pip install` and the wheel index above instead.
|
|
44
|
-
|
|
45
|
-
### PyG 2.7 notes
|
|
46
|
-
|
|
47
|
-
PyG 2.7 dropped Python 3.9 and PyTorch ≤2.5. See the [2.7.0 release notes](https://github.com/pyg-team/pytorch_geometric/releases/tag/2.7.0) for PyTorch 2.6–2.8 compatibility tables. `torch_geometric.distributed` is deprecated — use standard `torch.distributed` DDP (see `references/scaling.md`).
|
|
48
|
-
|
|
49
|
-
## Core Concepts
|
|
50
|
-
|
|
51
|
-
### Graph Data: `Data` and `HeteroData`
|
|
52
|
-
|
|
53
|
-
A graph lives in a `Data` object. The key attributes:
|
|
54
|
-
|
|
55
|
-
```python
|
|
56
|
-
from torch_geometric.data import Data
|
|
57
|
-
|
|
58
|
-
data = Data(
|
|
59
|
-
x=node_features, # [num_nodes, num_node_features]
|
|
60
|
-
edge_index=edge_index, # [2, num_edges] — COO format, dtype=torch.long
|
|
61
|
-
edge_attr=edge_features, # [num_edges, num_edge_features]
|
|
62
|
-
y=labels, # node-level [num_nodes, *] or graph-level [1, *]
|
|
63
|
-
pos=positions, # [num_nodes, num_dimensions] (for point clouds/spatial)
|
|
64
|
-
)
|
|
65
|
-
```
|
|
66
|
-
|
|
67
|
-
**`edge_index` format is critical**: it's a `[2, num_edges]` tensor where `edge_index[0]` = source nodes, `edge_index[1]` = target nodes. It is NOT a list of tuples. If you have edge pairs as rows, transpose and call `.contiguous()`:
|
|
68
|
-
|
|
69
|
-
```python
|
|
70
|
-
# If edges are [[src1, dst1], [src2, dst2], ...] — transpose first:
|
|
71
|
-
edge_index = edge_pairs.t().contiguous()
|
|
72
|
-
```
|
|
73
|
-
|
|
74
|
-
For undirected graphs, include both directions: edge (0,1) needs both `[0,1]` and `[1,0]` in edge_index.
|
|
75
|
-
|
|
76
|
-
For heterogeneous graphs, use `HeteroData` — see the Heterogeneous Graphs section below.
|
|
77
|
-
|
|
78
|
-
### Datasets
|
|
79
|
-
|
|
80
|
-
PyG bundles many standard datasets that auto-download and preprocess:
|
|
81
|
-
|
|
82
|
-
```python
|
|
83
|
-
from torch_geometric.datasets import Planetoid, TUDataset
|
|
84
|
-
|
|
85
|
-
# Single-graph node classification (Cora, Citeseer, Pubmed)
|
|
86
|
-
dataset = Planetoid(root='./data', name='Cora')
|
|
87
|
-
data = dataset[0] # single graph with train/val/test masks
|
|
88
|
-
|
|
89
|
-
# Multi-graph classification (ENZYMES, MUTAG, IMDB-BINARY, etc.)
|
|
90
|
-
dataset = TUDataset(root='./data', name='ENZYMES')
|
|
91
|
-
# dataset[0], dataset[1], ... are individual graphs
|
|
92
|
-
```
|
|
93
|
-
|
|
94
|
-
Common datasets by task:
|
|
95
|
-
- **Node classification**: Planetoid (Cora/Citeseer/Pubmed), OGB (ogbn-arxiv, ogbn-products, ogbn-mag)
|
|
96
|
-
- **Graph classification**: TUDataset (MUTAG, ENZYMES, PROTEINS, IMDB-BINARY), OGB (ogbg-molhiv)
|
|
97
|
-
- **Link prediction**: OGB (ogbl-collab, ogbl-citation2)
|
|
98
|
-
- **Molecular**: QM7, QM9, MoleculeNet
|
|
99
|
-
- **Point cloud/mesh**: ShapeNet, ModelNet10/40, FAUST
|
|
100
|
-
|
|
101
|
-
### Transforms
|
|
102
|
-
|
|
103
|
-
Transforms preprocess or augment graph data, analogous to torchvision transforms:
|
|
104
|
-
|
|
105
|
-
```python
|
|
106
|
-
import torch_geometric.transforms as T
|
|
107
|
-
|
|
108
|
-
# Common transforms
|
|
109
|
-
T.NormalizeFeatures() # Row-normalize node features to sum to 1
|
|
110
|
-
T.ToUndirected() # Add reverse edges to make graph undirected
|
|
111
|
-
T.AddSelfLoops() # Add self-loop edges
|
|
112
|
-
T.KNNGraph(k=6) # Build k-NN graph from point cloud positions
|
|
113
|
-
T.RandomJitter(0.01) # Random noise augmentation on positions
|
|
114
|
-
T.Compose([...]) # Chain multiple transforms
|
|
115
|
-
|
|
116
|
-
# Apply as pre_transform (once, saved to disk) or transform (every access)
|
|
117
|
-
dataset = ShapeNet(root='./data', pre_transform=T.KNNGraph(k=6),
|
|
118
|
-
transform=T.RandomJitter(0.01))
|
|
119
|
-
```
|
|
120
|
-
|
|
121
|
-
## Building GNN Models
|
|
122
|
-
|
|
123
|
-
### Quick Start: Using Built-in Layers
|
|
124
|
-
|
|
125
|
-
The fastest way to build a GNN — stack conv layers from `torch_geometric.nn`:
|
|
126
|
-
|
|
127
|
-
```python
|
|
128
|
-
import torch
|
|
129
|
-
import torch.nn.functional as F
|
|
130
|
-
from torch_geometric.nn import GCNConv
|
|
131
|
-
|
|
132
|
-
class GCN(torch.nn.Module):
|
|
133
|
-
def __init__(self, in_channels, hidden_channels, out_channels):
|
|
134
|
-
super().__init__()
|
|
135
|
-
self.conv1 = GCNConv(in_channels, hidden_channels)
|
|
136
|
-
self.conv2 = GCNConv(hidden_channels, out_channels)
|
|
137
|
-
|
|
138
|
-
def forward(self, x, edge_index):
|
|
139
|
-
x = self.conv1(x, edge_index).relu()
|
|
140
|
-
x = F.dropout(x, p=0.5, training=self.training)
|
|
141
|
-
x = self.conv2(x, edge_index)
|
|
142
|
-
return x
|
|
143
|
-
```
|
|
144
|
-
|
|
145
|
-
**Important**: PyG conv layers do NOT include activation functions — apply them yourself after each layer. This is by design for flexibility.
|
|
146
|
-
|
|
147
|
-
### Choosing a Conv Layer
|
|
148
|
-
|
|
149
|
-
Pick based on your task and graph structure:
|
|
150
|
-
|
|
151
|
-
| Layer | Best for | Key idea |
|
|
152
|
-
|-------|----------|----------|
|
|
153
|
-
| `GCNConv` | Homogeneous, semi-supervised node classification | Spectral-inspired, degree-normalized aggregation |
|
|
154
|
-
| `GATConv` / `GATv2Conv` | When neighbor importance varies | Attention-weighted messages |
|
|
155
|
-
| `SAGEConv` | Large graphs, inductive settings | Sampling-friendly, learnable aggregation |
|
|
156
|
-
| `GINConv` | Graph classification, maximizing expressiveness | As powerful as WL test |
|
|
157
|
-
| `TransformerConv` | Rich edge features, complex interactions | Multi-head attention with edge features |
|
|
158
|
-
| `EdgeConv` | Point clouds, dynamic graphs | MLP on edge features (x_i, x_j - x_i) |
|
|
159
|
-
| `RGCNConv` | Heterogeneous with many relation types | Relation-specific weight matrices |
|
|
160
|
-
| `HGTConv` | Heterogeneous graphs | Type-specific attention |
|
|
161
|
-
|
|
162
|
-
All conv layers accept `(x, edge_index)` at minimum. Many also accept `edge_attr` for edge features.
|
|
163
|
-
|
|
164
|
-
### Lazy Initialization
|
|
165
|
-
|
|
166
|
-
Use `-1` for input channels to let PyG infer dimensions automatically — especially useful for heterogeneous models:
|
|
167
|
-
|
|
168
|
-
```python
|
|
169
|
-
conv = SAGEConv((-1, -1), 64) # Input dims inferred on first forward pass
|
|
170
|
-
# Initialize lazy modules:
|
|
171
|
-
with torch.no_grad():
|
|
172
|
-
out = model(data.x, data.edge_index)
|
|
173
|
-
```
|
|
174
|
-
|
|
175
|
-
### High-Level Model APIs
|
|
176
|
-
|
|
177
|
-
For common architectures, PyG provides ready-made model classes:
|
|
178
|
-
|
|
179
|
-
```python
|
|
180
|
-
from torch_geometric.nn import GraphSAGE, GCN, GAT, GIN
|
|
181
|
-
|
|
182
|
-
model = GraphSAGE(
|
|
183
|
-
in_channels=dataset.num_features,
|
|
184
|
-
hidden_channels=64,
|
|
185
|
-
out_channels=dataset.num_classes,
|
|
186
|
-
num_layers=2,
|
|
187
|
-
)
|
|
188
|
-
```
|
|
189
|
-
|
|
190
|
-
### Custom Layers via MessagePassing
|
|
191
|
-
|
|
192
|
-
To implement a novel GNN layer, subclass `MessagePassing`. The framework is:
|
|
193
|
-
|
|
194
|
-
1. `propagate()` orchestrates the message passing
|
|
195
|
-
2. `message()` defines what info flows along each edge (the phi function)
|
|
196
|
-
3. `aggregate()` combines messages at each node (sum/mean/max)
|
|
197
|
-
4. `update()` transforms the aggregated result (the gamma function)
|
|
198
|
-
|
|
199
|
-
```python
|
|
200
|
-
from torch_geometric.nn import MessagePassing
|
|
201
|
-
from torch_geometric.utils import add_self_loops, degree
|
|
202
|
-
|
|
203
|
-
class MyConv(MessagePassing):
|
|
204
|
-
def __init__(self, in_channels, out_channels):
|
|
205
|
-
super().__init__(aggr='add') # "add", "mean", or "max"
|
|
206
|
-
self.lin = torch.nn.Linear(in_channels, out_channels)
|
|
207
|
-
|
|
208
|
-
def forward(self, x, edge_index):
|
|
209
|
-
# Pre-processing before message passing
|
|
210
|
-
x = self.lin(x)
|
|
211
|
-
# Start message passing
|
|
212
|
-
return self.propagate(edge_index, x=x)
|
|
213
|
-
|
|
214
|
-
def message(self, x_j):
|
|
215
|
-
# x_j: features of source nodes for each edge [num_edges, features]
|
|
216
|
-
# The _j suffix auto-indexes source nodes, _i indexes target nodes
|
|
217
|
-
return x_j
|
|
218
|
-
```
|
|
219
|
-
|
|
220
|
-
**The `_i` / `_j` convention**: any tensor passed to `propagate()` can be auto-indexed by appending `_i` (target/central node) or `_j` (source/neighbor node) in the `message()` signature. So if you pass `x=...` to propagate, you can access `x_i` and `x_j` in message().
|
|
221
|
-
|
|
222
|
-
Read `references/message_passing.md` for the full GCN and EdgeConv implementation examples.
|
|
223
|
-
|
|
224
|
-
## Task-Specific Patterns
|
|
225
|
-
|
|
226
|
-
### Node Classification
|
|
227
|
-
|
|
228
|
-
```python
|
|
229
|
-
# Full-batch training on a single graph (e.g., Cora)
|
|
230
|
-
model.train()
|
|
231
|
-
for epoch in range(200):
|
|
232
|
-
optimizer.zero_grad()
|
|
233
|
-
out = model(data.x, data.edge_index)
|
|
234
|
-
loss = F.cross_entropy(out[data.train_mask], data.y[data.train_mask])
|
|
235
|
-
loss.backward()
|
|
236
|
-
optimizer.step()
|
|
237
|
-
|
|
238
|
-
# Evaluation — train(False) puts the model in inference mode (disables dropout/BN)
|
|
239
|
-
model.train(False)
|
|
240
|
-
pred = model(data.x, data.edge_index).argmax(dim=1)
|
|
241
|
-
acc = (pred[data.test_mask] == data.y[data.test_mask]).float().mean()
|
|
242
|
-
```
|
|
243
|
-
|
|
244
|
-
### Graph Classification
|
|
245
|
-
|
|
246
|
-
Multiple graphs — use `DataLoader` for mini-batching and global pooling to get graph-level representations:
|
|
247
|
-
|
|
248
|
-
```python
|
|
249
|
-
from torch_geometric.loader import DataLoader
|
|
250
|
-
from torch_geometric.nn import GCNConv, global_mean_pool
|
|
251
|
-
|
|
252
|
-
loader = DataLoader(dataset, batch_size=32, shuffle=True)
|
|
253
|
-
|
|
254
|
-
class GraphClassifier(torch.nn.Module):
|
|
255
|
-
def __init__(self, in_ch, hidden_ch, out_ch):
|
|
256
|
-
super().__init__()
|
|
257
|
-
self.conv1 = GCNConv(in_ch, hidden_ch)
|
|
258
|
-
self.conv2 = GCNConv(hidden_ch, hidden_ch)
|
|
259
|
-
self.lin = torch.nn.Linear(hidden_ch, out_ch)
|
|
260
|
-
|
|
261
|
-
def forward(self, x, edge_index, batch):
|
|
262
|
-
x = self.conv1(x, edge_index).relu()
|
|
263
|
-
x = self.conv2(x, edge_index).relu()
|
|
264
|
-
x = global_mean_pool(x, batch) # [num_graphs_in_batch, hidden_ch]
|
|
265
|
-
return self.lin(x)
|
|
266
|
-
|
|
267
|
-
# Training loop
|
|
268
|
-
for data in loader:
|
|
269
|
-
out = model(data.x, data.edge_index, data.batch)
|
|
270
|
-
loss = F.cross_entropy(out, data.y)
|
|
271
|
-
```
|
|
272
|
-
|
|
273
|
-
PyG's `DataLoader` batches multiple graphs by creating block-diagonal adjacency matrices. The `batch` tensor maps each node to its graph index. Pooling ops (`global_mean_pool`, `global_max_pool`, `global_add_pool`) use this to aggregate per-graph.
|
|
274
|
-
|
|
275
|
-
### Link Prediction
|
|
276
|
-
|
|
277
|
-
Split edges into train/val/test, use negative sampling:
|
|
278
|
-
|
|
279
|
-
```python
|
|
280
|
-
from torch_geometric.transforms import RandomLinkSplit
|
|
281
|
-
|
|
282
|
-
transform = RandomLinkSplit(
|
|
283
|
-
num_val=0.1,
|
|
284
|
-
num_test=0.1,
|
|
285
|
-
is_undirected=True,
|
|
286
|
-
add_negative_train_samples=False,
|
|
287
|
-
)
|
|
288
|
-
train_data, val_data, test_data = transform(data)
|
|
289
|
-
|
|
290
|
-
# Encode nodes, then score edges
|
|
291
|
-
z = model.encode(train_data.x, train_data.edge_index)
|
|
292
|
-
# Positive edges
|
|
293
|
-
pos_score = (z[train_data.edge_label_index[0]] * z[train_data.edge_label_index[1]]).sum(dim=1)
|
|
294
|
-
```
|
|
295
|
-
|
|
296
|
-
Read `references/link_prediction.md` for the complete link prediction guide: GAE/VGAE autoencoders, full training loops, LinkNeighborLoader for large graphs, heterogeneous link prediction, and evaluation metrics.
|
|
297
|
-
|
|
298
|
-
## Scaling to Large Graphs
|
|
299
|
-
|
|
300
|
-
For graphs that don't fit in GPU memory, use neighbor sampling via `NeighborLoader`:
|
|
301
|
-
|
|
302
|
-
```python
|
|
303
|
-
from torch_geometric.loader import NeighborLoader
|
|
304
|
-
|
|
305
|
-
train_loader = NeighborLoader(
|
|
306
|
-
data,
|
|
307
|
-
num_neighbors=[15, 10], # Sample 15 neighbors in hop 1, 10 in hop 2
|
|
308
|
-
batch_size=128, # Number of seed nodes per batch
|
|
309
|
-
input_nodes=data.train_mask, # Which nodes to sample from
|
|
310
|
-
shuffle=True,
|
|
311
|
-
)
|
|
312
|
-
|
|
313
|
-
for batch in train_loader:
|
|
314
|
-
batch = batch.to(device)
|
|
315
|
-
out = model(batch.x, batch.edge_index)
|
|
316
|
-
# Only use first batch_size nodes for loss (these are the seed nodes)
|
|
317
|
-
loss = F.cross_entropy(out[:batch.batch_size], batch.y[:batch.batch_size])
|
|
318
|
-
```
|
|
319
|
-
|
|
320
|
-
**Key points about NeighborLoader**:
|
|
321
|
-
- `num_neighbors` list length should match GNN depth (number of message passing layers)
|
|
322
|
-
- Seed nodes are always the first `batch.batch_size` nodes in the output
|
|
323
|
-
- `batch.n_id` maps relabeled indices back to original node IDs
|
|
324
|
-
- Works for both `Data` and `HeteroData`
|
|
325
|
-
- For link prediction, use `LinkNeighborLoader` instead
|
|
326
|
-
- Sampling more than 2-3 hops is generally infeasible (exponential blowup)
|
|
327
|
-
|
|
328
|
-
Other scalability options: `ClusterLoader` (ClusterGCN), `GraphSAINTSampler`, `ShaDowKHopSampler`. For multi-GPU training, DDP, PyTorch Lightning integration, and `torch.compile` support, read `references/scaling.md`.
|
|
329
|
-
|
|
330
|
-
## Heterogeneous Graphs
|
|
331
|
-
|
|
332
|
-
For graphs with multiple node and edge types (social networks, knowledge graphs, recommendation):
|
|
333
|
-
|
|
334
|
-
```python
|
|
335
|
-
from torch_geometric.data import HeteroData
|
|
336
|
-
|
|
337
|
-
data = HeteroData()
|
|
338
|
-
|
|
339
|
-
# Node features — indexed by node type string
|
|
340
|
-
data['user'].x = torch.randn(1000, 64)
|
|
341
|
-
data['movie'].x = torch.randn(500, 128)
|
|
342
|
-
|
|
343
|
-
# Edge indices — indexed by (src_type, edge_type, dst_type) triplet
|
|
344
|
-
data['user', 'rates', 'movie'].edge_index = torch.randint(0, 500, (2, 3000))
|
|
345
|
-
data['user', 'follows', 'user'].edge_index = torch.randint(0, 1000, (2, 5000))
|
|
346
|
-
|
|
347
|
-
# Access convenience dicts
|
|
348
|
-
data.x_dict # {'user': tensor, 'movie': tensor}
|
|
349
|
-
data.edge_index_dict # {('user','rates','movie'): tensor, ...}
|
|
350
|
-
data.metadata() # ([node_types], [edge_types])
|
|
351
|
-
```
|
|
352
|
-
|
|
353
|
-
### Three ways to build heterogeneous GNNs
|
|
354
|
-
|
|
355
|
-
**1. Auto-convert with `to_hetero()`** — write a homogeneous model, convert automatically:
|
|
356
|
-
|
|
357
|
-
```python
|
|
358
|
-
from torch_geometric.nn import SAGEConv, to_hetero
|
|
359
|
-
|
|
360
|
-
class GNN(torch.nn.Module):
|
|
361
|
-
def __init__(self, hidden_channels, out_channels):
|
|
362
|
-
super().__init__()
|
|
363
|
-
self.conv1 = SAGEConv((-1, -1), hidden_channels)
|
|
364
|
-
self.conv2 = SAGEConv((-1, -1), out_channels)
|
|
365
|
-
|
|
366
|
-
def forward(self, x, edge_index):
|
|
367
|
-
x = self.conv1(x, edge_index).relu()
|
|
368
|
-
x = self.conv2(x, edge_index)
|
|
369
|
-
return x
|
|
370
|
-
|
|
371
|
-
model = GNN(64, dataset.num_classes)
|
|
372
|
-
model = to_hetero(model, data.metadata(), aggr='sum')
|
|
373
|
-
|
|
374
|
-
# Now accepts dicts:
|
|
375
|
-
out = model(data.x_dict, data.edge_index_dict)
|
|
376
|
-
```
|
|
377
|
-
|
|
378
|
-
Use `(-1, -1)` for bipartite input channels (source, target may differ). Lazy init handles the rest.
|
|
379
|
-
|
|
380
|
-
**2. `HeteroConv` wrapper** — different conv per edge type:
|
|
381
|
-
|
|
382
|
-
```python
|
|
383
|
-
from torch_geometric.nn import HeteroConv, GCNConv, SAGEConv, GATConv
|
|
384
|
-
|
|
385
|
-
conv = HeteroConv({
|
|
386
|
-
('paper', 'cites', 'paper'): GCNConv(-1, 64),
|
|
387
|
-
('author', 'writes', 'paper'): SAGEConv((-1, -1), 64),
|
|
388
|
-
('paper', 'rev_writes', 'author'): GATConv((-1, -1), 64, add_self_loops=False),
|
|
389
|
-
}, aggr='sum')
|
|
390
|
-
```
|
|
391
|
-
|
|
392
|
-
**3. Native heterogeneous operators** like `HGTConv`:
|
|
393
|
-
|
|
394
|
-
```python
|
|
395
|
-
from torch_geometric.nn import HGTConv
|
|
396
|
-
conv = HGTConv(hidden_channels, hidden_channels, data.metadata(), num_heads=4)
|
|
397
|
-
```
|
|
398
|
-
|
|
399
|
-
**Important for heterogeneous graphs**:
|
|
400
|
-
- Use `T.ToUndirected()` to add reverse edge types for bidirectional message flow
|
|
401
|
-
- Disable `add_self_loops` in bipartite conv layers (different source/dest types) — use skip connections instead: `conv(x, edge_index) + lin(x)`
|
|
402
|
-
- For NeighborLoader on HeteroData, specify `input_nodes` as `('node_type', mask)` tuple
|
|
403
|
-
- `num_neighbors` can be a dict keyed by edge type for fine-grained control
|
|
404
|
-
|
|
405
|
-
Read `references/heterogeneous.md` for complete examples including training loops and NeighborLoader usage with heterogeneous graphs.
|
|
406
|
-
|
|
407
|
-
## Custom Datasets
|
|
408
|
-
|
|
409
|
-
For loading your own data into PyG:
|
|
410
|
-
|
|
411
|
-
- **Quick (no class needed)**: Create `Data` objects directly and pass a list to `DataLoader`
|
|
412
|
-
- **Reusable (fits in RAM)**: Subclass `InMemoryDataset` — override `raw_file_names`, `processed_file_names`, `download()`, `process()`
|
|
413
|
-
- **Large (disk-backed)**: Subclass `Dataset` — also override `len()` and `get()`
|
|
414
|
-
- **From CSV**: Load node/edge tables with pandas, build mappings to consecutive indices, assemble into `Data` or `HeteroData`
|
|
415
|
-
- **From NetworkX**: `from_networkx(G)` converts a NetworkX graph directly
|
|
416
|
-
- **From scipy sparse**: `from_scipy_sparse_matrix(adj)` extracts edge_index
|
|
417
|
-
|
|
418
|
-
Read `references/custom_datasets.md` for complete examples with all patterns, CSV loading with encoders, and the MovieLens walkthrough.
|
|
419
|
-
|
|
420
|
-
## Explainability
|
|
421
|
-
|
|
422
|
-
PyG provides `torch_geometric.explain` for interpreting GNN predictions:
|
|
423
|
-
|
|
424
|
-
```python
|
|
425
|
-
from torch_geometric.explain import Explainer, GNNExplainer
|
|
426
|
-
|
|
427
|
-
explainer = Explainer(
|
|
428
|
-
model=model,
|
|
429
|
-
algorithm=GNNExplainer(epochs=200),
|
|
430
|
-
explanation_type='model',
|
|
431
|
-
node_mask_type='attributes',
|
|
432
|
-
edge_mask_type='object',
|
|
433
|
-
model_config=dict(
|
|
434
|
-
mode='multiclass_classification',
|
|
435
|
-
task_level='node',
|
|
436
|
-
return_type='log_probs',
|
|
437
|
-
),
|
|
438
|
-
)
|
|
439
|
-
|
|
440
|
-
explanation = explainer(data.x, data.edge_index, index=10)
|
|
441
|
-
explanation.visualize_graph() # Important subgraph
|
|
442
|
-
explanation.visualize_feature_importance(top_k=10) # Feature importance
|
|
443
|
-
```
|
|
444
|
-
|
|
445
|
-
Available algorithms: `GNNExplainer` (optimization-based), `PGExplainer` (parametric, trained), `CaptumExplainer` (gradient-based via Captum), `AttentionExplainer` (attention weights). Works for both homogeneous and heterogeneous graphs.
|
|
446
|
-
|
|
447
|
-
Read `references/explainability.md` for all algorithms, heterogeneous explanations, evaluation metrics, and PGExplainer training.
|
|
448
|
-
|
|
449
|
-
## Common Pitfalls
|
|
450
|
-
|
|
451
|
-
1. **edge_index shape**: Must be `[2, num_edges]`, not `[num_edges, 2]`. Transpose if needed.
|
|
452
|
-
2. **Forgetting activations**: Conv layers don't include ReLU/etc — add them manually.
|
|
453
|
-
3. **Self-loops in hetero bipartite**: Don't use `add_self_loops=True` when source and dest node types differ. Use skip connections instead.
|
|
454
|
-
4. **NeighborLoader slicing**: Only the first `batch.batch_size` nodes are your seed nodes. Slice predictions and labels accordingly.
|
|
455
|
-
5. **Undirected graphs**: If your graph is undirected, include edges in both directions in `edge_index`, or use `T.ToUndirected()`.
|
|
456
|
-
6. **Lazy init**: Models with `-1` input channels need one forward pass with `torch.no_grad()` before training to initialize parameters.
|
|
457
|
-
7. **Global pooling for graph tasks**: Use `global_mean_pool(x, batch)` (not manual reshape) to aggregate node features to graph-level.
|
|
458
|
-
8. **num_neighbors alignment**: Keep `len(num_neighbors)` equal to the number of GNN layers. More hops than layers wastes compute; fewer means wasted model capacity.
|
|
@@ -1,241 +0,0 @@
|
|
|
1
|
-
---
|
|
2
|
-
name: torchdrug
|
|
3
|
-
description: Build and troubleshoot TorchDrug 0.2.1 workflows for molecular graphs, property prediction, self-supervised pretraining, molecule generation, retrosynthesis, protein representation learning, and knowledge graph reasoning. Use when code imports torchdrug or needs its datasets, models, tasks, or Engine.
|
|
4
|
-
license: Apache-2.0 license
|
|
5
|
-
compatibility: TorchDrug 0.2.1 requires Python 3.7-3.10 and supports PyTorch 1.8-2.0. Apple Silicon is CPU-only; MPS is unsupported.
|
|
6
|
-
allowed-tools: Read Write Edit Bash
|
|
7
|
-
metadata:
|
|
8
|
-
version: "1.1"
|
|
9
|
-
skill-author: K-Dense Inc.
|
|
10
|
-
---
|
|
11
|
-
|
|
12
|
-
# TorchDrug
|
|
13
|
-
|
|
14
|
-
Use TorchDrug as a modular PyTorch graph-learning stack:
|
|
15
|
-
|
|
16
|
-
1. load a `datasets.*` dataset,
|
|
17
|
-
2. choose a `models.*` representation model,
|
|
18
|
-
3. wrap it in a `tasks.*` objective,
|
|
19
|
-
4. train and evaluate it with `core.Engine`.
|
|
20
|
-
|
|
21
|
-
The current official documentation and latest release are both **0.2.1**. Treat
|
|
22
|
-
newer Python or PyTorch combinations as unverified rather than silently assuming
|
|
23
|
-
compatibility.
|
|
24
|
-
|
|
25
|
-
## Start with the version guard
|
|
26
|
-
|
|
27
|
-
Before generating or debugging code, inspect the environment:
|
|
28
|
-
|
|
29
|
-
```bash
|
|
30
|
-
python --version
|
|
31
|
-
python -c "import torch; print(torch.__version__)"
|
|
32
|
-
python -c "import torchdrug; print(torchdrug.__version__)"
|
|
33
|
-
```
|
|
34
|
-
|
|
35
|
-
The supported matrix for TorchDrug 0.2.1 is:
|
|
36
|
-
|
|
37
|
-
- Python 3.7 through 3.10
|
|
38
|
-
- PyTorch 1.8 through 2.0
|
|
39
|
-
- Linux, Windows, or macOS
|
|
40
|
-
- Apple Silicon: PyTorch 1.13 or later, CPU only; no MPS support
|
|
41
|
-
|
|
42
|
-
If the project uses Python 3.11+ or PyTorch 2.1+, create a compatible environment
|
|
43
|
-
or explicitly test a source build. Do not present such combinations as supported.
|
|
44
|
-
|
|
45
|
-
## Installation
|
|
46
|
-
|
|
47
|
-
Prefer a dedicated Python 3.10 environment and pin the TorchDrug release:
|
|
48
|
-
|
|
49
|
-
```bash
|
|
50
|
-
uv venv --python 3.10
|
|
51
|
-
source .venv/bin/activate
|
|
52
|
-
uv pip install "torch==2.0.0"
|
|
53
|
-
```
|
|
54
|
-
|
|
55
|
-
Install `torch-scatter` and `torch-cluster` wheels matched to the exact PyTorch
|
|
56
|
-
and CUDA pair, following the
|
|
57
|
-
[official installation page](https://torchdrug.ai/docs/installation.html). For a
|
|
58
|
-
CPU-only PyTorch 2.0 environment, one reproducible wheel combination is:
|
|
59
|
-
|
|
60
|
-
```bash
|
|
61
|
-
uv pip install "torch-scatter==2.1.1" "torch-cluster==1.6.1" \
|
|
62
|
-
--find-links "https://data.pyg.org/whl/torch-2.0.0+cpu.html"
|
|
63
|
-
uv pip install "torchdrug==0.2.1"
|
|
64
|
-
```
|
|
65
|
-
|
|
66
|
-
Do not copy a CUDA wheel URL between environments. Match the PyTorch version,
|
|
67
|
-
CUDA build, Python ABI, and platform. On Apple Silicon, the official docs require
|
|
68
|
-
building `torch-scatter` and `torch-cluster` from source; pin reviewed source
|
|
69
|
-
revisions and expect CPU execution.
|
|
70
|
-
|
|
71
|
-
## Canonical property-prediction workflow
|
|
72
|
-
|
|
73
|
-
Use the documented ClinTox → GIN → `PropertyPrediction` → `Engine` pattern:
|
|
74
|
-
|
|
75
|
-
```python
|
|
76
|
-
import torch
|
|
77
|
-
from torchdrug import core, datasets, models, tasks
|
|
78
|
-
|
|
79
|
-
dataset = datasets.ClinTox("~/molecule-datasets/")
|
|
80
|
-
lengths = [int(0.8 * len(dataset)), int(0.1 * len(dataset))]
|
|
81
|
-
lengths.append(len(dataset) - sum(lengths))
|
|
82
|
-
train_set, valid_set, test_set = torch.utils.data.random_split(dataset, lengths)
|
|
83
|
-
|
|
84
|
-
model = models.GIN(
|
|
85
|
-
input_dim=dataset.node_feature_dim,
|
|
86
|
-
hidden_dims=[256, 256, 256, 256],
|
|
87
|
-
short_cut=True,
|
|
88
|
-
batch_norm=True,
|
|
89
|
-
concat_hidden=True,
|
|
90
|
-
)
|
|
91
|
-
task = tasks.PropertyPrediction(
|
|
92
|
-
model,
|
|
93
|
-
task=dataset.tasks,
|
|
94
|
-
criterion="bce",
|
|
95
|
-
metric=("auprc", "auroc"),
|
|
96
|
-
)
|
|
97
|
-
|
|
98
|
-
optimizer = torch.optim.Adam(task.parameters(), lr=1e-3)
|
|
99
|
-
solver = core.Engine(
|
|
100
|
-
task,
|
|
101
|
-
train_set,
|
|
102
|
-
valid_set,
|
|
103
|
-
test_set,
|
|
104
|
-
optimizer,
|
|
105
|
-
batch_size=1024,
|
|
106
|
-
)
|
|
107
|
-
solver.train(num_epoch=100)
|
|
108
|
-
solver.evaluate("valid")
|
|
109
|
-
```
|
|
110
|
-
|
|
111
|
-
Add `gpus=[0]` only when a supported CUDA device is available. Omit `gpus` for
|
|
112
|
-
CPU execution.
|
|
113
|
-
|
|
114
|
-
For binary classification, `task.predict(batch)` returns logits; apply
|
|
115
|
-
`torch.sigmoid` when probabilities are needed. In 0.2.1, normalized regression
|
|
116
|
-
predictions are returned on the original target scale, which is a breaking change
|
|
117
|
-
from older releases.
|
|
118
|
-
|
|
119
|
-
## Choose the official workflow
|
|
120
|
-
|
|
121
|
-
### Molecular property prediction
|
|
122
|
-
|
|
123
|
-
- Dataset: `datasets.ClinTox`, `BBBP`, `Tox21`, `QM9`, or another documented
|
|
124
|
-
molecule dataset.
|
|
125
|
-
- Model: start with `models.GIN`; use `edge_input_dim` when the selected feature
|
|
126
|
-
configuration supplies edge features.
|
|
127
|
-
- Task: `tasks.PropertyPrediction`.
|
|
128
|
-
- Read [molecular property prediction](references/molecular_property_prediction.md).
|
|
129
|
-
|
|
130
|
-
### Self-supervised molecular pretraining
|
|
131
|
-
|
|
132
|
-
- InfoGraph: `models.InfoGraph(gin_model, separate_model=False)` wrapped by
|
|
133
|
-
`tasks.Unsupervised`.
|
|
134
|
-
- Attribute masking: `tasks.AttributeMasking(model, mask_rate=0.15)`.
|
|
135
|
-
- Recreate the same encoder for fine-tuning, then load the checkpoint with
|
|
136
|
-
`strict=False` before training `tasks.PropertyPrediction`.
|
|
137
|
-
- Read [molecular property prediction](references/molecular_property_prediction.md).
|
|
138
|
-
|
|
139
|
-
### Molecule generation
|
|
140
|
-
|
|
141
|
-
- Dataset: `datasets.ZINC250k(..., kekulize=True, atom_feature="symbol")`.
|
|
142
|
-
- GCPN: an `models.RGCN` encoder wrapped by `tasks.GCPNGeneration`.
|
|
143
|
-
- GraphAF: node and edge `models.GraphAF` flows wrapped by
|
|
144
|
-
`tasks.AutoregressiveGeneration`.
|
|
145
|
-
- Supported optimization tasks in the tutorial are `"qed"` and `"plogp"`;
|
|
146
|
-
criteria are `"nll"` and/or `"ppo"`.
|
|
147
|
-
- Read [molecular generation](references/molecular_generation.md).
|
|
148
|
-
|
|
149
|
-
### Retrosynthesis
|
|
150
|
-
|
|
151
|
-
- Create two synchronized `datasets.USPTO50k` views: reaction mode for center
|
|
152
|
-
identification and `as_synthon=True` for synthon completion.
|
|
153
|
-
- Train `tasks.CenterIdentification` and `tasks.SynthonCompletion` separately.
|
|
154
|
-
- Combine the trained tasks with `tasks.Retrosynthesis`; do not pass raw models
|
|
155
|
-
directly to the end-to-end task.
|
|
156
|
-
- Read [retrosynthesis](references/retrosynthesis.md).
|
|
157
|
-
|
|
158
|
-
### Knowledge graph reasoning
|
|
159
|
-
|
|
160
|
-
- Embedding workflow: `datasets.FB15k237` → `models.RotatE` →
|
|
161
|
-
`tasks.KnowledgeGraphCompletion`.
|
|
162
|
-
- Neural reasoning workflow: `models.NeuralLP` with `fact_ratio=0.75`.
|
|
163
|
-
- Read [knowledge graph reasoning](references/knowledge_graphs.md).
|
|
164
|
-
|
|
165
|
-
### Protein modeling
|
|
166
|
-
|
|
167
|
-
- Build proteins with `data.Protein.from_sequence`, `from_pdb`, or
|
|
168
|
-
`from_molecule`.
|
|
169
|
-
- Sequence encoders include `models.ESM`, `ProteinCNN`, `ProteinResNet`,
|
|
170
|
-
`ProteinLSTM`, and `ProteinBERT`; structure encoders include `models.GearNet`.
|
|
171
|
-
- Use documented graph-construction layers rather than a nonexistent
|
|
172
|
-
`protein.residue_graph()` convenience method.
|
|
173
|
-
- Read [protein modeling](references/protein_modeling.md).
|
|
174
|
-
|
|
175
|
-
## Rules for reliable TorchDrug code
|
|
176
|
-
|
|
177
|
-
1. **Follow the 0.2.1 API.** The official docs are not a rolling latest-version
|
|
178
|
-
site.
|
|
179
|
-
2. **Prefer documented feature names.** Use `atom_feature`, `bond_feature`,
|
|
180
|
-
`residue_feature`, and `mol_feature`; `node_feature`, `edge_feature`, and
|
|
181
|
-
`graph_feature` are deprecated aliases in relevant dataset constructors.
|
|
182
|
-
3. **Let `Engine` preprocess tasks.** If composing pre-trained tasks without
|
|
183
|
-
constructing their solvers, call each task's `preprocess()` manually.
|
|
184
|
-
4. **Keep paired splits synchronized.** For retrosynthesis, reset the same random
|
|
185
|
-
seed before splitting reaction and synthon datasets.
|
|
186
|
-
5. **Use TorchDrug collation.** Use `data.graph_collate` or `core.Engine`;
|
|
187
|
-
generic PyTorch collation does not know how to pack TorchDrug graphs.
|
|
188
|
-
6. **Separate model, task, and engine arguments.** A common source of invented
|
|
189
|
-
code is passing task options to a model or passing raw models where a composed
|
|
190
|
-
task is required.
|
|
191
|
-
7. **Validate generated chemistry.** Treat model outputs as candidates, not as
|
|
192
|
-
experimentally valid or synthesizable compounds.
|
|
193
|
-
|
|
194
|
-
## Troubleshooting
|
|
195
|
-
|
|
196
|
-
### Installation or import failure
|
|
197
|
-
|
|
198
|
-
Check Python, PyTorch, `torch-scatter`, and `torch-cluster` as one compatibility
|
|
199
|
-
set. Most failures are binary-wheel mismatches, unsupported Python versions, or
|
|
200
|
-
attempts to use MPS.
|
|
201
|
-
|
|
202
|
-
### Feature dimension mismatch
|
|
203
|
-
|
|
204
|
-
Build model dimensions from the loaded dataset:
|
|
205
|
-
|
|
206
|
-
- `dataset.node_feature_dim`
|
|
207
|
-
- `dataset.edge_feature_dim`
|
|
208
|
-
- `dataset.num_bond_type`
|
|
209
|
-
- `dataset.num_entity` and `dataset.num_relation` for knowledge graphs
|
|
210
|
-
|
|
211
|
-
Do not hard-code dimensions copied from a different feature configuration.
|
|
212
|
-
|
|
213
|
-
### Device mismatch
|
|
214
|
-
|
|
215
|
-
Pass `gpus=[0]` to `core.Engine` for supported CUDA execution. For manual
|
|
216
|
-
prediction, collate first and move the entire nested batch with `utils.cuda`.
|
|
217
|
-
|
|
218
|
-
### Checkpoint mismatch
|
|
219
|
-
|
|
220
|
-
Recreate the same model and feature configuration. For pretraining-to-fine-tuning
|
|
221
|
-
transfer, load the checkpoint's `"model"` state with `strict=False`; for a complete
|
|
222
|
-
solver, use `solver.save()` and `solver.load()`.
|
|
223
|
-
|
|
224
|
-
## Reference index
|
|
225
|
-
|
|
226
|
-
- [Core concepts and data structures](references/core_concepts.md)
|
|
227
|
-
- [Datasets](references/datasets.md)
|
|
228
|
-
- [Models and architectures](references/models_architectures.md)
|
|
229
|
-
- [Molecular property prediction and pretraining](references/molecular_property_prediction.md)
|
|
230
|
-
- [Protein modeling](references/protein_modeling.md)
|
|
231
|
-
- [Molecular generation](references/molecular_generation.md)
|
|
232
|
-
- [Retrosynthesis](references/retrosynthesis.md)
|
|
233
|
-
- [Knowledge graph reasoning](references/knowledge_graphs.md)
|
|
234
|
-
|
|
235
|
-
## Upstream sources
|
|
236
|
-
|
|
237
|
-
- [TorchDrug 0.2.1 documentation](https://torchdrug.ai/docs/)
|
|
238
|
-
- [Tutorial index](https://torchdrug.ai/docs/tutorials/)
|
|
239
|
-
- [Installation](https://torchdrug.ai/docs/installation.html)
|
|
240
|
-
- [Package reference](https://torchdrug.ai/docs/api/)
|
|
241
|
-
- [TorchDrug 0.2.1 release notes](https://github.com/DeepGraphLearning/torchdrug/releases/tag/v0.2.1)
|