Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions .gitignore
Original file line number Diff line number Diff line change
@@ -1,3 +1,5 @@
__pycache__/
*.pyc
*.pyo
.claude
.envrc
8 changes: 4 additions & 4 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,7 @@ PyTorch is an **optional** dependency. The full inference pipeline (featurizatio
uv sync

# With PyTorch (needed for converting checkpoints from the original Protenix format)
uv sync --extra torch
uv sync --extra pytorch-cpu
```

## Serialized models
Expand All @@ -31,7 +31,7 @@ Pre-converted Equinox models skip the PyTorch dependency entirely and load in un
Models are hosted on [HuggingFace](https://huggingface.co/nickrb/protenij) and downloaded automatically on first use.

```python
from protenix.backend import load_model
from protenij.backend import load_model

# Downloads from HuggingFace, caches to ~/.protenix/
model = load_model("protenix_base_default_v1.0.0")
Expand All @@ -42,10 +42,10 @@ model = load_model("~/.protenix/protenix_base_default_v1.0.0")

### Translating from a PyTorch checkpoint

Requires the `torch` extra.
Requires the `pytorch-cpu` extra.

```bash
uv sync --extra torch
uv sync --extra pytorch-cpu
python translate_models.py
```

Expand Down
26 changes: 13 additions & 13 deletions bench_load.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
"""Benchmark: PyTorch-path vs Equinox-path model loading."""
import os
os.environ["PROTENIX_DATA_ROOT_DIR"] = os.path.expanduser("~/.protenix")
CACHE_DIR = os.environ.get("PROTENIJ_CACHE_DIR", os.path.expanduser("~/.protenix"))
os.environ["PROTENIX_DATA_ROOT_DIR"] = CACHE_DIR

import time
import jax
Expand All @@ -9,14 +10,13 @@
import torch
from ml_collections.config_dict import ConfigDict

from protenix.configs.configs_base import configs as configs_base
from protenix.configs.configs_data import data_configs
from protenix.configs.configs_inference import inference_configs
from protenix.configs.configs_model_type import model_configs
from protenix.config import parse_configs
from protenij.configs.configs_base import configs as configs_base
from protenij.configs.configs_data import data_configs
from protenij.configs.configs_inference import inference_configs
from protenij.configs.configs_model_type import model_configs
from protenij.config import parse_configs

MODEL_NAME = "protenix_base_default_v1.0.0"
CACHE_DIR = os.path.expanduser("~/.protenix")
EQX_PATH = os.path.join(CACHE_DIR, MODEL_NAME) # will produce .eqx + .skeleton.pkl


Expand All @@ -32,9 +32,9 @@ def build_configs():

def load_torch_path(configs):
"""Load model via PyTorch, return (jax_model, timings_dict)."""
from protenix.model.protenix import Protenix as TorchProtenix
import protenix.protenij
from protenix.backend import from_torch
from protenij.model.protenix import Protenix as TorchProtenix
import protenij.protenij
from protenij.backend import from_torch

# Step 1: torch.load + DDP strip
checkpoint_path = os.path.join(CACHE_DIR, f"{MODEL_NAME}.pt")
Expand Down Expand Up @@ -72,11 +72,11 @@ def load_torch_path(configs):

def load_eqx_path():
"""Load model via Equinox serialization, return (jax_model, timings_dict)."""
from protenix.backend import load_model
from protenij.backend import load_model

# Step 1: pickle.load skeleton
t0 = time.perf_counter()
import protenix.protenij # ensure pytree node types are registered
import protenij.protenij # ensure pytree node types are registered
import pickle
with open(f"{EQX_PATH}.skeleton.pkl", "rb") as f:
skeleton = pickle.load(f)
Expand Down Expand Up @@ -125,7 +125,7 @@ def main():

# -- Save if needed --
if not eqx_exists:
from protenix.backend import save_model
from protenij.backend import save_model

print(f"\nSaving Equinox model to {EQX_PATH}.eqx ...")
t0 = time.perf_counter()
Expand Down
File renamed without changes.
22 changes: 19 additions & 3 deletions protenix/backend.py → protenij/backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -148,7 +148,7 @@ def save_model(model, path):
"clusters-by-entity-40.txt",
]

