Downloads · 30 days
36
100% of all-time downloads
NisargOza/CausalCellJEPA
CausalCellJEPA is a feature extraction model from NisargOza. Use it when you need embeddings to search or compare text. It is set up for pytorch. The card lists the license as other.
CausalCellJEPA predicts unpaired single-cell perturbation populations from a baseline population and a biological action embedding. It combines a frozen JEPA cell encoder, ESM-2 action features, and an action-conditio…
Downloads · 30 days
36
100% of all-time downloads
All-time downloads
36
Public
Repo size
126 MB
Likes
0
Public
Click a slice to open those files.
.safetensors68.2 MB · 95%
From the Hugging Face model README
CausalCellJEPA predicts unpaired single-cell perturbation populations from a baseline population
and a biological action embedding. It combines a frozen JEPA cell encoder, ESM-2 action features,
and an action-conditioned set transition. This is a custom PyTorch research model, not a
Transformers AutoModel.
The completed study is a mixed result, not a validated global state-of-the-art result. Distributional calibration is strong, while latent effect direction and outcome-fit-unseen transfer remain limitations. Source and full evidence: NisargOza/CausalCellJEPA.
| Component | Purpose | File |
|---|---|---|
stage1_teacher | Expression tokens → 256D cell state | weights |
stage2_primary | Control set + 320D action → predicted set | weights |
transcriptomic_readout | 256D latent → 3,000-HVG expression | weights |
stage2_multiteacher_v4 | Exploratory ESM-2 + GO dynamics | weights |
MODEL_MANIFEST.json lists every weight, metadata file, configuration, original-artifact hash,
and exported SHA-256. Additional released components include the effect anchor, external-response
and STRING+GO predictors, control OOD gate, and frozen Replogle action caches. Raw data, optimizer
state, superseded checkpoints, State baseline weights, and third-party ESM-2 weights are excluded.
>=3.11,<3.14; validated with 3.12.132.13.0, NumPy 2.5.2, h5py 3.16.00.8.0, GeomLoss 0.3.1, PyYAML 6.0.31.32.0The exact loading environment is pinned in requirements.txt.
python3.12 -m venv .venv
source .venv/bin/activate
python -m pip install huggingface_hub==1.32.0 torch==2.13.0 numpy==2.5.2 \
h5py==3.16.0 safetensors==0.8.0 geomloss==0.3.1 PyYAML==6.0.3
from pathlib import Path
import sys
import torch
from huggingface_hub import snapshot_download
snapshot = Path(snapshot_download("NisargOza/CausalCellJEPA"))
sys.path.insert(0, str(snapshot))
from load_components import load_primary_dynamics, load_tensor_component
model, _ = load_primary_dynamics(snapshot)
model.eval()
# API smoke inputs; replace with normalized JEPA latents and ESM-2 action features.
control = torch.zeros(1, 32, 256)
action = torch.zeros(1, 320)
action_known = torch.ones(1, dtype=torch.bool)
with torch.inference_mode():
predicted_latents = model(control, action, action_known)
readout, metadata = load_tensor_component(snapshot, "transcriptomic_readout")
predicted_expression = (
predicted_latents @ readout["weights"] + readout["bias"]
).clamp_min(metadata["output_clamp_min"])
print(predicted_latents.shape, predicted_expression.shape)
# torch.Size([1, 32, 256]) torch.Size([1, 32, 3000])
load_components.py verifies the selected component weights and metadata against
MODEL_MANIFEST.json. Inputs are control [batch, 32, 256], action [batch, 320], and
boolean action_known [batch]; outputs are normalized predicted latents
[batch, 32, 256]. The zero tensors above test the API only and are not biologically meaningful.
Normalization, target vocabulary, and feature provenance are bundled under provenance/.
| Replogle double OOD, 199 targets | CausalCellJEPA | Pseudo-paired |
|---|---|---|
| Magnitude absolute error ↓ | 0.0591 | 0.2274 |
| MMD ↓ | 0.0294 | 0.0971 |
| Sinkhorn divergence ↓ | 0.1371 | 0.2735 |
| Latent effect Pearson ↑ | 0.0715 | 0.1873 |
On 27 Adamson targets, the frozen final candidate improved Systema Pearson over perturbed mean by
+0.2606 (95% target-bootstrap CI [+0.1563, +0.3700]) but trailed STRING+GO by -0.0201
([-0.0422, -0.0002]). The external confirmation therefore failed one of six preregistered
criteria and does not support a global SOTA claim.
Use this release for perturbation-model research, feature extraction, and reproduction of the
reported comparisons. Replogle, Adamson, and Nadig single-cell data are not redistributed;
identifiers, splits, checksums, roles, and leakage audits are under provenance/.
The model is not validated for clinical, diagnostic, treatment-selection, or safety-critical use. New targets require rebuilding biological features. Exploratory post-test components must not be interpreted as confirmatory replacements for the primary architecture.
Tensor weights use Safetensors. MODEL_MANIFEST.json records exported and source-artifact
hashes; the retained small effect-anchor checkpoint is loaded with weights_only=True.
No software or model-weight license was present in the source repository, so this card uses
license: other; public download does not itself grant reuse or redistribution rights. Cite the
GitHub repository, the exact Hub revision, and CITATION.cff.