Skip to content

Biological & Macromolecular Models

The k3_node.applications.bio module provides multi-backend Keras 3 ports of Uni-Mol's protein–ligand binding-pose prediction models, for docking a small-molecule ligand into a protein pocket given 3D structural input.


Models

UniMolDockingModel

The original Uni-Mol docking head: a UniMolModel subclass that adds a cross-attention coordinate-update head predicting an iterative pose correction for the ligand, alongside a pairwise distance head.

k3_node.models.unimol.UniMolDockingModel

Bases: UniMolModel

Uni-Mol Protein-Ligand Binding Pose Prediction Model.

Example
import numpy as np
from k3_node.models import UniMolDockingModel

tokens = np.random.randint(1, 64, size=(2, 6))  # atom tokens of 2 molecules with 6 atoms
coords = np.random.rand(2, 6, 3).astype("float32") * 3.0  # 3D conformations

model = UniMolDockingModel(vocab_size=64, encoder_layers=2, encoder_embed_dim=32, encoder_ffn_embed_dim=64,
                           encoder_attention_heads=4, num_kernel=16)
pose, pred_dist = model(tokens, src_coord=coords)
print(tuple(pose.shape), tuple(pred_dist.shape))  # (2, 6, 3) (2, 6, 6)

Usage:

from k3_node.models.unimol import UniMolDockingModel, download_unimol_checkpoint, load_unimol_weights

model = UniMolDockingModel(output_dim=2, data_type="molecule")
ckpt = download_unimol_checkpoint("binding_pose")
load_unimol_weights(model, checkpoint_path=ckpt)

# src_tokens: [batch, seq_len] combined ligand+pocket tokens; src_coord: [batch, seq_len, 3]
pose_update, pred_dist = model(src_tokens, src_coord=src_coord)

DockingPoseModelV2

Uni-Mol Docking V2: a joint ligand/pocket transformer (separate token vocabularies and embeddings for molecule vs. pocket, fused through shared pair-aware transformer layers) that predicts both the ligand binding pose and pairwise distance matrix in a single forward pass.

k3_node.models.unimol_docking_v2.DockingPoseModelV2

Bases: Model

Uni-Mol Docking V2 model for joint protein pocket and ligand holo binding pose prediction.

Parameters:

Name Type Description Default
mol_vocab_size int

Molecule vocabulary size. (default: 512)

512
pocket_vocab_size int

Pocket residue vocabulary size. (default: 512)

512
embed_dim int

Transformer feature embedding dimension. (default: 512)

512
pair_dim int

Pairwise feature dimension. (default: 128)

128
num_layers int

Number of joint transformer layers. (default: 12)

12
num_heads int

Number of multihead attention heads. (default: 32)

32
**kwargs

Additional model arguments.

{}
Example
import numpy as np
from k3_node.models import DockingPoseModelV2

mol_tokens = np.random.randint(0, 64, size=(2, 4))  # ligand atoms
pocket_tokens = np.random.randint(0, 64, size=(2, 6))  # protein pocket atoms
mol_coords = np.random.rand(2, 4, 3).astype("float32")
pocket_coords = np.random.rand(2, 6, 3).astype("float32")
model = DockingPoseModelV2(mol_vocab_size=64, pocket_vocab_size=64, embed_dim=32, pair_dim=16,
                           num_layers=2, num_heads=4)
docked_coords, pred_dist = model(mol_tokens, pocket_tokens, mol_coords=mol_coords, pocket_coords=pocket_coords)
print(tuple(docked_coords.shape), tuple(pred_dist.shape))  # (2, 4, 3) (2, 10, 10)

call(mol_tokens, pocket_tokens=None, mol_coords=None, pocket_coords=None, mol_mask=None, pocket_mask=None, training=False)

Forward pass for Docking V2.

Parameters:

Name Type Description Default
mol_tokens Tensor

Ligand tokens [batch_size, mol_len].

required
pocket_tokens Tensor

Pocket tokens [batch_size, pocket_len].

None
mol_coords Tensor

Initial ligand 3D coordinates [batch_size, mol_len, 3].

None
pocket_coords Tensor

Pocket 3D coordinates [batch_size, pocket_len, 3].

None
mol_mask Tensor

Boolean padding mask for ligand.

None
pocket_mask Tensor

Boolean padding mask for pocket.

None
training bool

Training flag.

False

Returns:

Type Description

Tuple[Tensor, Tensor]: Predicted docked ligand 3D coordinates and predicted complex distance matrix.

Usage:

from k3_node.models.unimol_docking_v2 import (
    DockingPoseModelV2,
    download_unimol_docking_checkpoint,
    load_unimol_docking_weights,
)

model = DockingPoseModelV2(embed_dim=512, pair_dim=128, num_layers=12, num_heads=32)
ckpt = download_unimol_docking_checkpoint()
load_unimol_docking_weights(model, checkpoint_path=ckpt)

pose, pair_dist = model(
    mol_tokens, pocket_tokens=pocket_tokens,
    mol_coords=mol_coords, pocket_coords=pocket_coords,
)


Pretrained Checkpoints

download_unimol_checkpoint

k3_node.models.unimol.download_unimol_checkpoint(name='mol_pre_no_h', folder='checkpoints', log=True)

Downloads a pre-trained Uni-Mol checkpoint (.pt).

Parameters:

Name Type Description Default
name str

Pretrained checkpoint name or alias (e.g. "mol_pre_no_h", "mol_pre_all_h", "pocket_pre", "mp_all_h", "oled_pre_no_h", "qm9", "drugs", "binding_pose").

'mol_pre_no_h'
folder str

Target directory to save the checkpoint. (default: "checkpoints")

'checkpoints'
log bool

Whether to print download progress. (default: True)

True

Returns:

Name Type Description
str str

Absolute path to the downloaded checkpoint file.

load_unimol_weights

k3_node.models.unimol.load_unimol_weights(model, checkpoint_path=None, pretrained_name=None, folder='checkpoints', download=True)

Loads pre-trained weights from a PyTorch checkpoint (.pt) into a Keras 3 UniMolModel.

Parameters:

Name Type Description Default
model UniMolModel

The target UniMolModel instance.

required
checkpoint_path str

Local path to .pt file or checkpoint name.

None
pretrained_name str

Pretrained model identifier.

None
folder str

Directory to store downloaded checkpoints. (default: "checkpoints")

'checkpoints'
download bool

Whether to download checkpoint if missing locally. (default: True)

True

Returns:

Name Type Description
UniMolModel UniMolModel

The model with loaded weights.

download_unimol_docking_checkpoint

k3_node.models.unimol_docking_v2.download_unimol_docking_checkpoint(version='v2', folder='checkpoints', log=True)

Downloads a pre-trained Uni-Mol Docking V2 checkpoint (.pt).

Parameters:

Name Type Description Default
version str

Docking model version ("v2").

'v2'
folder str

Target download folder. (default: "checkpoints")

'checkpoints'
log bool

Print download log. (default: True)

True

Returns:

Name Type Description
str str

Absolute path to the downloaded file.

load_unimol_docking_weights

k3_node.models.unimol_docking_v2.load_unimol_docking_weights(model, checkpoint_path=None, folder='checkpoints', download=True)

Loads pre-trained weights into a Keras 3 DockingPoseModelV2 model.

See the Fine-Tuning Recipes guide for a worked example of fine-tuning UniMolDockingModel on a custom docking dataset.