_CACHE_DIR = os.path.expanduser("~/.protenix")
_CACHE_DIR = os.environ.get("PROTENIJ_CACHE_DIR", os.path.expanduser("~/.protenix"))


def _hf_download(filename: str) -> None:
Expand Down Expand Up @@ -182,6 +182,21 @@ def _resolve_model_path(name_or_path: str) -> str:
return cache_path


class _ProtenixAliasUnpickler(pickle.Unpickler):
"""Resolve legacy `protenix.*` class paths to the renamed `protenij.*` package.

Existing skeleton.pkl files were pickled when this package was named
`protenix`. Rather than aliasing globally in sys.modules (which would
hijack the namespace for any concurrently-installed upstream `protenix`),
we redirect lookups only during this unpickle call.
"""

def find_class(self, module, name):
if module == "protenix" or module.startswith("protenix."):
module = "protenij" + module[len("protenix"):]
return super().find_class(module, name)


def load_model(name_or_path: str):
"""Load an Equinox model (no PyTorch needed).

Expand All @@ -190,12 +205,13 @@ def load_model(name_or_path: str):
or a model name (e.g. 'protenix_base_default_v1.0.0') which will be
resolved from ~/.protenix/ or downloaded from HuggingFace.
"""
import protenix.protenij # ensure pytree node types are registered
import protenij.protenij # ensure pytree node types are registered

download_data()
path = _resolve_model_path(name_or_path)
with open(f"{path}.skeleton.pkl", "rb") as f:
skeleton = pickle.load(f)
skeleton = _ProtenixAliasUnpickler(f).load()

return eqx.tree_deserialise_leaves(f"{path}.eqx", skeleton)


Expand Down
File renamed without changes.
2 changes: 1 addition & 1 deletion protenix/config/config.py → protenij/config/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,7 @@
import yaml
from ml_collections.config_dict import ConfigDict

