---
name: k-dense-ai/torchdrug
source: https://app.decimal.ai/s/k-dense-ai-torchdrug@1/SKILL.md
source_sha256: ebfadc174dcb
---

# TorchDrug

Use TorchDrug as a modular PyTorch graph-learning stack:

1. load a `datasets.*` dataset,
2. choose a `models.*` representation model,
3. wrap it in a `tasks.*` objective,
4. train and evaluate it with `core.Engine`.

The current official documentation and latest release are both **0.2.1**. Treat
newer Python or PyTorch combinations as unverified rather than silently assuming
compatibility.

## Start with the version guard

Before generating or debugging code, inspect the environment:

```bash
python --version
python -c "import torch; print(torch.__version__)"
python -c "import torchdrug; print(torchdrug.__version__)"
```

The supported matrix for TorchDrug 0.2.1 is:

- Python 3.7 through 3.10
- PyTorch 1.8 through 2.0
- Linux, Windows, or macOS
- Apple Silicon: PyTorch 1.13 or later, CPU only; no MPS support

If the project uses Python 3.11+ or PyTorch 2.1+, create a compatible environment
or explicitly test a source build. Do not present such combinations as supported.

## Installation

Prefer a dedicated Python 3.10 environment and pin the TorchDrug release:

```bash
uv venv --python 3.10
source .venv/bin/activate
uv pip install "torch==2.0.0"
```

