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
|
pocket_vocab_size
|
int
|
Pocket residue vocabulary size. (default: |
512
|
embed_dim
|
int
|
Transformer feature embedding dimension. (default: |
512
|
pair_dim
|
int
|
Pairwise feature dimension. (default: |
128
|
num_layers
|
int
|
Number of joint transformer layers. (default: |
12
|
num_heads
|
int
|
Number of multihead attention heads. (default: |
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 |
required |
pocket_tokens
|
Tensor
|
Pocket tokens |
None
|
mol_coords
|
Tensor
|
Initial ligand 3D coordinates |
None
|
pocket_coords
|
Tensor
|
Pocket 3D coordinates |
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'
|
folder
|
str
|
Target directory to save the checkpoint. (default: |
'checkpoints'
|
log
|
bool
|
Whether to print download progress. (default: |
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'
|
download
|
bool
|
Whether to download checkpoint if missing locally. (default: |
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'
|
folder
|
str
|
Target download folder. (default: |
'checkpoints'
|
log
|
bool
|
Print download log. (default: |
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.