from protenix.config.extend_types import (
from protenij.config.extend_types import (
DefaultNoneWithType,
GlobalConfigValue,
ListValue,
Expand Down
File renamed without changes.
File renamed without changes.
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@
# limitations under the License.

# pylint: disable=C0114,C0301
from protenix.config.extend_types import (
from protenij.config.extend_types import (
GlobalConfigValue,
ListValue,
RequiredValue,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,7 @@
from copy import deepcopy
from pathlib import Path

from protenix.config.extend_types import GlobalConfigValue, ListValue
from protenij.config.extend_types import GlobalConfigValue, ListValue

default_test_configs = {
"sampler_configs": {
Expand Down Expand Up @@ -122,7 +122,8 @@
},
}
# HARDCODE cache path.
DATA_ROOT_DIR = str(Path("~/.protenix").expanduser())#os.environ.get("PROTENIX_DATA_ROOT_DIR", str(Path("~/.protenix").expanduser()))
#DATA_ROOT_DIR = str(Path("~/.protenix").expanduser())
DATA_ROOT_DIR = os.environ.get("PROTENIJ_CACHE_DIR", str(Path("~/.protenix").expanduser()))

# Use CCD cache created by scripts/gen_ccd_cache.py priority. (without date in filename)
# See: docs/prepare_data.md
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@
# pylint: disable=C0114
import os

from protenix.config.extend_types import ListValue, RequiredValue
from protenij.config.extend_types import ListValue, RequiredValue

current_file_path = os.path.abspath(__file__)
current_directory = os.path.dirname(current_file_path)
Expand Down
File renamed without changes.
4 changes: 2 additions & 2 deletions protenix/data/ccd.py → protenij/data/ccd.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,8 +26,8 @@
from biotite.structure import AtomArray
from rdkit import Chem

from protenix.configs.configs_data import data_configs
from protenix.data.substructure_perms import get_substructure_perms
from protenij.configs.configs_data import data_configs
from protenij.data.substructure_perms import get_substructure_perms

logger = logging.getLogger(__name__)

Expand Down
File renamed without changes.
File renamed without changes.
Original file line number Diff line number Diff line change
Expand Up @@ -29,10 +29,10 @@
from biotite.structure import AtomArray
from scipy.spatial.distance import cdist

from protenix.data.constants import ELEMS, STD_RESIDUES
from protenix.data.tokenizer import Token, TokenArray
from protenix.data.utils import get_atom_mask_by_name
from protenix.utils.logger import get_logger
from protenij.data.constants import ELEMS, STD_RESIDUES
from protenij.data.tokenizer import Token, TokenArray
from protenij.data.utils import get_atom_mask_by_name
from protenij.utils.logger import get_logger

logger = get_logger(__name__)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -23,11 +23,11 @@
import pandas as pd
from biotite.structure import AtomArray

from protenix.data.msa_featurizer import MSAFeaturizer
from protenix.data.parser import DistillationMMCIFParser, MMCIFParser
from protenix.data.tokenizer import AtomArrayTokenizer, TokenArray
from protenix.utils.cropping import CropData
from protenix.utils.file_io import load_gzip_pickle
from protenij.data.msa_featurizer import MSAFeaturizer
from protenij.data.parser import DistillationMMCIFParser, MMCIFParser
from protenij.data.tokenizer import AtomArrayTokenizer, TokenArray
from protenij.utils.cropping import CropData
from protenij.utils.file_io import load_gzip_pickle


class DataPipeline(object):
Expand Down
6 changes: 3 additions & 3 deletions protenix/data/dataloader.py → protenij/data/dataloader.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,9 +20,9 @@
from ml_collections.config_dict import ConfigDict
from torch.utils.data import DataLoader, DistributedSampler, Sampler

from protenix.data.dataset import Dataset, get_datasets
from protenix.utils.logger import get_logger
from protenix.utils.torch_utils import collate_fn_first
from protenij.data.dataset import Dataset, get_datasets
from protenij.utils.logger import get_logger
from protenij.utils.torch_utils import collate_fn_first

logger = get_logger(__name__)

Expand Down
24 changes: 12 additions & 12 deletions protenix/data/dataset.py → protenij/data/dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,21 +27,21 @@
from ml_collections.config_dict import ConfigDict
from torch.utils.data import Dataset

from protenix.data.constants import EvaluationChainInterface
from protenix.data.constraint_featurizer import ConstraintFeatureGenerator
from protenix.data.data_pipeline import DataPipeline
from protenix.data.featurizer import Featurizer
from protenix.data.msa_featurizer import MSAFeaturizer
from protenix.data.tokenizer import TokenArray
from protenix.data.utils import (
from protenij.data.constants import EvaluationChainInterface
from protenij.data.constraint_featurizer import ConstraintFeatureGenerator
from protenij.data.data_pipeline import DataPipeline
from protenij.data.featurizer import Featurizer
from protenij.data.msa_featurizer import MSAFeaturizer
from protenij.data.tokenizer import TokenArray
from protenij.data.utils import (
data_type_transform,
get_antibody_clusters,
make_dummy_feature,
)
from protenix.utils.cropping import CropData
from protenix.utils.file_io import read_indices_csv
from protenix.utils.logger import get_logger
from protenix.utils.torch_utils import dict_to_tensor
from protenij.utils.cropping import CropData
from protenij.utils.file_io import read_indices_csv
from protenij.utils.logger import get_logger
from protenij.utils.torch_utils import dict_to_tensor

logger = get_logger(__name__)

Expand All @@ -50,7 +50,7 @@ class BaseSingleDataset(Dataset):
"""
dataset for a single data source
data = self.__item__(idx)
return a dict of features and labels, the keys and the shape are defined in protenix.data.utils
return a dict of features and labels, the keys and the shape are defined in protenij.data.utils
"""

def __init__(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -24,8 +24,8 @@
except ImportError:
_HAS_TORCH = False

from protenix.data.compute_esm import compute_ESM_embeddings, load_esm_model
from protenix.utils.logger import get_logger
from protenij.data.compute_esm import compute_ESM_embeddings, load_esm_model
from protenij.utils.logger import get_logger

logger = get_logger(__name__)

Expand Down
8 changes: 4 additions & 4 deletions protenix/data/featurizer.py → protenij/data/featurizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,10 +20,10 @@
from biotite.structure import Atom, AtomArray, get_residue_starts
from sklearn.neighbors import KDTree

from protenix.data.constants import STD_RESIDUES, STD_RESIDUES_WITH_GAP, get_all_elems
from protenix.data.tokenizer import Token, TokenArray
from protenix.data.utils import get_atom_level_token_mask, get_ligand_polymer_bond_mask
from protenix.utils.geometry import angle_3p, random_transform
from protenij.data.constants import STD_RESIDUES, STD_RESIDUES_WITH_GAP, get_all_elems
from protenij.data.tokenizer import Token, TokenArray
from protenij.data.utils import get_atom_level_token_mask, get_ligand_polymer_bond_mask
from protenij.utils.geometry import angle_3p, random_transform


class Featurizer(object):
Expand Down
2 changes: 1 addition & 1 deletion protenix/data/filter.py → protenij/data/filter.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,7 @@
from biotite.structure import AtomArray, get_molecule_indices
from scipy.spatial.distance import cdist

from protenix.data.constants import CRYSTALLIZATION_AIDS
from protenij.data.constants import CRYSTALLIZATION_AIDS


class Filter(object):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -24,13 +24,13 @@
from biotite.structure import AtomArray
from torch.utils.data import DataLoader, Dataset, DistributedSampler

from protenix.data.data_pipeline import DataPipeline
from protenix.data.esm_featurizer import ESMFeaturizer
from protenix.data.json_to_feature import SampleDictToFeatures
from protenix.data.msa_featurizer import InferenceMSAFeaturizer
from protenix.data.utils import data_type_transform, make_dummy_feature
from protenix.utils.distributed import DIST_WRAPPER
from protenix.utils.torch_utils import collate_fn_identity, dict_to_tensor
from protenij.data.data_pipeline import DataPipeline
from protenij.data.esm_featurizer import ESMFeaturizer
from protenij.data.json_to_feature import SampleDictToFeatures
from protenij.data.msa_featurizer import InferenceMSAFeaturizer
from protenij.data.utils import data_type_transform, make_dummy_feature
from protenij.utils.distributed import DIST_WRAPPER
from protenij.utils.torch_utils import collate_fn_identity, dict_to_tensor

logger = logging.getLogger(__name__)

Expand Down
8 changes: 4 additions & 4 deletions protenix/data/json_maker.py → protenij/data/json_maker.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,10 +22,10 @@
import numpy as np
from biotite.structure import AtomArray, get_chain_starts, get_residue_starts

from protenix.data.constants import STD_RESIDUES
from protenix.data.filter import Filter
from protenix.data.parser import AddAtomArrayAnnot, MMCIFParser
from protenix.data.utils import get_lig_lig_bonds, get_ligand_polymer_bond_mask
from protenij.data.constants import STD_RESIDUES
from protenij.data.filter import Filter
from protenij.data.parser import AddAtomArrayAnnot, MMCIFParser
from protenij.data.utils import get_lig_lig_bonds, get_ligand_polymer_bond_mask


def merge_covalent_bonds(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -26,7 +26,7 @@
from rdkit import Chem
from rdkit.Chem import AllChem

from protenix.data import ccd
from protenij.data import ccd

logger = logging.getLogger(__name__)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -18,12 +18,12 @@
import numpy as np
from biotite.structure import AtomArray

from protenix.data.constraint_featurizer import ConstraintFeatureGenerator
from protenix.data.featurizer import Featurizer
from protenix.data.json_parser import add_entity_atom_array, remove_leaving_atoms
from protenix.data.parser import AddAtomArrayAnnot
from protenix.data.tokenizer import AtomArrayTokenizer, TokenArray
from protenix.data.utils import int_to_letters
from protenij.data.constraint_featurizer import ConstraintFeatureGenerator
from protenij.data.featurizer import Featurizer
from protenij.data.json_parser import add_entity_atom_array, remove_leaving_atoms
from protenij.data.parser import AddAtomArrayAnnot
from protenij.data.tokenizer import AtomArrayTokenizer, TokenArray
from protenij.data.utils import int_to_letters

logger = logging.getLogger(__name__)

Expand Down
Loading