Install `torch-scatter` and `torch-cluster` wheels matched to the exact PyTorch
and CUDA pair, following the
[official installation page](https://torchdrug.ai/docs/installation.html). For a
CPU-only PyTorch 2.0 environment, one reproducible wheel combination is:

```bash
uv pip install "torch-scatter==2.1.1" "torch-cluster==1.6.1" \
  --find-links "https://data.pyg.org/whl/torch-2.0.0+cpu.html"
uv pip install "torchdrug==0.2.1"
```

Do not copy a CUDA wheel URL between environments. Match the PyTorch version,
CUDA build, Python ABI, and platform. On Apple Silicon, the official docs require
building `torch-scatter` and `torch-cluster` from source; pin reviewed source
revisions and expect CPU execution.

## Canonical property-prediction workflow

Use the documented ClinTox → GIN → `PropertyPrediction` → `Engine` pattern:

```python
import torch
from torchdrug import core, datasets, models, tasks

dataset = datasets.ClinTox("~/molecule-datasets/")
lengths = [int(0.8 * len(dataset)), int(0.1 * len(dataset))]
lengths.append(len(dataset) - sum(lengths))
train_set, valid_set, test_set = torch.utils.data.random_split(dataset, lengths)

model = models.GIN(
    input_dim=dataset.node_feature_dim,
    hidden_dims=[256, 256, 256, 256],
    short_cut=True,
    batch_norm=True,
    concat_hidden=True,
)
task = tasks.PropertyPrediction(
    model,
    task=dataset.tasks,
    criterion="bce",
    metric=("auprc", "auroc"),
)

optimizer = torch.optim.Adam(task.parameters(), lr=1e-3)
solver = core.Engine(
    task,
    train_set,
    valid_set,
    test_set,
    optimizer,
    batch_size=1024,
)
solver.train(num_epoch=100)
solver.evaluate("valid")
```

Add `gpus=[0]` only when a supported CUDA device is available. Omit `gpus` for
CPU execution.

For binary classification, `task.predict(batch)` returns logits; apply
`torch.sigmoid` when probabilities are needed. In 0.2.1, normalized regression
predictions are returned on the original target scale, which is a breaking change
from older releases.

## Choose the official workflow

### Molecular property prediction

- Dataset: `datasets.ClinTox`, `BBBP`, `Tox21`, `QM9`, or another documented
  molecule dataset.
- Model: start with `models.GIN`; use `edge_input_dim` when the selected feature
  configuration supplies edge features.
- Task: `tasks.PropertyPrediction`.
- Read [molecular property prediction](references/molecular_property_prediction.md).

### Self-supervised molecular pretraining

- InfoGraph: `models.InfoGraph(gin_model, separate_model=False)` wrapped by
  `tasks.Unsupervised`.
- Attribute masking: `tasks.AttributeMasking(model, mask_rate=0.15)`.
- Recreate the same encoder for fine-tuning, then load the checkpoint with
  `strict=False` before training `tasks.PropertyPrediction`.
- Read [molecular property prediction](references/molecular_property_prediction.md).

### Molecule generation

- Dataset: `datasets.ZINC250k(..., kekulize=True, atom_feature="symbol")`.
- GCPN: an `models.RGCN` encoder wrapped by `tasks.GCPNGeneration`.
- GraphAF: node and edge `models.GraphAF` flows wrapped by
  `tasks.AutoregressiveGeneration`.
- Supported optimization tasks in the tutorial are `"qed"` and `"plogp"`;
  criteria are `"nll"` and/or `"ppo"`.
- Read [molecular generation](references/molecular_generation.md).

### Retrosynthesis

- Create two synchronized `datasets.USPTO50k` views: reaction mode for center
  identification and `as_synthon=True` for synthon completion.
- Train `tasks.CenterIdentification` and `tasks.SynthonCompletion` separately.
- Combine the trained tasks with `tasks.Retrosynthesis`; do not pass raw models
  directly to the end-to-end task.
- Read [retrosynthesis](references/retrosynthesis.md).

### Knowledge graph reasoning

- Embedding workflow: `datasets.FB15k237` → `models.RotatE` →
  `tasks.KnowledgeGraphCompletion`.
- Neural reasoning workflow: `models.NeuralLP` with `fact_ratio=0.75`.
- Read [knowledge graph reasoning](references/knowledge_graphs.md).

### Protein modeling

- Build proteins with `data.Protein.from_sequence`, `from_pdb`, or
  `from_molecule`.
- Sequence encoders include `models.ESM`, `ProteinCNN`, `ProteinResNet`,
  `ProteinLSTM`, and `ProteinBERT`; structure encoders include `models.GearNet`.
- Use documented graph-construction layers rather than a nonexistent
  `protein.residue_graph()` convenience method.
- Read [protein modeling](references/protein_modeling.md).

## Rules for reliable TorchDrug code

1. **Follow the 0.2.1 API.** The official docs are not a rolling latest-version
   site.
2. **Prefer documented feature names.** Use `atom_feature`, `bond_feature`,
   `residue_feature`, and `mol_feature`; `node_feature`, `edge_feature`, and
   `graph_feature` are deprecated aliases in relevant dataset constructors.
3. **Let `Engine` preprocess tasks.** If composing pre-trained tasks without
   constructing their solvers, call each task's `preprocess()` manually.
4. **Keep paired splits synchronized.** For retrosynthesis, reset the same random
   seed before splitting reaction and synthon datasets.
5. **Use TorchDrug collation.** Use `data.graph_collate` or `core.Engine`;
   generic PyTorch collation does not know how to pack TorchDrug graphs.
6. **Separate model, task, and engine arguments.** A common source of invented
   code is passing task options to a model or passing raw models where a composed
   task is required.
7. **Validate generated chemistry.** Treat model outputs as candidates, not as
   experimentally valid or synthesizable compounds.

## Troubleshooting

### Installation or import failure

Check Python, PyTorch, `torch-scatter`, and `torch-cluster` as one compatibility
set. Most failures are binary-wheel mismatches, unsupported Python versions, or
attempts to use MPS.

### Feature dimension mismatch

Build model dimensions from the loaded dataset:

- `dataset.node_feature_dim`
- `dataset.edge_feature_dim`
- `dataset.num_bond_type`
- `dataset.num_entity` and `dataset.num_relation` for knowledge graphs

Do not hard-code dimensions copied from a different feature configuration.

### Device mismatch

Pass `gpus=[0]` to `core.Engine` for supported CUDA execution. For manual
prediction, collate first and move the entire nested batch with `utils.cuda`.

### Checkpoint mismatch

Recreate the same model and feature configuration. For pretraining-to-fine-tuning
transfer, load the checkpoint's `"model"` state with `strict=False`; for a complete
solver, use `solver.save()` and `solver.load()`.

## Reference index

- [Core concepts and data structures](references/core_concepts.md)
- [Datasets](references/datasets.md)
- [Models and architectures](references/models_architectures.md)
- [Molecular property prediction and pretraining](references/molecular_property_prediction.md)
- [Protein modeling](references/protein_modeling.md)
- [Molecular generation](references/molecular_generation.md)
- [Retrosynthesis](references/retrosynthesis.md)
- [Knowledge graph reasoning](references/knowledge_graphs.md)

## Upstream sources

- [TorchDrug 0.2.1 documentation](https://torchdrug.ai/docs/)
- [Tutorial index](https://torchdrug.ai/docs/tutorials/)
- [Installation](https://torchdrug.ai/docs/installation.html)
- [Package reference](https://torchdrug.ai/docs/api/)
- [TorchDrug 0.2.1 release notes](https://github.com/DeepGraphLearning/torchdrug/releases/tag/v0.2.1)