JAX port of CDDD (Continuous Data-Driven Descriptors) for molecular SMILES embedding/reconstruction
Project description
jax-cddd 🧪
A JAX port of CDDD (Continuous and Data-Driven Descriptors) — a SMILES autoencoder that produces a fixed-length, continuous molecular descriptor by translating between equivalent SMILES representations of the same molecule.
Winter, R., Montanari, F., Noé, F. & Clevert, D.-A. Learning continuous and data-driven molecular descriptors by translating equivalent chemical representations, Chem. Sci., 2019.
This project reimplements the pretrained default_model (encoder + decoder) in
JAX/Flax, loading the original TensorFlow 1.x checkpoint weights unmodified. It
is validated both against an independent from-scratch recomputation of the
model's math and against the real, unmodified original code running under a
legacy TensorFlow 1.15, to within 1e-4 numerical tolerance — see
Validation.
Only inference (embedding and reconstruction) is in scope. Training and the auxiliary QSAR property head are not ported.
✨ What you get
seq_to_emb: SMILES → 512-dim descriptor (the encoder)emb_to_seq: 512-dim descriptor → reconstructed SMILES (the decoder, greedy or beam search)- A memory-efficient batching path for embedding/reconstructing millions of molecules without exhausting GPU memory
📦 Installation
pip install jax-cddd
jax-cddd is on PyPI. For development
(editable install from a clone), instead:
micromamba create -n jax-cddd python=3.11 pip -c conda-forge -y
micromamba run -n jax-cddd pip install -e .
Getting the pretrained weights
Two data files — the converted weights
(src/jax_cddd/data/default_model_params.npz, ~200 MB) and the vocabulary
(src/jax_cddd/data/indices_char.npy) — are downloaded automatically,
with checksum verification, the first time you instantiate CDDDModel(),
if they aren't already present locally — see Quickstart below.
No action needed; they're fetched from this repo's GitHub release.
To point at a different location instead (e.g. a fork's own release), override
via environment variables:
export JAX_CDDD_WEIGHTS_URL="https://.../default_model_params.npz"
export JAX_CDDD_VOCAB_URL="https://.../indices_char.npy"
🚀 Quickstart
from jax_cddd.inference import CDDDModel
model = CDDDModel() # downloads the converted default_model weights on first use
# Embed
embedding = model.seq_to_emb("CC(=O)Oc1ccccc1C(=O)O") # aspirin -> [512] array
# Reconstruct
smiles = model.emb_to_seq(embedding) # beam search, returns a single string
print(smiles) # "CC(=O)Oc1ccccc1C(=O)O"
CDDDModel loads once and can be reused for as many calls as you like — hold
onto the instance rather than re-creating it per molecule.
⚡ Performance
Both seq_to_emb and emb_to_seq are JIT-compiled under the hood. CDDDModel()
eagerly warms up the shapes a single-molecule workload hits during
construction (a few extra seconds, on top of loading the weights), so a
freshly constructed model's first real call is already fast rather than
paying that compile cost then:
seq_to_emb: ~10-30 ms per molecule.emb_to_seq: ~70-90 ms per molecule (defaultbeam_width=10).
Molecules with an unusual shape (SMILES longer than ~64 tokens, or a
non-default beam_width/max_len) fall outside the warmed-up set and pay a
one-time compile the first time that particular shape is used — every
subsequent call with the same shape is fast again. Pass CDDDModel(warmup=False)
to skip warmup and get a near-instant construction instead, if you'd rather
pay the compile cost on first use than at startup.
📚 Batches, at any size
Both methods accept lists directly, from a handful of molecules up to entire multi-million-molecule datasets — large inputs are automatically processed in memory-bounded chunks under the hood (sorted by length first, to minimize the compute wasted padding short sequences up to one long outlier's length), so you never need to chunk your own input:
smiles_list = [...] # any size
embeddings = model.seq_to_emb(smiles_list) # [n, 512] array
reconstructed = model.emb_to_seq(embeddings) # list[str]
# Top-k hypotheses per molecule instead of just the best one
top3 = model.emb_to_seq(embeddings, beam_width=10, num_top=3) # list[list[str]]
Knobs that matter for quality/speed/memory:
beam_width(default10): higher explores more reconstruction hypotheses;beam_width=1degenerates to plain greedy decoding and is fastest.max_len(default1000): decoding step cap; lower it if you know your molecules are small, to save time.chunk_size(default512forseq_to_emb,256foremb_to_seq): maximum number of molecules/embeddings processed per internal batch. Lower it if you hit GPU memory limits (beam search holdschunk_size × beam_width × max_lenworth of state, so it's more memory-hungry than encoding); raise it for more throughput on a bigger GPU.
A couple of tips beyond that for very large runs:
- 🔹 Prefer greedy (
beam_width=1) if top-1 accuracy is good enough — beam search is several times more expensive in both memory and time. - 🔹 Cap GPU preallocation on small GPUs via
XLA_PYTHON_CLIENT_MEM_FRACTIONif you share the GPU with other processes (the package already requests a conservative0.3by default).
✅ Validation
Fidelity is checked two independent ways:
- Against the math, from scratch.
tests/test_against_tf1_reference.pycompares the JAX encoder/decoder to reference activations recomputed directly from the checkpoint's raw tensors, using a from-scratch NumPy reimplementation of the TF1GRUCellformula (not a copy of this package's code). On 30 diverse molecules: max encoder error1.2e-6, max decoder logit error1.9e-5(target was1e-4), exact-argmax agreement on every token. - Against the actual original code.
tests/test_against_original_code.pycompares the JAX port to the real, unmodified original implementation — the genuinetf.contrib.rnn.MultiRNNCell/tf.contrib.seq2seq.BeamSearchDecodergraph, checkpoint restored viatf.train.Saver, executed under a legacy TensorFlow 1.15. On a handful of real drugs: max embedding error1.4e-6, and every reconstructed SMILES matches exactly, including the same two near-misses the original code also makes (e.g. acetic acid reconstructed as its acetate anion).
micromamba run -n jax-cddd python -m pytest tests/
Both fixtures are pre-generated and committed (tests/fixtures/), so the test
suite above never needs TensorFlow. Regenerating them requires a legacy
environment (an old Python for the old TF wheel) and the original
default_model.zip checkpoint placed under src/jax_cddd/_cddd_ref_/:
micromamba create -n jax-cddd-convert python=3.7 pip -c conda-forge -y
micromamba run -n jax-cddd-convert pip install "tensorflow==1.15.5" "protobuf==3.20.3"
micromamba run -n jax-cddd-convert python scripts/validate_against_tf1_reference.py
micromamba run -n jax-cddd-convert python scripts/run_original_model.py
(That environment is never pip install -e .'d — this package requires
Python ≥3.10 — its scripts add src/ to sys.path directly to reuse
jax_cddd's dependency-free submodules.)
A round-trip smoke test on real drugs (aspirin, caffeine, ibuprofen, ...) is also available standalone:
micromamba run -n jax-cddd python scripts/reconstruct_smoke_test.py
🗂️ Project layout
src/jax_cddd/
vocab.py # SMILES tokenizer + vocabulary
gru.py # GRU cell matching the original TF1 math exactly
modules.py # encoder / decoder forward passes
decoding.py # greedy + beam search decoding
params.py # weight pytree, .npz (de)serialization, model loading
inference.py # CDDDModel: the public API
_legacy_tf1/ # original TF1 reference code, kept for history
scripts/
convert_checkpoint.py # maintainer-only: how the published weights were produced
validate_against_tf1_reference.py # from-scratch NumPy fidelity fixture
run_original_model.py # runs the real original TF1 graph, for fidelity fixture
reconstruct_smoke_test.py # end-to-end embed/reconstruct demo
tests/
📄 License & citation
The original CDDD implementation is MIT-licensed (© 2018 Jan Robin Winter). If you use this in published work, please cite the original paper above.
Project details
Release history Release notifications | RSS feed
Download files
Download the file for your platform. If you're not sure which to choose, learn more about installing packages.
Source Distribution
Built Distribution
Filter files by name, interpreter, ABI, and platform.
If you're not sure about the file name format, learn more about wheel file names.
Copy a direct link to the current filters
File details
Details for the file jax_cddd-0.1.0.post0.tar.gz.
File metadata
- Download URL: jax_cddd-0.1.0.post0.tar.gz
- Upload date:
- Size: 30.4 kB
- Tags: Source
- Uploaded using Trusted Publishing? No
- Uploaded via: twine/6.2.0 CPython/3.11.15
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
be6f903351fd6e569a4abbc44204450f9a736dbb40056a0f324a4187dda59953
|
|
| MD5 |
b089d1e34234f4d9e3ac08e37120e7c4
|
|
| BLAKE2b-256 |
24630f820c606d0bed59f4c5007b8dc691c725e766750ea2210353c32d3f7d4b
|
File details
Details for the file jax_cddd-0.1.0.post0-py3-none-any.whl.
File metadata
- Download URL: jax_cddd-0.1.0.post0-py3-none-any.whl
- Upload date:
- Size: 24.3 kB
- Tags: Python 3
- Uploaded using Trusted Publishing? No
- Uploaded via: twine/6.2.0 CPython/3.11.15
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
943d9d2b8209800771f84d498becfef8417fa27399084fc501b07c2460b4699d
|
|
| MD5 |
11724380cf45809bb0c6ddf320a69170
|
|
| BLAKE2b-256 |
9a89616ec99b771348fd63297eeaa3620027fa77cc1d2e99a552ddeca243b8d6
|