Initial upload: JEPA-Anything full project + domain checkpoints
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- .DS_Store +0 -0
- 01_Vision_Controlled_Binding/README.md +12 -0
- 01_Vision_Controlled_Binding/Shapes3D_JEPA_Pretraining/README.md +36 -0
- 01_Vision_Controlled_Binding/Shapes3D_JEPA_Pretraining/checkpoints/Shapes3D_JEPA_Pretraining_final_checkpoint.pt +3 -0
- 01_Vision_Controlled_Binding/Shapes3D_Main_Factor_Prediction/README.md +36 -0
- 01_Vision_Controlled_Binding/Shapes3D_Main_Factor_Prediction/checkpoints/Shapes3D_Main_Factor_Prediction_final_checkpoint.pt +3 -0
- 02_Single_Cell/Norman_Single_Cell_Perturbation_Prediction/README.md +44 -0
- 02_Single_Cell/Norman_Single_Cell_Perturbation_Prediction/checkpoints/Norman_Single_Cell_Perturbation_Prediction_final_checkpoint.pt +3 -0
- 02_Single_Cell/Norman_Single_Cell_Perturbation_Prediction/model/Norman_Single_Cell_Perturbation_Prediction_vocab.json +0 -0
- 02_Single_Cell/Norman_Single_Cell_Perturbation_Prediction/model/cell_jepa/__init__.py +12 -0
- 02_Single_Cell/Norman_Single_Cell_Perturbation_Prediction/model/cell_jepa/objectives.py +220 -0
- 02_Single_Cell/Norman_Single_Cell_Perturbation_Prediction/model/cell_jepa/objectives_anything.py +558 -0
- 02_Single_Cell/Norman_Single_Cell_Perturbation_Prediction/test/Norman_Single_Cell_Perturbation_Prediction_test.py +143 -0
- 02_Single_Cell/Norman_Single_Cell_Perturbation_Prediction/test/Norman_Single_Cell_Perturbation_Prediction_test_utils.py +67 -0
- 02_Single_Cell/README.md +11 -0
- 04_CITRIS_Intervention/CITRIS_Interventional_Pong/README.md +36 -0
- 04_CITRIS_Intervention/CITRIS_Interventional_Pong/checkpoints/CITRIS_Interventional_Pong_final_checkpoint.pt +3 -0
- 04_CITRIS_Intervention/README.md +12 -0
- 04_CITRIS_Intervention/six_step_one_step/README.md +35 -0
- 04_CITRIS_Intervention/six_step_one_step/checkpoints/seed11/h6_dense.pt +3 -0
- 04_CITRIS_Intervention/six_step_one_step/checkpoints/seed11/h6_opf.pt +3 -0
- 04_CITRIS_Intervention/six_step_one_step/checkpoints/seed23/h6_dense.pt +3 -0
- 04_CITRIS_Intervention/six_step_one_step/checkpoints/seed23/h6_opf.pt +3 -0
- 04_CITRIS_Intervention/six_step_one_step/checkpoints/seed37/h6_dense.pt +3 -0
- 04_CITRIS_Intervention/six_step_one_step/checkpoints/seed37/h6_opf.pt +3 -0
- 04_CITRIS_Intervention/six_step_one_step/checkpoints/seed53/h6_dense.pt +3 -0
- 04_CITRIS_Intervention/six_step_one_step/checkpoints/seed53/h6_opf.pt +3 -0
- 04_CITRIS_Intervention/six_step_one_step/checkpoints/seed71/h6_dense.pt +3 -0
- 04_CITRIS_Intervention/six_step_one_step/checkpoints/seed71/h6_opf.pt +3 -0
- 04_CITRIS_Intervention/six_step_one_step/test/evaluate.py +60 -0
- 05_Ten_Task_Dynamics/CLEVRER_Visual_Dynamics/README.md +36 -0
- 05_Ten_Task_Dynamics/CLEVRER_Visual_Dynamics/checkpoints/CLEVRER_Visual_Dynamics_final_checkpoint.pt +3 -0
- 05_Ten_Task_Dynamics/CausalWorld_Pushing/README.md +36 -0
- 05_Ten_Task_Dynamics/CausalWorld_Pushing/checkpoints/CausalWorld_Pushing_final_checkpoint.pt +3 -0
- 05_Ten_Task_Dynamics/CausalWorld_Pushing_Long_Horizon/README.md +36 -0
- 05_Ten_Task_Dynamics/CausalWorld_Pushing_Long_Horizon/checkpoints/CausalWorld_Pushing_Long_Horizon_final_checkpoint.pt +3 -0
- 05_Ten_Task_Dynamics/DMC_Cartpole_Pixels/README.md +36 -0
- 05_Ten_Task_Dynamics/DMC_Cartpole_Pixels/checkpoints/DMC_Cartpole_Pixels_final_checkpoint.pt +3 -0
- 05_Ten_Task_Dynamics/DMC_VB_Cheetah_Run/README.md +36 -0
- 05_Ten_Task_Dynamics/DMC_VB_Cheetah_Run/checkpoints/DMC_VB_Cheetah_Run_final_checkpoint.pt +3 -0
- 05_Ten_Task_Dynamics/Distracting_Control_Cartpole/README.md +36 -0
- 05_Ten_Task_Dynamics/Distracting_Control_Cartpole/checkpoints/Distracting_Control_Cartpole_final_checkpoint.pt +3 -0
- 05_Ten_Task_Dynamics/PDEBench_Burgers/README.md +36 -0
- 05_Ten_Task_Dynamics/PDEBench_Burgers/checkpoints/PDEBench_Burgers_final_checkpoint.pt +3 -0
- 05_Ten_Task_Dynamics/PDEBench_Navier_Stokes/README.md +36 -0
- 05_Ten_Task_Dynamics/PDEBench_Navier_Stokes/checkpoints/PDEBench_Navier_Stokes_final_checkpoint.pt +3 -0
- 05_Ten_Task_Dynamics/PDEBench_Reaction_Diffusion_1D/README.md +36 -0
- 05_Ten_Task_Dynamics/PDEBench_Reaction_Diffusion_1D/checkpoints/PDEBench_Reaction_Diffusion_1D_final_checkpoint.pt +3 -0
- 05_Ten_Task_Dynamics/PDEBench_Reaction_Diffusion_2D/README.md +36 -0
- 05_Ten_Task_Dynamics/PDEBench_Reaction_Diffusion_2D/checkpoints/PDEBench_Reaction_Diffusion_2D_final_checkpoint.pt +3 -0
.DS_Store
ADDED
|
Binary file (6.15 kB). View file
|
|
|
01_Vision_Controlled_Binding/README.md
ADDED
|
@@ -0,0 +1,12 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Vision and Controlled Binding
|
| 2 |
+
|
| 3 |
+
This domain asks whether a visual representation can preserve the structure of a controlled change, not just the appearance of an image. Visible patches provide context and masked patch states provide predictive targets. After JEPA pretraining, a structured readout compares a source scene with its intervention and identifies both the affected spatial support and the operation that occurred. OPF is used to separate complementary visual change modes so the representation can recombine them for support-operation pairs that were not used to fit the readout.
|
| 4 |
+
|
| 5 |
+
Open a model page for its research background, checkpoint, input format, and inference command. Additional checkpoints from this domain will be released progressively.
|
| 6 |
+
|
| 7 |
+
## Models
|
| 8 |
+
|
| 9 |
+
- [Shapes3D_JEPA_Pretraining](Shapes3D_JEPA_Pretraining/README.md)
|
| 10 |
+
- [Shapes3D_Main_Factor_Prediction](Shapes3D_Main_Factor_Prediction/README.md)
|
| 11 |
+
|
| 12 |
+
[Back to the project overview](../README.md)
|
01_Vision_Controlled_Binding/Shapes3D_JEPA_Pretraining/README.md
ADDED
|
@@ -0,0 +1,36 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Shapes3D JEPA Pretraining
|
| 2 |
+
|
| 3 |
+
## Background
|
| 4 |
+
|
| 5 |
+
This model learns token-level visual representations from controlled changes in object shape, color, scale, orientation, and scene appearance. The archived checkpoint provides the encoder used before factor binding.
|
| 6 |
+
|
| 7 |
+
## Files
|
| 8 |
+
|
| 9 |
+
- Model implementation: [`../../model`](<../../model>)
|
| 10 |
+
- Checkpoint: [`checkpoints/Shapes3D_JEPA_Pretraining_final_checkpoint.pt`](checkpoints/Shapes3D_JEPA_Pretraining_final_checkpoint.pt)
|
| 11 |
+
- Inference command: [`../../test/infer_vision.py`](<../../test/infer_vision.py>)
|
| 12 |
+
|
| 13 |
+
## Run inference
|
| 14 |
+
|
| 15 |
+
Run from this experiment directory:
|
| 16 |
+
|
| 17 |
+
```bash
|
| 18 |
+
python ../../test/infer_vision.py \
|
| 19 |
+
--checkpoint checkpoints/Shapes3D_JEPA_Pretraining_final_checkpoint.pt \
|
| 20 |
+
--input <preprocessed_images.npz> \
|
| 21 |
+
--output <new_output_dir>/tokens.npz
|
| 22 |
+
```
|
| 23 |
+
|
| 24 |
+
Provide an `images` array with shape `[N, 64, 64, 3]` or `[N, 3, 64, 64]`. Values may be uint8 in `[0, 255]` or floats in `[0, 1]`. The command writes visual token embeddings. The Shapes3D release will progressively add the task-specific SO-OPF readout checkpoint alongside the encoder checkpoints.
|
| 25 |
+
|
| 26 |
+
For a quick checkpoint and tensor-shape check, omit `--input`:
|
| 27 |
+
|
| 28 |
+
```bash
|
| 29 |
+
python ../../test/infer_vision.py \
|
| 30 |
+
--checkpoint checkpoints/Shapes3D_JEPA_Pretraining_final_checkpoint.pt \
|
| 31 |
+
--output /tmp/jepa_inference_smoke/tokens.npz
|
| 32 |
+
```
|
| 33 |
+
|
| 34 |
+
This command reports the predicted tensor shape and finite-output status. Dataset runs also report MSE when targets are supplied.
|
| 35 |
+
|
| 36 |
+
[Back to the project overview](<../../README.md>)
|
01_Vision_Controlled_Binding/Shapes3D_JEPA_Pretraining/checkpoints/Shapes3D_JEPA_Pretraining_final_checkpoint.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:74138fc4a4b537d89ac902e89ca140fc99159e2d1e0c6f883d3ca7952a1dba7d
|
| 3 |
+
size 13441074
|
01_Vision_Controlled_Binding/Shapes3D_Main_Factor_Prediction/README.md
ADDED
|
@@ -0,0 +1,36 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Shapes3D Factor Prediction
|
| 2 |
+
|
| 3 |
+
## Background
|
| 4 |
+
|
| 5 |
+
This stage connects the Shapes3D visual encoder to support-by-operation factor prediction. It asks whether a representation can identify what changed and where the change occurred.
|
| 6 |
+
|
| 7 |
+
## Files
|
| 8 |
+
|
| 9 |
+
- Model implementation: [`../../model`](<../../model>)
|
| 10 |
+
- Checkpoint: [`checkpoints/Shapes3D_Main_Factor_Prediction_final_checkpoint.pt`](checkpoints/Shapes3D_Main_Factor_Prediction_final_checkpoint.pt)
|
| 11 |
+
- Inference command: [`../../test/infer_vision.py`](<../../test/infer_vision.py>)
|
| 12 |
+
|
| 13 |
+
## Run inference
|
| 14 |
+
|
| 15 |
+
Run from this experiment directory:
|
| 16 |
+
|
| 17 |
+
```bash
|
| 18 |
+
python ../../test/infer_vision.py \
|
| 19 |
+
--checkpoint checkpoints/Shapes3D_Main_Factor_Prediction_final_checkpoint.pt \
|
| 20 |
+
--input <preprocessed_images.npz> \
|
| 21 |
+
--output <new_output_dir>/tokens.npz
|
| 22 |
+
```
|
| 23 |
+
|
| 24 |
+
Provide an `images` array with shape `[N, 64, 64, 3]` or `[N, 3, 64, 64]`. Values may be uint8 in `[0, 255]` or floats in `[0, 1]`. The command writes visual token embeddings. The Shapes3D release will progressively add the task-specific SO-OPF readout checkpoint alongside the encoder checkpoints.
|
| 25 |
+
|
| 26 |
+
For a quick checkpoint and tensor-shape check, omit `--input`:
|
| 27 |
+
|
| 28 |
+
```bash
|
| 29 |
+
python ../../test/infer_vision.py \
|
| 30 |
+
--checkpoint checkpoints/Shapes3D_Main_Factor_Prediction_final_checkpoint.pt \
|
| 31 |
+
--output /tmp/jepa_inference_smoke/tokens.npz
|
| 32 |
+
```
|
| 33 |
+
|
| 34 |
+
This command reports the predicted tensor shape and finite-output status. Dataset runs also report MSE when targets are supplied.
|
| 35 |
+
|
| 36 |
+
[Back to the project overview](<../../README.md>)
|
01_Vision_Controlled_Binding/Shapes3D_Main_Factor_Prediction/checkpoints/Shapes3D_Main_Factor_Prediction_final_checkpoint.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:cb45ceb88443d21d3f62e95b8a5e9495cb30b5d04e8aa9725954caecbca020e7
|
| 3 |
+
size 11157882
|
02_Single_Cell/Norman_Single_Cell_Perturbation_Prediction/README.md
ADDED
|
@@ -0,0 +1,44 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Norman Single-Cell Perturbation Prediction
|
| 2 |
+
|
| 3 |
+
## Background
|
| 4 |
+
|
| 5 |
+
This model predicts how gene expression changes after a perturbation. It combines a single-cell backbone with factorized perturbation-response representations, then decodes the predicted expression profile.
|
| 6 |
+
|
| 7 |
+
## Files and requirements
|
| 8 |
+
|
| 9 |
+
- Model code: [`model/cell_jepa`](model/cell_jepa)
|
| 10 |
+
- Vocabulary: [`model/Norman_Single_Cell_Perturbation_Prediction_vocab.json`](model/Norman_Single_Cell_Perturbation_Prediction_vocab.json)
|
| 11 |
+
- Checkpoint: [`checkpoints/Norman_Single_Cell_Perturbation_Prediction_final_checkpoint.pt`](checkpoints/Norman_Single_Cell_Perturbation_Prediction_final_checkpoint.pt)
|
| 12 |
+
- Inference and metric entry point: [`test/Norman_Single_Cell_Perturbation_Prediction_test.py`](test/Norman_Single_Cell_Perturbation_Prediction_test.py)
|
| 13 |
+
|
| 14 |
+
The model uses scGPT for the expression backbone and GEARS for the Norman evaluation pipeline.
|
| 15 |
+
|
| 16 |
+
## Verify checkpoint loading
|
| 17 |
+
|
| 18 |
+
Use checkpoint mode to reconstruct the model and report its parameter and factor configuration:
|
| 19 |
+
|
| 20 |
+
```bash
|
| 21 |
+
python test/Norman_Single_Cell_Perturbation_Prediction_test.py \
|
| 22 |
+
--repo <cell_jepa_source_repo> \
|
| 23 |
+
--checkpoint checkpoints/Norman_Single_Cell_Perturbation_Prediction_final_checkpoint.pt \
|
| 24 |
+
--vocab model/Norman_Single_Cell_Perturbation_Prediction_vocab.json \
|
| 25 |
+
--checkpoint-only \
|
| 26 |
+
--output <new_output_dir>/norman_checkpoint.json \
|
| 27 |
+
--device cpu
|
| 28 |
+
```
|
| 29 |
+
|
| 30 |
+
## Evaluate on Norman
|
| 31 |
+
|
| 32 |
+
```bash
|
| 33 |
+
python test/Norman_Single_Cell_Perturbation_Prediction_test.py \
|
| 34 |
+
--repo <cell_jepa_source_repo> \
|
| 35 |
+
--data-root <norman_data_root> \
|
| 36 |
+
--checkpoint checkpoints/Norman_Single_Cell_Perturbation_Prediction_final_checkpoint.pt \
|
| 37 |
+
--vocab model/Norman_Single_Cell_Perturbation_Prediction_vocab.json \
|
| 38 |
+
--output <new_output_dir>/norman_metrics.json \
|
| 39 |
+
--device cuda
|
| 40 |
+
```
|
| 41 |
+
|
| 42 |
+
The evaluation command predicts expression profiles on the prepared Norman split and writes the perturbation metrics to JSON.
|
| 43 |
+
|
| 44 |
+
[Back to the project overview](../../README.md)
|
02_Single_Cell/Norman_Single_Cell_Perturbation_Prediction/checkpoints/Norman_Single_Cell_Perturbation_Prediction_final_checkpoint.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:ee029201f537faea600ec9d8f7c5dfad9b0a7eaac747eb9bc60b407f70f3dcb7
|
| 3 |
+
size 419246117
|
02_Single_Cell/Norman_Single_Cell_Perturbation_Prediction/model/Norman_Single_Cell_Perturbation_Prediction_vocab.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
02_Single_Cell/Norman_Single_Cell_Perturbation_Prediction/model/cell_jepa/__init__.py
ADDED
|
@@ -0,0 +1,12 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Cell-JEPA modules used by the reproduction workflow."""
|
| 2 |
+
|
| 3 |
+
from .objectives import CellJEPA, PerturbationJEPA, supervised_ecs_loss
|
| 4 |
+
from .objectives_anything import CellJEPAAnything, PerturbationJEPAAnything
|
| 5 |
+
|
| 6 |
+
__all__ = [
|
| 7 |
+
"CellJEPA",
|
| 8 |
+
"CellJEPAAnything",
|
| 9 |
+
"PerturbationJEPA",
|
| 10 |
+
"PerturbationJEPAAnything",
|
| 11 |
+
"supervised_ecs_loss",
|
| 12 |
+
]
|
02_Single_Cell/Norman_Single_Cell_Perturbation_Prediction/model/cell_jepa/objectives.py
ADDED
|
@@ -0,0 +1,220 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""JEPA losses wrapped around the official scGPT model implementations.
|
| 2 |
+
|
| 3 |
+
The paper's public artifact is a manuscript, rather than an official code
|
| 4 |
+
repository. This module implements Eqs. 2.3 and 2.5 using the public scGPT
|
| 5 |
+
backbone and GEARS perturbation data protocol.
|
| 6 |
+
"""
|
| 7 |
+
|
| 8 |
+
from __future__ import annotations
|
| 9 |
+
|
| 10 |
+
import copy
|
| 11 |
+
from typing import Any, Optional
|
| 12 |
+
|
| 13 |
+
import torch
|
| 14 |
+
import torch.nn.functional as F
|
| 15 |
+
from torch import Tensor, nn
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
def _mlp(width: int) -> nn.Sequential:
|
| 19 |
+
return nn.Sequential(nn.Linear(width, width), nn.GELU(), nn.Linear(width, width))
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
def supervised_ecs_loss(embeddings: Tensor, labels: Tensor, temperature: float = 0.5) -> Tensor:
|
| 23 |
+
"""Label-aware ECS used for PBMC representation learning (paper Sec. 2.4.2).
|
| 24 |
+
|
| 25 |
+
Each cell treats all same-label cells in its minibatch as positives. Anchors
|
| 26 |
+
without another same-label cell are ignored rather than producing NaNs.
|
| 27 |
+
"""
|
| 28 |
+
if embeddings.ndim != 2 or labels.ndim != 1:
|
| 29 |
+
raise ValueError("embeddings must be [batch, width] and labels must be [batch]")
|
| 30 |
+
if embeddings.shape[0] != labels.shape[0]:
|
| 31 |
+
raise ValueError("labels and embeddings must have the same batch dimension")
|
| 32 |
+
if embeddings.shape[0] < 2:
|
| 33 |
+
return embeddings.new_zeros(())
|
| 34 |
+
logits = F.normalize(embeddings, dim=-1) @ F.normalize(embeddings, dim=-1).T
|
| 35 |
+
logits = logits / temperature
|
| 36 |
+
diagonal = torch.eye(len(labels), dtype=torch.bool, device=labels.device)
|
| 37 |
+
logits = logits.masked_fill(diagonal, float("-inf"))
|
| 38 |
+
log_prob = F.log_softmax(logits, dim=1)
|
| 39 |
+
positive = labels[:, None].eq(labels[None, :]) & ~diagonal
|
| 40 |
+
count = positive.sum(dim=1)
|
| 41 |
+
valid = count > 0
|
| 42 |
+
if not valid.any():
|
| 43 |
+
return embeddings.new_zeros(())
|
| 44 |
+
return -((log_prob.masked_fill(~positive, 0).sum(dim=1) / count.clamp_min(1))[valid]).mean()
|
| 45 |
+
|
| 46 |
+
|
| 47 |
+
def ecs_loss(embeddings: Tensor, threshold: float = 0.3) -> Tensor:
|
| 48 |
+
"""The original scGPT embedding-consistency penalty used for perturbation.
|
| 49 |
+
|
| 50 |
+
Perturbation batches contain a single cell line rather than PBMC cell-type
|
| 51 |
+
labels, so Sec. 2.5 uses the non-label-aware ECS term. This is equivalent
|
| 52 |
+
to the public scGPT implementation but is exposed here because the JEPA
|
| 53 |
+
wrapper needs the intermediate cell embeddings.
|
| 54 |
+
"""
|
| 55 |
+
if embeddings.shape[0] < 2:
|
| 56 |
+
return embeddings.new_zeros(())
|
| 57 |
+
similarity = F.normalize(embeddings, dim=-1) @ F.normalize(embeddings, dim=-1).T
|
| 58 |
+
similarity = similarity - torch.diag_embed(torch.diagonal(similarity))
|
| 59 |
+
similarity = F.relu(similarity)
|
| 60 |
+
return torch.mean(1 - (similarity - threshold).pow(2))
|
| 61 |
+
|
| 62 |
+
|
| 63 |
+
class CellJEPA(nn.Module):
|
| 64 |
+
"""Student--EMA teacher wrapper for ``scgpt.model.TransformerModel``.
|
| 65 |
+
|
| 66 |
+
Student inputs may contain a mask sentinel; the teacher always receives
|
| 67 |
+
the corresponding unmasked values. ``losses`` returns masked-value MSE
|
| 68 |
+
plus the latent cosine-distance JEPA objective from Eq. 2.3.
|
| 69 |
+
"""
|
| 70 |
+
|
| 71 |
+
def __init__(self, student: nn.Module, momentum: float = 0.996) -> None:
|
| 72 |
+
super().__init__()
|
| 73 |
+
if not 0.0 <= momentum < 1.0:
|
| 74 |
+
raise ValueError("momentum must be in [0, 1)")
|
| 75 |
+
self.student = student
|
| 76 |
+
self.teacher = copy.deepcopy(student)
|
| 77 |
+
self.predictor = _mlp(student.d_model)
|
| 78 |
+
self.momentum = momentum
|
| 79 |
+
self.freeze_teacher()
|
| 80 |
+
|
| 81 |
+
def freeze_teacher(self) -> None:
|
| 82 |
+
self.teacher.eval()
|
| 83 |
+
for parameter in self.teacher.parameters():
|
| 84 |
+
parameter.requires_grad_(False)
|
| 85 |
+
|
| 86 |
+
def train(self, mode: bool = True) -> "CellJEPA":
|
| 87 |
+
super().train(mode)
|
| 88 |
+
self.teacher.eval() # the EMA target must never use dropout updates
|
| 89 |
+
return self
|
| 90 |
+
|
| 91 |
+
@torch.no_grad()
|
| 92 |
+
def update_teacher(self) -> None:
|
| 93 |
+
"""EMA update performed once after each student optimizer step."""
|
| 94 |
+
for teacher_parameter, student_parameter in zip(
|
| 95 |
+
self.teacher.parameters(), self.student.parameters()
|
| 96 |
+
):
|
| 97 |
+
teacher_parameter.mul_(self.momentum).add_(student_parameter, alpha=1 - self.momentum)
|
| 98 |
+
for teacher_buffer, student_buffer in zip(self.teacher.buffers(), self.student.buffers()):
|
| 99 |
+
teacher_buffer.copy_(student_buffer)
|
| 100 |
+
|
| 101 |
+
def forward(
|
| 102 |
+
self,
|
| 103 |
+
gene_ids: Tensor,
|
| 104 |
+
masked_values: Tensor,
|
| 105 |
+
target_values: Tensor,
|
| 106 |
+
padding_mask: Tensor,
|
| 107 |
+
batch_labels: Optional[Tensor] = None,
|
| 108 |
+
include_gepc: bool = False,
|
| 109 |
+
) -> tuple[dict[str, Tensor], Tensor]:
|
| 110 |
+
student_output = self.student(
|
| 111 |
+
gene_ids, masked_values, padding_mask, batch_labels=batch_labels,
|
| 112 |
+
CLS=False, CCE=False, MVC=include_gepc, ECS=False,
|
| 113 |
+
)
|
| 114 |
+
with torch.no_grad():
|
| 115 |
+
teacher_output = self.teacher(
|
| 116 |
+
gene_ids, target_values, padding_mask, batch_labels=batch_labels,
|
| 117 |
+
CLS=False, CCE=False, MVC=False, ECS=False,
|
| 118 |
+
)
|
| 119 |
+
return student_output, teacher_output["cell_emb"]
|
| 120 |
+
|
| 121 |
+
def losses(
|
| 122 |
+
self,
|
| 123 |
+
gene_ids: Tensor,
|
| 124 |
+
masked_values: Tensor,
|
| 125 |
+
target_values: Tensor,
|
| 126 |
+
padding_mask: Tensor,
|
| 127 |
+
masked_positions: Tensor,
|
| 128 |
+
batch_labels: Optional[Tensor] = None,
|
| 129 |
+
cell_labels: Optional[Tensor] = None,
|
| 130 |
+
jepa_weight: float = 1.0,
|
| 131 |
+
reconstruction_weight: float = 1.0,
|
| 132 |
+
ecs_weight: float = 0.0,
|
| 133 |
+
gepc_weight: float = 0.0,
|
| 134 |
+
) -> dict[str, Tensor]:
|
| 135 |
+
student_output, teacher_embedding = self(
|
| 136 |
+
gene_ids, masked_values, target_values, padding_mask, batch_labels,
|
| 137 |
+
include_gepc=bool(gepc_weight),
|
| 138 |
+
)
|
| 139 |
+
prediction = student_output["mlm_output"]
|
| 140 |
+
reconstruction = F.mse_loss(prediction[masked_positions], target_values[masked_positions])
|
| 141 |
+
jepa = 1 - F.cosine_similarity(self.predictor(student_output["cell_emb"]), teacher_embedding, dim=-1).mean()
|
| 142 |
+
ecs = (
|
| 143 |
+
supervised_ecs_loss(student_output["cell_emb"], cell_labels)
|
| 144 |
+
if cell_labels is not None and ecs_weight
|
| 145 |
+
else reconstruction.new_zeros(())
|
| 146 |
+
)
|
| 147 |
+
gepc = (
|
| 148 |
+
F.mse_loss(student_output["mvc_output"][masked_positions], target_values[masked_positions])
|
| 149 |
+
if gepc_weight
|
| 150 |
+
else reconstruction.new_zeros(())
|
| 151 |
+
)
|
| 152 |
+
total = reconstruction_weight * reconstruction + jepa_weight * jepa + ecs_weight * ecs + gepc_weight * gepc
|
| 153 |
+
return {"total": total, "reconstruction": reconstruction, "gepc": gepc, "jepa": jepa, "ecs": ecs}
|
| 154 |
+
|
| 155 |
+
|
| 156 |
+
class PerturbationJEPA(nn.Module):
|
| 157 |
+
"""Cell-JEPA perturbation head for scGPT's ``TransformerGenerator``.
|
| 158 |
+
|
| 159 |
+
The student receives control expression and perturbation indicators. The
|
| 160 |
+
EMA teacher receives the observed post-perturbation profile and the same
|
| 161 |
+
perturbation indicators. This follows paper Sec. 2.5 with no input masking.
|
| 162 |
+
"""
|
| 163 |
+
|
| 164 |
+
def __init__(self, student: nn.Module, momentum: float = 0.996) -> None:
|
| 165 |
+
super().__init__()
|
| 166 |
+
self.student = student
|
| 167 |
+
self.teacher = copy.deepcopy(student)
|
| 168 |
+
self.predictor = _mlp(student.d_model)
|
| 169 |
+
self.momentum = momentum
|
| 170 |
+
self.freeze_teacher()
|
| 171 |
+
|
| 172 |
+
def freeze_teacher(self) -> None:
|
| 173 |
+
self.teacher.eval()
|
| 174 |
+
for parameter in self.teacher.parameters():
|
| 175 |
+
parameter.requires_grad_(False)
|
| 176 |
+
|
| 177 |
+
def train(self, mode: bool = True) -> "PerturbationJEPA":
|
| 178 |
+
super().train(mode)
|
| 179 |
+
self.teacher.eval()
|
| 180 |
+
return self
|
| 181 |
+
|
| 182 |
+
@torch.no_grad()
|
| 183 |
+
def update_teacher(self) -> None:
|
| 184 |
+
for teacher_parameter, student_parameter in zip(self.teacher.parameters(), self.student.parameters()):
|
| 185 |
+
teacher_parameter.mul_(self.momentum).add_(student_parameter, alpha=1 - self.momentum)
|
| 186 |
+
for teacher_buffer, student_buffer in zip(self.teacher.buffers(), self.student.buffers()):
|
| 187 |
+
teacher_buffer.copy_(student_buffer)
|
| 188 |
+
|
| 189 |
+
@staticmethod
|
| 190 |
+
def _run(model: nn.Module, gene_ids: Tensor, values: Tensor, perturbations: Tensor, padding_mask: Tensor) -> tuple[Tensor, Tensor]:
|
| 191 |
+
# The public TransformerGenerator keeps these intermediates private;
|
| 192 |
+
# accessing them avoids modifying the official baseline source.
|
| 193 |
+
hidden = model._encode(gene_ids, values, perturbations, padding_mask)
|
| 194 |
+
decoded = model.decoder(hidden, values)["pred"]
|
| 195 |
+
embedding = model._get_cell_emb_from_layer(hidden, values)
|
| 196 |
+
return decoded, embedding
|
| 197 |
+
|
| 198 |
+
def losses(
|
| 199 |
+
self,
|
| 200 |
+
gene_ids: Tensor,
|
| 201 |
+
control_values: Tensor,
|
| 202 |
+
perturbations: Tensor,
|
| 203 |
+
observed_values: Tensor,
|
| 204 |
+
padding_mask: Tensor,
|
| 205 |
+
jepa_weight: float = 1.0,
|
| 206 |
+
reconstruction_weight: float = 1.0,
|
| 207 |
+
ecs_weight: float = 0.8,
|
| 208 |
+
) -> dict[str, Tensor]:
|
| 209 |
+
predicted_values, student_embedding = self._run(
|
| 210 |
+
self.student, gene_ids, control_values, perturbations, padding_mask
|
| 211 |
+
)
|
| 212 |
+
with torch.no_grad():
|
| 213 |
+
_, teacher_embedding = self._run(
|
| 214 |
+
self.teacher, gene_ids, observed_values, perturbations, padding_mask
|
| 215 |
+
)
|
| 216 |
+
reconstruction = F.mse_loss(predicted_values, observed_values)
|
| 217 |
+
jepa = 1 - F.cosine_similarity(self.predictor(student_embedding), teacher_embedding, dim=-1).mean()
|
| 218 |
+
consistency = ecs_loss(student_embedding) if ecs_weight else reconstruction.new_zeros(())
|
| 219 |
+
total = reconstruction_weight * reconstruction + jepa_weight * jepa + ecs_weight * consistency
|
| 220 |
+
return {"total": total, "reconstruction": reconstruction, "jepa": jepa, "ecs": consistency}
|
02_Single_Cell/Norman_Single_Cell_Perturbation_Prediction/model/cell_jepa/objectives_anything.py
ADDED
|
@@ -0,0 +1,558 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""JEPA-Anything with response factors and activation-level diversity."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
import copy
|
| 6 |
+
import math
|
| 7 |
+
|
| 8 |
+
import torch
|
| 9 |
+
import torch.nn.functional as F
|
| 10 |
+
from torch import Tensor, nn
|
| 11 |
+
|
| 12 |
+
from .objectives import ecs_loss, supervised_ecs_loss
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
class PerturbationJEPAAnything(nn.Module):
|
| 16 |
+
"""Predict orthogonal perturbation-response factors and use them for decoding.
|
| 17 |
+
|
| 18 |
+
The objective targets the paired response
|
| 19 |
+
``teacher(observed) - teacher(control)``, penalizes cross-factor
|
| 20 |
+
activation covariance, learns sparse balanced routing, and injects the
|
| 21 |
+
routed factor prediction into the expression decoder.
|
| 22 |
+
"""
|
| 23 |
+
|
| 24 |
+
def __init__(
|
| 25 |
+
self,
|
| 26 |
+
student: nn.Module,
|
| 27 |
+
num_factors: int = 4,
|
| 28 |
+
momentum: float = 0.996,
|
| 29 |
+
) -> None:
|
| 30 |
+
super().__init__()
|
| 31 |
+
width = int(student.d_model)
|
| 32 |
+
if num_factors < 2 or width % num_factors:
|
| 33 |
+
raise ValueError("num_factors must be >=2 and divide student.d_model")
|
| 34 |
+
if not 0.0 <= momentum < 1.0:
|
| 35 |
+
raise ValueError("momentum must be in [0, 1)")
|
| 36 |
+
|
| 37 |
+
self.student = student
|
| 38 |
+
self.teacher = copy.deepcopy(student)
|
| 39 |
+
self.width = width
|
| 40 |
+
self.num_factors = num_factors
|
| 41 |
+
self.factor_width = width // num_factors
|
| 42 |
+
self.momentum = momentum
|
| 43 |
+
|
| 44 |
+
identity = torch.eye(width).reshape(width, num_factors, self.factor_width)
|
| 45 |
+
self.projectors = nn.Parameter(identity.permute(1, 0, 2).contiguous())
|
| 46 |
+
self.predictor_trunk = nn.Sequential(nn.Linear(width, width), nn.GELU())
|
| 47 |
+
self.factor_heads = nn.ModuleList(
|
| 48 |
+
nn.Linear(width, self.factor_width) for _ in range(num_factors)
|
| 49 |
+
)
|
| 50 |
+
self.gate = nn.Linear(width, num_factors)
|
| 51 |
+
self.response_to_hidden = nn.Linear(width, width)
|
| 52 |
+
nn.init.zeros_(self.response_to_hidden.weight)
|
| 53 |
+
nn.init.zeros_(self.response_to_hidden.bias)
|
| 54 |
+
self.freeze_teacher()
|
| 55 |
+
|
| 56 |
+
def freeze_teacher(self) -> None:
|
| 57 |
+
self.teacher.eval()
|
| 58 |
+
for parameter in self.teacher.parameters():
|
| 59 |
+
parameter.requires_grad_(False)
|
| 60 |
+
|
| 61 |
+
def train(self, mode: bool = True) -> "PerturbationJEPAAnything":
|
| 62 |
+
super().train(mode)
|
| 63 |
+
self.teacher.eval()
|
| 64 |
+
return self
|
| 65 |
+
|
| 66 |
+
@torch.no_grad()
|
| 67 |
+
def update_teacher(self) -> None:
|
| 68 |
+
for teacher_parameter, student_parameter in zip(
|
| 69 |
+
self.teacher.parameters(), self.student.parameters()
|
| 70 |
+
):
|
| 71 |
+
teacher_parameter.mul_(self.momentum).add_(
|
| 72 |
+
student_parameter, alpha=1 - self.momentum
|
| 73 |
+
)
|
| 74 |
+
for teacher_buffer, student_buffer in zip(
|
| 75 |
+
self.teacher.buffers(), self.student.buffers()
|
| 76 |
+
):
|
| 77 |
+
teacher_buffer.copy_(student_buffer)
|
| 78 |
+
|
| 79 |
+
@torch.no_grad()
|
| 80 |
+
def retract_projectors(self) -> None:
|
| 81 |
+
"""Project the concatenated blocks back onto the orthogonal group."""
|
| 82 |
+
full = torch.cat(tuple(self.projectors), dim=1)
|
| 83 |
+
q, r = torch.linalg.qr(full)
|
| 84 |
+
signs = torch.sign(torch.diagonal(r))
|
| 85 |
+
signs[signs == 0] = 1
|
| 86 |
+
q = q * signs.unsqueeze(0)
|
| 87 |
+
blocks = q.split(self.factor_width, dim=1)
|
| 88 |
+
self.projectors.copy_(torch.stack(blocks, dim=0))
|
| 89 |
+
|
| 90 |
+
@staticmethod
|
| 91 |
+
def _encode(
|
| 92 |
+
model: nn.Module,
|
| 93 |
+
gene_ids: Tensor,
|
| 94 |
+
values: Tensor,
|
| 95 |
+
perturbations: Tensor,
|
| 96 |
+
padding_mask: Tensor,
|
| 97 |
+
) -> tuple[Tensor, Tensor]:
|
| 98 |
+
hidden = model._encode(gene_ids, values, perturbations, padding_mask)
|
| 99 |
+
embedding = model._get_cell_emb_from_layer(hidden, values)
|
| 100 |
+
return hidden, embedding
|
| 101 |
+
|
| 102 |
+
def _predict_from_hidden(
|
| 103 |
+
self,
|
| 104 |
+
hidden: Tensor,
|
| 105 |
+
embedding: Tensor,
|
| 106 |
+
control_values: Tensor,
|
| 107 |
+
) -> tuple[Tensor, Tensor, Tensor]:
|
| 108 |
+
shared = self.predictor_trunk(embedding)
|
| 109 |
+
predicted_factors = torch.stack(
|
| 110 |
+
[head(shared) for head in self.factor_heads], dim=1
|
| 111 |
+
)
|
| 112 |
+
gates = torch.softmax(self.gate(embedding), dim=1)
|
| 113 |
+
routed = predicted_factors * gates.unsqueeze(-1) * self.num_factors
|
| 114 |
+
response_context = self.response_to_hidden(routed.flatten(1))
|
| 115 |
+
predicted_values = self.student.decoder(
|
| 116 |
+
hidden + response_context.unsqueeze(1), control_values
|
| 117 |
+
)["pred"]
|
| 118 |
+
return predicted_values, predicted_factors, gates
|
| 119 |
+
|
| 120 |
+
def predict_expression(
|
| 121 |
+
self,
|
| 122 |
+
gene_ids: Tensor,
|
| 123 |
+
control_values: Tensor,
|
| 124 |
+
perturbations: Tensor,
|
| 125 |
+
padding_mask: Tensor,
|
| 126 |
+
) -> Tensor:
|
| 127 |
+
hidden, embedding = self._encode(
|
| 128 |
+
self.student, gene_ids, control_values, perturbations, padding_mask
|
| 129 |
+
)
|
| 130 |
+
predicted_values, _, _ = self._predict_from_hidden(
|
| 131 |
+
hidden, embedding, control_values
|
| 132 |
+
)
|
| 133 |
+
return predicted_values
|
| 134 |
+
|
| 135 |
+
def _block_orthogonality_loss(self) -> Tensor:
|
| 136 |
+
identity = torch.eye(
|
| 137 |
+
self.factor_width,
|
| 138 |
+
device=self.projectors.device,
|
| 139 |
+
dtype=self.projectors.dtype,
|
| 140 |
+
)
|
| 141 |
+
within = torch.stack(
|
| 142 |
+
[
|
| 143 |
+
(projector.T @ projector - identity).pow(2).mean()
|
| 144 |
+
for projector in self.projectors
|
| 145 |
+
]
|
| 146 |
+
).mean()
|
| 147 |
+
cross_terms = []
|
| 148 |
+
for i in range(self.num_factors):
|
| 149 |
+
for j in range(i + 1, self.num_factors):
|
| 150 |
+
cross_terms.append(
|
| 151 |
+
(self.projectors[i].T @ self.projectors[j]).pow(2).mean()
|
| 152 |
+
)
|
| 153 |
+
cross = torch.stack(cross_terms).mean()
|
| 154 |
+
return within + cross
|
| 155 |
+
|
| 156 |
+
@staticmethod
|
| 157 |
+
def _cross_factor_covariance_loss(factors: Tensor) -> Tensor:
|
| 158 |
+
"""VICReg-style cross-covariance between different factor blocks."""
|
| 159 |
+
batch_size = factors.shape[0]
|
| 160 |
+
centered = factors - factors.mean(dim=0, keepdim=True)
|
| 161 |
+
std = centered.std(dim=0, unbiased=False, keepdim=True).clamp_min(1e-4)
|
| 162 |
+
normalized = centered / std
|
| 163 |
+
covariance = torch.einsum("bkr,bls->klrs", normalized, normalized)
|
| 164 |
+
covariance = covariance / max(batch_size - 1, 1)
|
| 165 |
+
n_factors = factors.shape[1]
|
| 166 |
+
mask = ~torch.eye(n_factors, dtype=torch.bool, device=factors.device)
|
| 167 |
+
return covariance[mask].pow(2).mean()
|
| 168 |
+
|
| 169 |
+
def losses(
|
| 170 |
+
self,
|
| 171 |
+
gene_ids: Tensor,
|
| 172 |
+
control_values: Tensor,
|
| 173 |
+
perturbations: Tensor,
|
| 174 |
+
observed_values: Tensor,
|
| 175 |
+
padding_mask: Tensor,
|
| 176 |
+
reconstruction_weight: float = 1.0,
|
| 177 |
+
prediction_weight: float = 1.0,
|
| 178 |
+
ecs_weight: float = 0.8,
|
| 179 |
+
orthogonality_weight: float = 1.0,
|
| 180 |
+
cross_covariance_weight: float = 1.0,
|
| 181 |
+
factor_activity_weight: float = 1.0,
|
| 182 |
+
encoder_variance_weight: float = 1.0,
|
| 183 |
+
gate_entropy_weight: float = 0.05,
|
| 184 |
+
gate_balance_weight: float = 0.05,
|
| 185 |
+
factor_std_floor: float = 0.02,
|
| 186 |
+
encoder_std_floor: float = 0.1,
|
| 187 |
+
) -> dict[str, Tensor]:
|
| 188 |
+
student_hidden, student_embedding = self._encode(
|
| 189 |
+
self.student, gene_ids, control_values, perturbations, padding_mask
|
| 190 |
+
)
|
| 191 |
+
predicted_values, predicted_factors, gates = self._predict_from_hidden(
|
| 192 |
+
student_hidden, student_embedding, control_values
|
| 193 |
+
)
|
| 194 |
+
# The JEPA target is a difference between two nearby teacher
|
| 195 |
+
# embeddings. Computing either embedding in BF16 causes cancellation
|
| 196 |
+
# and quantization noise in that small response, so the teacher branch
|
| 197 |
+
# deliberately remains FP32 even when the student uses autocast.
|
| 198 |
+
with torch.no_grad(), torch.autocast(
|
| 199 |
+
device_type=control_values.device.type, enabled=False
|
| 200 |
+
):
|
| 201 |
+
_, teacher_control = self._encode(
|
| 202 |
+
self.teacher,
|
| 203 |
+
gene_ids,
|
| 204 |
+
control_values.float(),
|
| 205 |
+
perturbations,
|
| 206 |
+
padding_mask,
|
| 207 |
+
)
|
| 208 |
+
_, teacher_observed = self._encode(
|
| 209 |
+
self.teacher,
|
| 210 |
+
gene_ids,
|
| 211 |
+
observed_values.float(),
|
| 212 |
+
perturbations,
|
| 213 |
+
padding_mask,
|
| 214 |
+
)
|
| 215 |
+
teacher_response = teacher_observed - teacher_control
|
| 216 |
+
|
| 217 |
+
# Keep the expensive transformer/decoder path under the caller's
|
| 218 |
+
# autocast, but compute all covariance, variance, routing, and scalar
|
| 219 |
+
# losses in FP32. These statistics are sensitive to BF16 rounding and
|
| 220 |
+
# otherwise drift even when the forward activations remain stable.
|
| 221 |
+
with torch.autocast(device_type=control_values.device.type, enabled=False):
|
| 222 |
+
predicted_values_fp32 = predicted_values.float()
|
| 223 |
+
predicted_factors_fp32 = predicted_factors.float()
|
| 224 |
+
observed_values_fp32 = observed_values.float()
|
| 225 |
+
student_embedding_fp32 = student_embedding.float()
|
| 226 |
+
gates_fp32 = gates.float()
|
| 227 |
+
teacher_response_fp32 = teacher_response.detach().float()
|
| 228 |
+
target_factors = torch.einsum(
|
| 229 |
+
"bd,kdr->bkr", teacher_response_fp32, self.projectors.float()
|
| 230 |
+
)
|
| 231 |
+
reconstruction = F.mse_loss(
|
| 232 |
+
predicted_values_fp32, observed_values_fp32
|
| 233 |
+
)
|
| 234 |
+
prediction = F.mse_loss(predicted_factors_fp32, target_factors)
|
| 235 |
+
orthogonality = self._block_orthogonality_loss()
|
| 236 |
+
cross_covariance = self._cross_factor_covariance_loss(target_factors)
|
| 237 |
+
factor_std = target_factors.std(dim=0, unbiased=False)
|
| 238 |
+
encoder_std = student_embedding_fp32.std(dim=0, unbiased=False)
|
| 239 |
+
factor_activity = F.relu(factor_std_floor - factor_std).mean()
|
| 240 |
+
encoder_variance = F.relu(encoder_std_floor - encoder_std).mean()
|
| 241 |
+
consistency = (
|
| 242 |
+
ecs_loss(student_embedding_fp32)
|
| 243 |
+
if ecs_weight
|
| 244 |
+
else reconstruction.new_zeros(())
|
| 245 |
+
)
|
| 246 |
+
gate_entropy = -(
|
| 247 |
+
gates_fp32.clamp_min(1e-8)
|
| 248 |
+
* gates_fp32.clamp_min(1e-8).log()
|
| 249 |
+
).sum(dim=1).mean() / math.log(self.num_factors)
|
| 250 |
+
mean_gate = gates_fp32.mean(dim=0)
|
| 251 |
+
gate_balance = self.num_factors * (
|
| 252 |
+
mean_gate - 1 / self.num_factors
|
| 253 |
+
).pow(2).sum()
|
| 254 |
+
effective_factors = torch.exp(
|
| 255 |
+
-(
|
| 256 |
+
gates_fp32.clamp_min(1e-8)
|
| 257 |
+
* gates_fp32.clamp_min(1e-8).log()
|
| 258 |
+
).sum(dim=1)
|
| 259 |
+
).mean()
|
| 260 |
+
|
| 261 |
+
total = (
|
| 262 |
+
reconstruction_weight * reconstruction
|
| 263 |
+
+ prediction_weight * prediction
|
| 264 |
+
+ ecs_weight * consistency
|
| 265 |
+
+ orthogonality_weight * orthogonality
|
| 266 |
+
+ cross_covariance_weight * cross_covariance
|
| 267 |
+
+ factor_activity_weight * factor_activity
|
| 268 |
+
+ encoder_variance_weight * encoder_variance
|
| 269 |
+
+ gate_entropy_weight * gate_entropy
|
| 270 |
+
+ gate_balance_weight * gate_balance
|
| 271 |
+
)
|
| 272 |
+
return {
|
| 273 |
+
"total": total,
|
| 274 |
+
"reconstruction": reconstruction,
|
| 275 |
+
"prediction": prediction,
|
| 276 |
+
"ecs": consistency,
|
| 277 |
+
"orthogonality": orthogonality,
|
| 278 |
+
"cross_covariance": cross_covariance,
|
| 279 |
+
"factor_activity": factor_activity,
|
| 280 |
+
"encoder_variance": encoder_variance,
|
| 281 |
+
"gate_entropy": gate_entropy,
|
| 282 |
+
"gate_balance": gate_balance,
|
| 283 |
+
"effective_factors": effective_factors,
|
| 284 |
+
"factor_std_min": factor_std.min(),
|
| 285 |
+
"factor_std_mean": factor_std.mean(),
|
| 286 |
+
"encoder_std_min": encoder_std.min(),
|
| 287 |
+
"encoder_std_mean": encoder_std.mean(),
|
| 288 |
+
}
|
| 289 |
+
|
| 290 |
+
@torch.no_grad()
|
| 291 |
+
def factor_diagnostics(self) -> dict[str, float]:
|
| 292 |
+
full = torch.cat(tuple(self.projectors), dim=1).float()
|
| 293 |
+
gram = full.T @ full
|
| 294 |
+
identity = torch.eye(self.width, device=full.device)
|
| 295 |
+
singular = torch.linalg.svdvals(full)
|
| 296 |
+
return {
|
| 297 |
+
"orthogonality_rmse": float((gram - identity).pow(2).mean().sqrt()),
|
| 298 |
+
"orthogonality_max_abs": float((gram - identity).abs().max()),
|
| 299 |
+
"projection_singular_min": float(singular.min()),
|
| 300 |
+
"projection_singular_max": float(singular.max()),
|
| 301 |
+
"projection_condition_number": float(singular.max() / singular.min()),
|
| 302 |
+
}
|
| 303 |
+
|
| 304 |
+
|
| 305 |
+
class CellJEPAAnything(nn.Module):
|
| 306 |
+
"""JEPA-Anything objective for cell-type representation learning.
|
| 307 |
+
|
| 308 |
+
Unlike the perturbation objective above, PBMC has no observed/control pair.
|
| 309 |
+
The EMA teacher therefore supplies the target cell representation while the
|
| 310 |
+
student sees the masked expression profile. The representation is routed
|
| 311 |
+
through orthogonal response factors so this mode uses the same
|
| 312 |
+
JEPA-Anything decomposition and diversity constraints as the perturbation
|
| 313 |
+
experiments.
|
| 314 |
+
"""
|
| 315 |
+
|
| 316 |
+
def __init__(self, student: nn.Module, num_factors: int = 4, momentum: float = 0.996) -> None:
|
| 317 |
+
super().__init__()
|
| 318 |
+
width = int(student.d_model)
|
| 319 |
+
if num_factors < 2 or width % num_factors:
|
| 320 |
+
raise ValueError("num_factors must be >=2 and divide student.d_model")
|
| 321 |
+
self.student = student
|
| 322 |
+
self.teacher = copy.deepcopy(student)
|
| 323 |
+
self.width = width
|
| 324 |
+
self.num_factors = num_factors
|
| 325 |
+
self.factor_width = width // num_factors
|
| 326 |
+
self.momentum = momentum
|
| 327 |
+
identity = torch.eye(width).reshape(width, num_factors, self.factor_width)
|
| 328 |
+
self.projectors = nn.Parameter(identity.permute(1, 0, 2).contiguous())
|
| 329 |
+
self.predictor_trunk = nn.Sequential(nn.Linear(width, width), nn.GELU())
|
| 330 |
+
self.factor_heads = nn.ModuleList(nn.Linear(width, self.factor_width) for _ in range(num_factors))
|
| 331 |
+
self.gate = nn.Linear(width, num_factors)
|
| 332 |
+
self.freeze_teacher()
|
| 333 |
+
|
| 334 |
+
def freeze_teacher(self) -> None:
|
| 335 |
+
self.teacher.eval()
|
| 336 |
+
for parameter in self.teacher.parameters():
|
| 337 |
+
parameter.requires_grad_(False)
|
| 338 |
+
|
| 339 |
+
def train(self, mode: bool = True) -> "CellJEPAAnything":
|
| 340 |
+
super().train(mode)
|
| 341 |
+
self.teacher.eval()
|
| 342 |
+
return self
|
| 343 |
+
|
| 344 |
+
@torch.no_grad()
|
| 345 |
+
def update_teacher(self) -> None:
|
| 346 |
+
for teacher_parameter, student_parameter in zip(self.teacher.parameters(), self.student.parameters()):
|
| 347 |
+
teacher_parameter.mul_(self.momentum).add_(student_parameter, alpha=1 - self.momentum)
|
| 348 |
+
for teacher_buffer, student_buffer in zip(self.teacher.buffers(), self.student.buffers()):
|
| 349 |
+
teacher_buffer.copy_(student_buffer)
|
| 350 |
+
|
| 351 |
+
@torch.no_grad()
|
| 352 |
+
def retract_projectors(self) -> None:
|
| 353 |
+
full = torch.cat(tuple(self.projectors), dim=1)
|
| 354 |
+
q, r = torch.linalg.qr(full)
|
| 355 |
+
signs = torch.sign(torch.diagonal(r)).masked_fill(torch.diagonal(r) == 0, 1)
|
| 356 |
+
q = q * signs.unsqueeze(0)
|
| 357 |
+
self.projectors.copy_(torch.stack(q.split(self.factor_width, dim=1), dim=0))
|
| 358 |
+
|
| 359 |
+
@staticmethod
|
| 360 |
+
def _cross_factor_covariance_loss(factors: Tensor) -> Tensor:
|
| 361 |
+
centered = factors - factors.mean(dim=0, keepdim=True)
|
| 362 |
+
std = centered.std(dim=0, unbiased=False, keepdim=True).clamp_min(1e-4)
|
| 363 |
+
normalized = centered / std
|
| 364 |
+
covariance = torch.einsum("bkr,bls->klrs", normalized, normalized)
|
| 365 |
+
covariance = covariance / max(factors.shape[0] - 1, 1)
|
| 366 |
+
mask = ~torch.eye(factors.shape[1], dtype=torch.bool, device=factors.device)
|
| 367 |
+
return covariance[mask].pow(2).mean()
|
| 368 |
+
|
| 369 |
+
def losses(
|
| 370 |
+
self, gene_ids: Tensor, masked_values: Tensor, target_values: Tensor,
|
| 371 |
+
padding_mask: Tensor, masked_positions: Tensor, batch_labels: Tensor | None = None,
|
| 372 |
+
cell_labels: Tensor | None = None, reconstruction_weight: float = 1.0,
|
| 373 |
+
jepa_weight: float = 1.0, ecs_weight: float = 1.0,
|
| 374 |
+
) -> dict[str, Tensor]:
|
| 375 |
+
# The official scGPT checkpoint used for PBMC does not enable either
|
| 376 |
+
# batch-label conditioning mode. Passing labels in that configuration
|
| 377 |
+
# makes scGPT reject the batch, so only forward them when supported.
|
| 378 |
+
batch_kwargs = {}
|
| 379 |
+
if batch_labels is not None and (
|
| 380 |
+
getattr(self.student, "use_batch_labels", False)
|
| 381 |
+
or getattr(self.student, "domain_spec_batchnorm", False)
|
| 382 |
+
):
|
| 383 |
+
batch_kwargs["batch_labels"] = batch_labels
|
| 384 |
+
student_output = self.student(
|
| 385 |
+
gene_ids, masked_values, padding_mask, **batch_kwargs,
|
| 386 |
+
CLS=False, CCE=False, MVC=True, ECS=False,
|
| 387 |
+
)
|
| 388 |
+
with torch.no_grad():
|
| 389 |
+
teacher_output = self.teacher(
|
| 390 |
+
gene_ids, target_values, padding_mask, **batch_kwargs,
|
| 391 |
+
CLS=False, CCE=False, MVC=False, ECS=False,
|
| 392 |
+
)
|
| 393 |
+
student_embedding = student_output["cell_emb"]
|
| 394 |
+
teacher_embedding = teacher_output["cell_emb"].detach()
|
| 395 |
+
predicted_factors = torch.stack(
|
| 396 |
+
[head(self.predictor_trunk(student_embedding)) for head in self.factor_heads], dim=1
|
| 397 |
+
)
|
| 398 |
+
gates = torch.softmax(self.gate(student_embedding), dim=1)
|
| 399 |
+
routed = predicted_factors * gates.unsqueeze(-1) * self.num_factors
|
| 400 |
+
target_factors = torch.einsum("bd,kdr->bkr", teacher_embedding.float(), self.projectors.float())
|
| 401 |
+
with torch.autocast(device_type=masked_values.device.type, enabled=False):
|
| 402 |
+
reconstruction = F.mse_loss(student_output["mlm_output"].float()[masked_positions], target_values.float()[masked_positions])
|
| 403 |
+
gepc = F.mse_loss(student_output["mvc_output"].float()[masked_positions], target_values.float()[masked_positions])
|
| 404 |
+
prediction = F.mse_loss(routed.float(), target_factors)
|
| 405 |
+
identity = torch.eye(self.factor_width, device=self.projectors.device)
|
| 406 |
+
within = torch.stack([(p.T @ p - identity).pow(2).mean() for p in self.projectors]).mean()
|
| 407 |
+
cross = []
|
| 408 |
+
for i in range(self.num_factors):
|
| 409 |
+
for j in range(i + 1, self.num_factors):
|
| 410 |
+
cross.append((self.projectors[i].T @ self.projectors[j]).pow(2).mean())
|
| 411 |
+
orthogonality = within + torch.stack(cross).mean()
|
| 412 |
+
cross_covariance = self._cross_factor_covariance_loss(target_factors)
|
| 413 |
+
factor_std = target_factors.std(dim=0, unbiased=False)
|
| 414 |
+
encoder_std = student_embedding.float().std(dim=0, unbiased=False)
|
| 415 |
+
factor_activity = F.relu(0.02 - factor_std).mean()
|
| 416 |
+
encoder_variance = F.relu(0.1 - encoder_std).mean()
|
| 417 |
+
ecs = supervised_ecs_loss(student_embedding.float(), cell_labels) if cell_labels is not None else student_embedding.new_zeros(())
|
| 418 |
+
entropy = -(gates.float().clamp_min(1e-8) * gates.float().clamp_min(1e-8).log()).sum(dim=1).mean() / math.log(self.num_factors)
|
| 419 |
+
mean_gate = gates.float().mean(dim=0)
|
| 420 |
+
balance = self.num_factors * (mean_gate - 1 / self.num_factors).pow(2).sum()
|
| 421 |
+
effective = torch.exp(-(gates.float().clamp_min(1e-8) * gates.float().clamp_min(1e-8).log()).sum(dim=1)).mean()
|
| 422 |
+
total = reconstruction_weight * reconstruction + jepa_weight * (prediction + 0.5 * gepc) + ecs_weight * ecs + orthogonality + cross_covariance + factor_activity + encoder_variance - 0.05 * entropy + 0.05 * balance
|
| 423 |
+
return {"total": total, "reconstruction": reconstruction, "gepc": gepc, "prediction": prediction, "ecs": ecs, "orthogonality": orthogonality, "cross_covariance": cross_covariance, "factor_activity": factor_activity, "encoder_variance": encoder_variance, "effective_factors": effective}
|
| 424 |
+
|
| 425 |
+
@torch.no_grad()
|
| 426 |
+
def factor_diagnostics(self) -> dict[str, float]:
|
| 427 |
+
full = torch.cat(tuple(self.projectors), dim=1).float()
|
| 428 |
+
gram = full.T @ full
|
| 429 |
+
singular = torch.linalg.svdvals(full)
|
| 430 |
+
return {"orthogonality_rmse": float((gram - torch.eye(self.width, device=full.device)).pow(2).mean().sqrt()), "projection_singular_min": float(singular.min()), "projection_singular_max": float(singular.max()), "effective_factors": float(self.num_factors)}
|
| 431 |
+
|
| 432 |
+
|
| 433 |
+
class CellJEPAAnythingStrict(nn.Module):
|
| 434 |
+
"""Paper-aligned OPF objective for the PBMC representation experiment.
|
| 435 |
+
|
| 436 |
+
This is the direct Eq. (3)--(10) implementation: an unmasked EMA target,
|
| 437 |
+
one independent predictor per orthogonal factor, and the orthogonality,
|
| 438 |
+
target-activity, and online-encoder-variance terms. It intentionally has
|
| 439 |
+
no routing gate, label-supervised ECS, or observation-space reconstruction;
|
| 440 |
+
those are separate extensions and are not part of the paper objective.
|
| 441 |
+
"""
|
| 442 |
+
|
| 443 |
+
def __init__(self, student: nn.Module, num_factors: int = 8, momentum: float = 0.996) -> None:
|
| 444 |
+
super().__init__()
|
| 445 |
+
width = int(student.d_model)
|
| 446 |
+
if num_factors < 2 or width % num_factors:
|
| 447 |
+
raise ValueError("num_factors must be >=2 and divide student.d_model")
|
| 448 |
+
self.student = student
|
| 449 |
+
self.teacher = copy.deepcopy(student)
|
| 450 |
+
self.width = width
|
| 451 |
+
self.num_factors = num_factors
|
| 452 |
+
self.factor_width = width // num_factors
|
| 453 |
+
self.momentum = momentum
|
| 454 |
+
identity = torch.eye(width).reshape(width, num_factors, self.factor_width)
|
| 455 |
+
self.projectors = nn.Parameter(identity.permute(1, 0, 2).contiguous())
|
| 456 |
+
self.predictor_trunk = nn.Sequential(nn.Linear(width, width), nn.GELU())
|
| 457 |
+
self.factor_heads = nn.ModuleList(
|
| 458 |
+
nn.Linear(width, self.factor_width) for _ in range(num_factors)
|
| 459 |
+
)
|
| 460 |
+
self.freeze_teacher()
|
| 461 |
+
|
| 462 |
+
def freeze_teacher(self) -> None:
|
| 463 |
+
self.teacher.eval()
|
| 464 |
+
for parameter in self.teacher.parameters():
|
| 465 |
+
parameter.requires_grad_(False)
|
| 466 |
+
|
| 467 |
+
def train(self, mode: bool = True) -> "CellJEPAAnythingStrict":
|
| 468 |
+
super().train(mode)
|
| 469 |
+
self.teacher.eval()
|
| 470 |
+
return self
|
| 471 |
+
|
| 472 |
+
@torch.no_grad()
|
| 473 |
+
def update_teacher(self) -> None:
|
| 474 |
+
for teacher_parameter, student_parameter in zip(
|
| 475 |
+
self.teacher.parameters(), self.student.parameters()
|
| 476 |
+
):
|
| 477 |
+
teacher_parameter.mul_(self.momentum).add_(
|
| 478 |
+
student_parameter, alpha=1 - self.momentum
|
| 479 |
+
)
|
| 480 |
+
for teacher_buffer, student_buffer in zip(self.teacher.buffers(), self.student.buffers()):
|
| 481 |
+
teacher_buffer.copy_(student_buffer)
|
| 482 |
+
|
| 483 |
+
@torch.no_grad()
|
| 484 |
+
def retract_projectors(self) -> None:
|
| 485 |
+
full = torch.cat(tuple(self.projectors), dim=1)
|
| 486 |
+
q, r = torch.linalg.qr(full)
|
| 487 |
+
signs = torch.sign(torch.diagonal(r)).masked_fill(torch.diagonal(r) == 0, 1)
|
| 488 |
+
self.projectors.copy_(torch.stack(q.mul(signs.unsqueeze(0)).split(self.factor_width, dim=1), dim=0))
|
| 489 |
+
|
| 490 |
+
def losses(
|
| 491 |
+
self,
|
| 492 |
+
gene_ids: Tensor,
|
| 493 |
+
masked_values: Tensor,
|
| 494 |
+
target_values: Tensor,
|
| 495 |
+
padding_mask: Tensor,
|
| 496 |
+
masked_positions: Tensor | None = None,
|
| 497 |
+
batch_labels: Tensor | None = None,
|
| 498 |
+
cell_labels: Tensor | None = None,
|
| 499 |
+
) -> dict[str, Tensor]:
|
| 500 |
+
batch_kwargs = {}
|
| 501 |
+
if batch_labels is not None and (
|
| 502 |
+
getattr(self.student, "use_batch_labels", False)
|
| 503 |
+
or getattr(self.student, "domain_spec_batchnorm", False)
|
| 504 |
+
):
|
| 505 |
+
batch_kwargs["batch_labels"] = batch_labels
|
| 506 |
+
student_output = self.student(
|
| 507 |
+
gene_ids, masked_values, padding_mask, **batch_kwargs,
|
| 508 |
+
CLS=False, CCE=False, MVC=False, ECS=False,
|
| 509 |
+
)
|
| 510 |
+
with torch.no_grad():
|
| 511 |
+
teacher_output = self.teacher(
|
| 512 |
+
gene_ids, target_values, padding_mask, **batch_kwargs,
|
| 513 |
+
CLS=False, CCE=False, MVC=False, ECS=False,
|
| 514 |
+
)
|
| 515 |
+
student_embedding = student_output["cell_emb"]
|
| 516 |
+
teacher_embedding = teacher_output["cell_emb"].detach().float()
|
| 517 |
+
predicted_factors = torch.stack(
|
| 518 |
+
[head(self.predictor_trunk(student_embedding)) for head in self.factor_heads],
|
| 519 |
+
dim=1,
|
| 520 |
+
)
|
| 521 |
+
target_factors = torch.einsum(
|
| 522 |
+
"bd,kdr->bkr", teacher_embedding, self.projectors.float()
|
| 523 |
+
)
|
| 524 |
+
with torch.autocast(device_type=masked_values.device.type, enabled=False):
|
| 525 |
+
prediction = F.mse_loss(predicted_factors.float(), target_factors)
|
| 526 |
+
identity = torch.eye(self.factor_width, device=self.projectors.device)
|
| 527 |
+
within = torch.stack(
|
| 528 |
+
[(p.T @ p - identity).pow(2).mean() for p in self.projectors]
|
| 529 |
+
).mean()
|
| 530 |
+
cross_terms = []
|
| 531 |
+
for i in range(self.num_factors):
|
| 532 |
+
for j in range(i + 1, self.num_factors):
|
| 533 |
+
cross_terms.append((self.projectors[i].T @ self.projectors[j]).pow(2).mean())
|
| 534 |
+
orthogonality = within + torch.stack(cross_terms).mean()
|
| 535 |
+
factor_std = target_factors.std(dim=0, unbiased=False)
|
| 536 |
+
encoder_std = student_embedding.float().std(dim=0, unbiased=False)
|
| 537 |
+
factor_activity = F.relu(0.02 - factor_std).mean()
|
| 538 |
+
encoder_variance = F.relu(0.1 - encoder_std).mean()
|
| 539 |
+
total = prediction + orthogonality + factor_activity + encoder_variance
|
| 540 |
+
return {
|
| 541 |
+
"total": total,
|
| 542 |
+
"prediction": prediction,
|
| 543 |
+
"orthogonality": orthogonality,
|
| 544 |
+
"factor_activity": factor_activity,
|
| 545 |
+
"encoder_variance": encoder_variance,
|
| 546 |
+
}
|
| 547 |
+
|
| 548 |
+
@torch.no_grad()
|
| 549 |
+
def factor_diagnostics(self) -> dict[str, float]:
|
| 550 |
+
full = torch.cat(tuple(self.projectors), dim=1).float()
|
| 551 |
+
gram = full.T @ full
|
| 552 |
+
singular = torch.linalg.svdvals(full)
|
| 553 |
+
return {
|
| 554 |
+
"orthogonality_rmse": float((gram - torch.eye(self.width, device=full.device)).pow(2).mean().sqrt()),
|
| 555 |
+
"projection_singular_min": float(singular.min()),
|
| 556 |
+
"projection_singular_max": float(singular.max()),
|
| 557 |
+
"effective_factors": float(self.num_factors),
|
| 558 |
+
}
|
02_Single_Cell/Norman_Single_Cell_Perturbation_Prediction/test/Norman_Single_Cell_Perturbation_Prediction_test.py
ADDED
|
@@ -0,0 +1,143 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python3
|
| 2 |
+
"""Evaluate the archived JEPA-Anything Norman checkpoint.
|
| 3 |
+
|
| 4 |
+
The archived package contains this script, its model implementation, the
|
| 5 |
+
checkpoint, and the GeneVocab. GEARS and scGPT remain runtime dependencies in
|
| 6 |
+
the original repository supplied through --repo; the Norman dataset is never
|
| 7 |
+
copied into the archive.
|
| 8 |
+
"""
|
| 9 |
+
|
| 10 |
+
from __future__ import annotations
|
| 11 |
+
|
| 12 |
+
import argparse
|
| 13 |
+
import json
|
| 14 |
+
import sys
|
| 15 |
+
from pathlib import Path
|
| 16 |
+
|
| 17 |
+
import numpy as np
|
| 18 |
+
import torch
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
def parse_args() -> argparse.Namespace:
|
| 22 |
+
parser = argparse.ArgumentParser()
|
| 23 |
+
parser.add_argument("--repo", type=Path, required=True)
|
| 24 |
+
parser.add_argument("--data-root", type=Path)
|
| 25 |
+
parser.add_argument("--checkpoint", type=Path, required=True)
|
| 26 |
+
parser.add_argument("--vocab", type=Path, required=True)
|
| 27 |
+
parser.add_argument("--output", type=Path, required=True)
|
| 28 |
+
parser.add_argument("--seed", type=int, default=1)
|
| 29 |
+
parser.add_argument("--batch-size", type=int, default=64)
|
| 30 |
+
parser.add_argument("--device", default="cuda")
|
| 31 |
+
parser.add_argument("--checkpoint-only", action="store_true",
|
| 32 |
+
help="Load and validate the frozen model without a dataset")
|
| 33 |
+
return parser.parse_args()
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
def main() -> None:
|
| 37 |
+
args = parse_args()
|
| 38 |
+
archive = Path(__file__).resolve().parents[1]
|
| 39 |
+
sys.path[:0] = [
|
| 40 |
+
str(archive / "model"),
|
| 41 |
+
str(archive / "test"),
|
| 42 |
+
str(args.repo / "src" / "GEARS"),
|
| 43 |
+
str(args.repo / "src" / "scGPT"),
|
| 44 |
+
str(args.repo / "third_party" / "GEARS"),
|
| 45 |
+
str(args.repo / "third_party" / "scGPT"),
|
| 46 |
+
]
|
| 47 |
+
|
| 48 |
+
from cell_jepa.objectives_anything import PerturbationJEPAAnything
|
| 49 |
+
from Norman_Single_Cell_Perturbation_Prediction_test_utils import (
|
| 50 |
+
add_top20_non_dropout_metrics,
|
| 51 |
+
build_model,
|
| 52 |
+
perturb_flags,
|
| 53 |
+
set_seed,
|
| 54 |
+
)
|
| 55 |
+
from scgpt.tokenizer import GeneVocab
|
| 56 |
+
|
| 57 |
+
if args.device == "cuda" and not torch.cuda.is_available():
|
| 58 |
+
raise RuntimeError("--device cuda was requested but CUDA is unavailable")
|
| 59 |
+
device = torch.device(args.device)
|
| 60 |
+
torch.set_num_threads(8)
|
| 61 |
+
torch.set_num_interop_threads(1)
|
| 62 |
+
set_seed(args.seed)
|
| 63 |
+
|
| 64 |
+
payload = torch.load(args.checkpoint, map_location="cpu", weights_only=True)
|
| 65 |
+
config = payload.get("args", {})
|
| 66 |
+
num_factors = int(config.get("num_factors", 4))
|
| 67 |
+
vocab = GeneVocab.from_file(args.vocab)
|
| 68 |
+
for token in ("<pad>", "<cls>", "<eoc>"):
|
| 69 |
+
if token not in vocab:
|
| 70 |
+
vocab.append_token(token)
|
| 71 |
+
student = build_model(vocab, use_fast=False)
|
| 72 |
+
model = PerturbationJEPAAnything(student, num_factors=num_factors).to(device)
|
| 73 |
+
model.load_state_dict(payload["model"], strict=True)
|
| 74 |
+
model.eval()
|
| 75 |
+
if args.checkpoint_only:
|
| 76 |
+
output = {"status": "checkpoint_loaded", "checkpoint": str(args.checkpoint),
|
| 77 |
+
"parameters": sum(parameter.numel() for parameter in model.parameters()),
|
| 78 |
+
"num_factors": num_factors}
|
| 79 |
+
args.output.parent.mkdir(parents=True, exist_ok=True)
|
| 80 |
+
args.output.write_text(json.dumps(output, indent=2) + "\n")
|
| 81 |
+
print(json.dumps(output, indent=2))
|
| 82 |
+
return
|
| 83 |
+
if args.data_root is None:
|
| 84 |
+
raise ValueError("--data-root is required unless --checkpoint-only is used")
|
| 85 |
+
from gears import PertData
|
| 86 |
+
from scgpt.utils import compute_perturbation_metrics, map_raw_id_to_vocab_id
|
| 87 |
+
|
| 88 |
+
pert_data = PertData(str(args.data_root))
|
| 89 |
+
pert_data.default_pert_graph = False
|
| 90 |
+
pert_data.load(data_path=str(args.data_root / "norman"))
|
| 91 |
+
pert_data.ctrl_adata = pert_data.adata[pert_data.adata.obs["condition"] == "ctrl"]
|
| 92 |
+
pert_data.prepare_split(split="simulation", seed=args.seed)
|
| 93 |
+
pert_data.get_dataloader(batch_size=args.batch_size, test_batch_size=4)
|
| 94 |
+
|
| 95 |
+
gene_ids = np.asarray(
|
| 96 |
+
[vocab[gene] if gene in vocab else vocab["<pad>"]
|
| 97 |
+
for gene in pert_data.adata.var["gene_name"].tolist()],
|
| 98 |
+
dtype=int,
|
| 99 |
+
)
|
| 100 |
+
|
| 101 |
+
results: dict[str, list] = {
|
| 102 |
+
"pert_cat": [], "pred": [], "truth": [], "pred_de": [], "truth_de": [],
|
| 103 |
+
}
|
| 104 |
+
with torch.no_grad():
|
| 105 |
+
for batch in pert_data.dataloader["test_loader"]:
|
| 106 |
+
batch = batch.to(device)
|
| 107 |
+
batch_size, n_genes = len(batch.y), batch.y.shape[1]
|
| 108 |
+
control = batch.x[:, 0].view(batch_size, n_genes)
|
| 109 |
+
flags = perturb_flags(batch, n_genes, device)
|
| 110 |
+
raw_ids = torch.arange(n_genes, device=device)
|
| 111 |
+
genes = map_raw_id_to_vocab_id(raw_ids, gene_ids).repeat(batch_size, 1)
|
| 112 |
+
padding = torch.zeros_like(control, dtype=torch.bool)
|
| 113 |
+
prediction = model.predict_expression(genes, control, flags, padding).float()
|
| 114 |
+
results["pert_cat"].extend(batch.pert)
|
| 115 |
+
results["pred"].append(prediction.cpu())
|
| 116 |
+
results["truth"].append(batch.y.cpu())
|
| 117 |
+
for index, de_idx in enumerate(batch.de_idx):
|
| 118 |
+
results["pred_de"].append(prediction[index, de_idx].cpu())
|
| 119 |
+
results["truth_de"].append(batch.y[index, de_idx].cpu())
|
| 120 |
+
|
| 121 |
+
results = {
|
| 122 |
+
key: (np.asarray(value) if key == "pert_cat" else torch.stack(value).numpy()
|
| 123 |
+
if key in {"pred_de", "truth_de"} else torch.cat(value).numpy())
|
| 124 |
+
for key, value in results.items()
|
| 125 |
+
}
|
| 126 |
+
metrics = compute_perturbation_metrics(results, pert_data.ctrl_adata)
|
| 127 |
+
metrics = add_top20_non_dropout_metrics(metrics, pert_data.adata, results)
|
| 128 |
+
output = {
|
| 129 |
+
"status": "completed",
|
| 130 |
+
"checkpoint": str(args.checkpoint),
|
| 131 |
+
"dataset": "norman",
|
| 132 |
+
"seed": args.seed,
|
| 133 |
+
"test_cells": int(len(results["pred"])),
|
| 134 |
+
"test_conditions": int(len(np.unique(results["pert_cat"]))),
|
| 135 |
+
"metrics": {key: float(value) for key, value in metrics.items()},
|
| 136 |
+
}
|
| 137 |
+
args.output.parent.mkdir(parents=True, exist_ok=True)
|
| 138 |
+
args.output.write_text(json.dumps(output, indent=2) + "\n")
|
| 139 |
+
print(json.dumps(output, indent=2))
|
| 140 |
+
|
| 141 |
+
|
| 142 |
+
if __name__ == "__main__":
|
| 143 |
+
main()
|
02_Single_Cell/Norman_Single_Cell_Perturbation_Prediction/test/Norman_Single_Cell_Perturbation_Prediction_test_utils.py
ADDED
|
@@ -0,0 +1,67 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Evaluation helpers for Norman single-cell perturbation prediction."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
import random
|
| 6 |
+
|
| 7 |
+
import numpy as np
|
| 8 |
+
import torch
|
| 9 |
+
from scgpt.model import TransformerGenerator
|
| 10 |
+
from scgpt.tokenizer import GeneVocab
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
def set_seed(seed: int) -> None:
|
| 14 |
+
random.seed(seed)
|
| 15 |
+
np.random.seed(seed)
|
| 16 |
+
torch.manual_seed(seed)
|
| 17 |
+
torch.cuda.manual_seed_all(seed)
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
def build_model(vocab: GeneVocab, use_fast: bool) -> TransformerGenerator:
|
| 21 |
+
return TransformerGenerator(
|
| 22 |
+
ntoken=len(vocab),
|
| 23 |
+
d_model=512,
|
| 24 |
+
nhead=8,
|
| 25 |
+
d_hid=512,
|
| 26 |
+
nlayers=12,
|
| 27 |
+
nlayers_cls=3,
|
| 28 |
+
n_cls=1,
|
| 29 |
+
vocab=vocab,
|
| 30 |
+
dropout=0.0,
|
| 31 |
+
pad_token="<pad>",
|
| 32 |
+
pad_value=0,
|
| 33 |
+
pert_pad_id=0,
|
| 34 |
+
use_fast_transformer=use_fast,
|
| 35 |
+
)
|
| 36 |
+
|
| 37 |
+
|
| 38 |
+
def perturb_flags(batch, n_genes: int, device: torch.device) -> torch.Tensor:
|
| 39 |
+
"""Support both historic and current GEARS batch layouts."""
|
| 40 |
+
batch_size = len(batch.y)
|
| 41 |
+
if batch.x.shape[1] > 1:
|
| 42 |
+
return batch.x[:, 1].long().view(batch_size, n_genes)
|
| 43 |
+
flags = torch.zeros((batch_size, n_genes), dtype=torch.long, device=device)
|
| 44 |
+
for row, indices in enumerate(batch.pert_idx):
|
| 45 |
+
flags[row, torch.as_tensor(indices, dtype=torch.long, device=device)] = 1
|
| 46 |
+
return flags
|
| 47 |
+
|
| 48 |
+
|
| 49 |
+
def add_top20_non_dropout_metrics(metrics: dict, adata, results: dict) -> dict:
|
| 50 |
+
"""Add GEARS top-20 non-dropout Pearson metrics."""
|
| 51 |
+
from gears.inference import non_dropout_analysis
|
| 52 |
+
|
| 53 |
+
per_perturbation = non_dropout_analysis(adata, results)
|
| 54 |
+
for output_key, gears_key in (
|
| 55 |
+
("top20_de_non_dropout", "pearson_top20_de_non_dropout"),
|
| 56 |
+
("pearson_delta_top20_de_non_dropout", "pearson_delta_top20_de_non_dropout"),
|
| 57 |
+
):
|
| 58 |
+
values = [
|
| 59 |
+
value[gears_key]
|
| 60 |
+
for value in per_perturbation.values()
|
| 61 |
+
if gears_key in value
|
| 62 |
+
]
|
| 63 |
+
if values:
|
| 64 |
+
metrics[output_key] = float(np.mean(values))
|
| 65 |
+
if "top20_de_non_dropout" in metrics:
|
| 66 |
+
metrics["top20"] = metrics["top20_de_non_dropout"]
|
| 67 |
+
return metrics
|
02_Single_Cell/README.md
ADDED
|
@@ -0,0 +1,11 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Single-Cell Perturbation
|
| 2 |
+
|
| 3 |
+
Single-cell RNA sequencing describes a cell with a high-dimensional and sparse expression profile. JEPA pretraining predicts a complete cell representation from masked expression, encouraging the encoder to retain biological state while reducing sensitivity to gene-level dropout and count noise. Perturbation experiments use the difference between perturbed and control embeddings as the target response. OPF divides that response into complementary factors that are recombined to predict the post-perturbation expression profile.
|
| 4 |
+
|
| 5 |
+
Open a model page for its research background, checkpoint, input format, and inference command. Additional checkpoints from this domain will be released progressively.
|
| 6 |
+
|
| 7 |
+
## Models
|
| 8 |
+
|
| 9 |
+
- [Norman_Single_Cell_Perturbation_Prediction](Norman_Single_Cell_Perturbation_Prediction/README.md)
|
| 10 |
+
|
| 11 |
+
[Back to the project overview](../README.md)
|
04_CITRIS_Intervention/CITRIS_Interventional_Pong/README.md
ADDED
|
@@ -0,0 +1,36 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# CITRIS Interventional Pong — Legacy Model
|
| 2 |
+
|
| 3 |
+
## Background
|
| 4 |
+
|
| 5 |
+
Interventional Pong exposes interventions on ball motion, ball position, and paddle position. The task studies whether single-intervention dynamics can be recombined when several variables change together. This checkpoint belongs to the earlier three-channel protocol.
|
| 6 |
+
|
| 7 |
+
## Files
|
| 8 |
+
|
| 9 |
+
- Model implementation: [`../../model`](<../../model>)
|
| 10 |
+
- Checkpoint: [`checkpoints/CITRIS_Interventional_Pong_final_checkpoint.pt`](checkpoints/CITRIS_Interventional_Pong_final_checkpoint.pt)
|
| 11 |
+
- Inference command: [`../../test/infer_dynamics.py`](<../../test/infer_dynamics.py>)
|
| 12 |
+
|
| 13 |
+
## Run inference
|
| 14 |
+
|
| 15 |
+
Run from this experiment directory:
|
| 16 |
+
|
| 17 |
+
```bash
|
| 18 |
+
python ../../test/infer_dynamics.py \
|
| 19 |
+
--checkpoint checkpoints/CITRIS_Interventional_Pong_final_checkpoint.pt \
|
| 20 |
+
--input <preprocessed_transitions.npz> \
|
| 21 |
+
--output <new_output_dir>/predictions.npz
|
| 22 |
+
```
|
| 23 |
+
|
| 24 |
+
Provide `x` and `action` for one-step prediction, or `x` and `actions` for a recurrent rollout. Optional `y` or `targets` arrays add overall and per-step MSE to the output report. Use the dataset-specific flattening, channel order, normalization, and action scaling described for the experiment.
|
| 25 |
+
|
| 26 |
+
For a quick checkpoint and tensor-shape check, omit `--input`:
|
| 27 |
+
|
| 28 |
+
```bash
|
| 29 |
+
python ../../test/infer_dynamics.py \
|
| 30 |
+
--checkpoint checkpoints/CITRIS_Interventional_Pong_final_checkpoint.pt \
|
| 31 |
+
--output /tmp/jepa_inference_smoke/predictions.npz
|
| 32 |
+
```
|
| 33 |
+
|
| 34 |
+
This command reports the predicted tensor shape and finite-output status. Dataset runs also report MSE when targets are supplied.
|
| 35 |
+
|
| 36 |
+
[Back to the project overview](<../../README.md>)
|
04_CITRIS_Intervention/CITRIS_Interventional_Pong/checkpoints/CITRIS_Interventional_Pong_final_checkpoint.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:80c890f601d4af364850737298307312f42203cef9fd2c0baaee5cb782feeb54
|
| 3 |
+
size 14709115
|
04_CITRIS_Intervention/README.md
ADDED
|
@@ -0,0 +1,12 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Interventional Dynamics
|
| 2 |
+
|
| 3 |
+
CITRIS Interventional Pong separates ordinary temporal evolution from explicit interventions on underlying game variables such as ball motion, ball position, and paddle position. The model receives the pre-intervention observation together with a label describing which variables changed, then predicts the resulting next state. This makes the task a direct test of compositional dynamics: factors learned from individual interventions should remain useful when several interventions occur together.
|
| 4 |
+
|
| 5 |
+
Open a model page for its research background, checkpoint, input format, and inference command. Additional checkpoints from this domain will be released progressively.
|
| 6 |
+
|
| 7 |
+
## Models
|
| 8 |
+
|
| 9 |
+
- [CITRIS_Interventional_Pong](CITRIS_Interventional_Pong/README.md)
|
| 10 |
+
- [six_step_one_step](six_step_one_step/README.md)
|
| 11 |
+
|
| 12 |
+
[Back to the project overview](../README.md)
|
04_CITRIS_Intervention/six_step_one_step/README.md
ADDED
|
@@ -0,0 +1,35 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# CITRIS Six-Step Models: Strict One-Step Evaluation
|
| 2 |
+
|
| 3 |
+
## Background
|
| 4 |
+
|
| 5 |
+
These models were trained with recurrent six-step feedback and are evaluated here on one-step prediction. The experiment asks whether a model trained for short rollout can predict held-out single interventions and simultaneous intervention combinations.
|
| 6 |
+
|
| 7 |
+
The protocol uses all four original channels: RGB plus the movement-projection channel. It is separate from the legacy three-channel checkpoint.
|
| 8 |
+
|
| 9 |
+
## Files
|
| 10 |
+
|
| 11 |
+
- Inference model: [`../../model`](../../model)
|
| 12 |
+
- Checkpoints: [`checkpoints`](checkpoints), with Dense and OPF weights for seeds 11, 23, 37, 53, and 71
|
| 13 |
+
- Test entry point: [`test/evaluate.py`](test/evaluate.py)
|
| 14 |
+
|
| 15 |
+
## Run the test
|
| 16 |
+
|
| 17 |
+
```bash
|
| 18 |
+
python test/evaluate.py \
|
| 19 |
+
--data-root <citris_data_directory> \
|
| 20 |
+
--output <new_output_directory> \
|
| 21 |
+
--device cuda
|
| 22 |
+
```
|
| 23 |
+
|
| 24 |
+
Provide the original four-channel `single_test.npz` and `combo_test.npz` files. The script loads ten checkpoints, evaluates exactly-one-intervention transitions and transitions with at least two interventions, and writes `results.json` plus `per_seed.csv`.
|
| 25 |
+
|
| 26 |
+
Five-seed four-channel MSE from the verified checkpoint run:
|
| 27 |
+
|
| 28 |
+
| One-step test split | Dense | OPF |
|
| 29 |
+
|---|---:|---:|
|
| 30 |
+
| Exactly one intervention | 0.00952107 | 0.00682598 |
|
| 31 |
+
| At least two interventions | 0.00942106 | 0.00822350 |
|
| 32 |
+
|
| 33 |
+
The table reports one-step evaluation for models trained with recurrent six-step feedback.
|
| 34 |
+
|
| 35 |
+
[Back to the project overview](../../README.md)
|
04_CITRIS_Intervention/six_step_one_step/checkpoints/seed11/h6_dense.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:3f688358c41f5f3b2ed3e78c844c0e3883b421d083e6ff40d809a312c28dcad7
|
| 3 |
+
size 18642739
|
04_CITRIS_Intervention/six_step_one_step/checkpoints/seed11/h6_opf.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:bbbe28b515c5152b679c4e763bddb47bac58639324d366e5c9257dfee2a4e1bd
|
| 3 |
+
size 18640945
|
04_CITRIS_Intervention/six_step_one_step/checkpoints/seed23/h6_dense.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:1bf5b89faacf28cf99ab82b85f8f14a31276eeeb01d173aaee69e42fde9085d5
|
| 3 |
+
size 18642739
|
04_CITRIS_Intervention/six_step_one_step/checkpoints/seed23/h6_opf.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:3d8138d3c13fd448e7a2986cd1fcce18fceeb3c31362f3e7a897e0771217f105
|
| 3 |
+
size 18640945
|
04_CITRIS_Intervention/six_step_one_step/checkpoints/seed37/h6_dense.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:fa880b7b2e6610670da9b7b49f98dd24b2100dd98091d19d78bac5add140a688
|
| 3 |
+
size 18642739
|
04_CITRIS_Intervention/six_step_one_step/checkpoints/seed37/h6_opf.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:1170a0b22b2048b761cb07b56175d2508a8c0c27fc57277ac040e9fac1ea2f88
|
| 3 |
+
size 18640945
|
04_CITRIS_Intervention/six_step_one_step/checkpoints/seed53/h6_dense.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:fe47c9b5826b5d6547d324b736dcff46894152b45c4a262b229b5e37077882d1
|
| 3 |
+
size 18642739
|
04_CITRIS_Intervention/six_step_one_step/checkpoints/seed53/h6_opf.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:862a2d7db70b4ed8992ae681c26815216d0d687998198eff13f01b4308a15c1f
|
| 3 |
+
size 18640945
|
04_CITRIS_Intervention/six_step_one_step/checkpoints/seed71/h6_dense.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:6e0953e2d690607efd5c04c58903d42be2e4643448e50c25de6c1332505d2f7a
|
| 3 |
+
size 18642739
|
04_CITRIS_Intervention/six_step_one_step/checkpoints/seed71/h6_opf.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:86e1904460c57e6360951d62277431ab3273500e3493225cc91dca3a173793f6
|
| 3 |
+
size 18640945
|
04_CITRIS_Intervention/six_step_one_step/test/evaluate.py
ADDED
|
@@ -0,0 +1,60 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python3
|
| 2 |
+
"""Evaluate the frozen six-step CITRIS checkpoints on strict one-step splits. No training."""
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
import argparse,csv,json,sys
|
| 5 |
+
from pathlib import Path
|
| 6 |
+
import numpy as np
|
| 7 |
+
import torch
|
| 8 |
+
|
| 9 |
+
PROJECT=Path(__file__).resolve().parents[3]
|
| 10 |
+
sys.path.insert(0,str(PROJECT))
|
| 11 |
+
from model import load_dynamics_checkpoint
|
| 12 |
+
|
| 13 |
+
GROUPS=['ball_vel_dir','ball_x','ball_y','paddle_left_y','paddle_right_y']
|
| 14 |
+
SEEDS=[11,23,37,53,71]
|
| 15 |
+
|
| 16 |
+
def load_split(path):
|
| 17 |
+
with np.load(path,allow_pickle=False) as data:
|
| 18 |
+
ids=[data['keys'].tolist().index(k) for k in GROUPS]
|
| 19 |
+
x=torch.from_numpy(data['images'].copy()).permute(0,3,1,2).float().div_(255).flatten(1)
|
| 20 |
+
actions=torch.from_numpy(data['targets'][:,ids].astype('float32'))
|
| 21 |
+
if x.shape[1]!=4096 or len(actions)!=len(x)-1:raise ValueError('Expected original four-channel CITRIS arrays')
|
| 22 |
+
return x,actions
|
| 23 |
+
|
| 24 |
+
def mse(model,x,actions,indices,device,batch=256):
|
| 25 |
+
total=0.0
|
| 26 |
+
with torch.inference_mode():
|
| 27 |
+
for start in range(0,len(indices),batch):
|
| 28 |
+
ix=indices[start:start+batch]
|
| 29 |
+
prediction=model(x[ix].to(device),actions[ix].to(device))
|
| 30 |
+
total+=float((prediction-x[ix+1].to(device)).square().sum())
|
| 31 |
+
return total/(len(indices)*x.shape[1])
|
| 32 |
+
|
| 33 |
+
def main():
|
| 34 |
+
p=argparse.ArgumentParser(description=__doc__);p.add_argument('--data-root',type=Path,required=True);p.add_argument('--checkpoint-root',type=Path,default=Path(__file__).resolve().parents[1]/'checkpoints');p.add_argument('--output',type=Path,required=True);p.add_argument('--device',default='cpu');a=p.parse_args()
|
| 35 |
+
rows=[]
|
| 36 |
+
splits=[('single_test',a.data_root/'single_test.npz',lambda active:active==1),('combo_test',a.data_root/'combo_test.npz',lambda active:active>=2)]
|
| 37 |
+
loaded_splits={}
|
| 38 |
+
for name,path,rule in splits:
|
| 39 |
+
x,actions=load_split(path);loaded_splits[name]=(x,actions,torch.where(rule(actions.sum(1)))[0])
|
| 40 |
+
for seed in SEEDS:
|
| 41 |
+
for kind in ['dense','opf']:
|
| 42 |
+
path=a.checkpoint_root/f'seed{seed}'/f'h6_{kind}.pt'
|
| 43 |
+
loaded=load_dynamics_checkpoint(path,'model',a.device)
|
| 44 |
+
if loaded.metadata.get('horizon')!=6 or loaded.metadata.get('kind')!=kind:raise ValueError(f'Unexpected metadata in {path}')
|
| 45 |
+
for split,(x,actions,indices) in loaded_splits.items():
|
| 46 |
+
rows.append({'seed':seed,'model':kind,'split':split,'count':len(indices),'mse_all4':mse(loaded.model,x,actions,indices,a.device)})
|
| 47 |
+
summary={}
|
| 48 |
+
for split,_,_ in splits:
|
| 49 |
+
summary[split]={}
|
| 50 |
+
for kind in ['dense','opf']:
|
| 51 |
+
values=np.array([row['mse_all4'] for row in rows if row['split']==split and row['model']==kind])
|
| 52 |
+
summary[split][kind]={'mean':float(values.mean()),'sample_sd':float(values.std(ddof=1)),'values':values.tolist()}
|
| 53 |
+
summary[split]['opf_reduction_percent']=float(100*(1-summary[split]['opf']['mean']/summary[split]['dense']['mean']))
|
| 54 |
+
a.output.mkdir(parents=True,exist_ok=False)
|
| 55 |
+
with (a.output/'per_seed.csv').open('w',newline='') as f:
|
| 56 |
+
writer=csv.DictWriter(f,fieldnames=list(rows[0]));writer.writeheader();writer.writerows(rows)
|
| 57 |
+
(a.output/'results.json').write_text(json.dumps({'status':'complete','training_horizon':6,'evaluation_horizon':1,'summary':summary},indent=2))
|
| 58 |
+
print(json.dumps(summary,indent=2))
|
| 59 |
+
|
| 60 |
+
if __name__=='__main__':main()
|
05_Ten_Task_Dynamics/CLEVRER_Visual_Dynamics/README.md
ADDED
|
@@ -0,0 +1,36 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# CLEVRER Visual Dynamics
|
| 2 |
+
|
| 3 |
+
## Background
|
| 4 |
+
|
| 5 |
+
CLEVRER represents dynamics as controlled video transitions. The checkpoint predicts the next visual state and can be rolled forward to study error accumulation.
|
| 6 |
+
|
| 7 |
+
## Files
|
| 8 |
+
|
| 9 |
+
- Model implementation: [`../../model`](<../../model>)
|
| 10 |
+
- Checkpoint: [`checkpoints/CLEVRER_Visual_Dynamics_final_checkpoint.pt`](checkpoints/CLEVRER_Visual_Dynamics_final_checkpoint.pt)
|
| 11 |
+
- Inference command: [`../../test/infer_dynamics.py`](<../../test/infer_dynamics.py>)
|
| 12 |
+
|
| 13 |
+
## Run inference
|
| 14 |
+
|
| 15 |
+
Run from this experiment directory:
|
| 16 |
+
|
| 17 |
+
```bash
|
| 18 |
+
python ../../test/infer_dynamics.py \
|
| 19 |
+
--checkpoint checkpoints/CLEVRER_Visual_Dynamics_final_checkpoint.pt \
|
| 20 |
+
--input <preprocessed_transitions.npz> \
|
| 21 |
+
--output <new_output_dir>/predictions.npz
|
| 22 |
+
```
|
| 23 |
+
|
| 24 |
+
Provide `x` and `action` for one-step prediction, or `x` and `actions` for a recurrent rollout. Optional `y` or `targets` arrays add overall and per-step MSE to the output report. Use the dataset-specific flattening, channel order, normalization, and action scaling described for the experiment.
|
| 25 |
+
|
| 26 |
+
For a quick checkpoint and tensor-shape check, omit `--input`:
|
| 27 |
+
|
| 28 |
+
```bash
|
| 29 |
+
python ../../test/infer_dynamics.py \
|
| 30 |
+
--checkpoint checkpoints/CLEVRER_Visual_Dynamics_final_checkpoint.pt \
|
| 31 |
+
--output /tmp/jepa_inference_smoke/predictions.npz
|
| 32 |
+
```
|
| 33 |
+
|
| 34 |
+
This command reports the predicted tensor shape and finite-output status. Dataset runs also report MSE when targets are supplied.
|
| 35 |
+
|
| 36 |
+
[Back to the project overview](<../../README.md>)
|
05_Ten_Task_Dynamics/CLEVRER_Visual_Dynamics/checkpoints/CLEVRER_Visual_Dynamics_final_checkpoint.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:8251b743f84ba43c8ab94307a9f44b6242970f36df1ffbf27268cc031c5d9e9d
|
| 3 |
+
size 39657359
|
05_Ten_Task_Dynamics/CausalWorld_Pushing/README.md
ADDED
|
@@ -0,0 +1,36 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# CausalWorld Pushing
|
| 2 |
+
|
| 3 |
+
## Background
|
| 4 |
+
|
| 5 |
+
This task predicts how a robot action changes object state. It tests whether action-conditioned latent factors retain the physical motion signal across controlled conditions.
|
| 6 |
+
|
| 7 |
+
## Files
|
| 8 |
+
|
| 9 |
+
- Model implementation: [`../../model`](<../../model>)
|
| 10 |
+
- Checkpoint: [`checkpoints/CausalWorld_Pushing_final_checkpoint.pt`](checkpoints/CausalWorld_Pushing_final_checkpoint.pt)
|
| 11 |
+
- Inference command: [`../../test/infer_dynamics.py`](<../../test/infer_dynamics.py>)
|
| 12 |
+
|
| 13 |
+
## Run inference
|
| 14 |
+
|
| 15 |
+
Run from this experiment directory:
|
| 16 |
+
|
| 17 |
+
```bash
|
| 18 |
+
python ../../test/infer_dynamics.py \
|
| 19 |
+
--checkpoint checkpoints/CausalWorld_Pushing_final_checkpoint.pt \
|
| 20 |
+
--input <preprocessed_transitions.npz> \
|
| 21 |
+
--output <new_output_dir>/predictions.npz
|
| 22 |
+
```
|
| 23 |
+
|
| 24 |
+
Provide `x` and `action` for one-step prediction, or `x` and `actions` for a recurrent rollout. Optional `y` or `targets` arrays add overall and per-step MSE to the output report. Use the dataset-specific flattening, channel order, normalization, and action scaling described for the experiment.
|
| 25 |
+
|
| 26 |
+
For a quick checkpoint and tensor-shape check, omit `--input`:
|
| 27 |
+
|
| 28 |
+
```bash
|
| 29 |
+
python ../../test/infer_dynamics.py \
|
| 30 |
+
--checkpoint checkpoints/CausalWorld_Pushing_final_checkpoint.pt \
|
| 31 |
+
--output /tmp/jepa_inference_smoke/predictions.npz
|
| 32 |
+
```
|
| 33 |
+
|
| 34 |
+
This command reports the predicted tensor shape and finite-output status. Dataset runs also report MSE when targets are supplied.
|
| 35 |
+
|
| 36 |
+
[Back to the project overview](<../../README.md>)
|
05_Ten_Task_Dynamics/CausalWorld_Pushing/checkpoints/CausalWorld_Pushing_final_checkpoint.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:cd7b198d91465e6284cddafbfa89ecc1d9cf2ede1678eeb89714f485b61b4d35
|
| 3 |
+
size 2039887
|
05_Ten_Task_Dynamics/CausalWorld_Pushing_Long_Horizon/README.md
ADDED
|
@@ -0,0 +1,36 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# CausalWorld Long-Horizon Pushing
|
| 2 |
+
|
| 3 |
+
## Background
|
| 4 |
+
|
| 5 |
+
This variant repeatedly feeds model predictions back into the dynamics model. It is intended for studying long-horizon state prediction and model-based control.
|
| 6 |
+
|
| 7 |
+
## Files
|
| 8 |
+
|
| 9 |
+
- Model implementation: [`../../model`](<../../model>)
|
| 10 |
+
- Checkpoint: [`checkpoints/CausalWorld_Pushing_Long_Horizon_final_checkpoint.pt`](checkpoints/CausalWorld_Pushing_Long_Horizon_final_checkpoint.pt)
|
| 11 |
+
- Inference command: [`../../test/infer_dynamics.py`](<../../test/infer_dynamics.py>)
|
| 12 |
+
|
| 13 |
+
## Run inference
|
| 14 |
+
|
| 15 |
+
Run from this experiment directory:
|
| 16 |
+
|
| 17 |
+
```bash
|
| 18 |
+
python ../../test/infer_dynamics.py \
|
| 19 |
+
--checkpoint checkpoints/CausalWorld_Pushing_Long_Horizon_final_checkpoint.pt \
|
| 20 |
+
--input <preprocessed_transitions.npz> \
|
| 21 |
+
--output <new_output_dir>/predictions.npz
|
| 22 |
+
```
|
| 23 |
+
|
| 24 |
+
Provide `x` and `action` for one-step prediction, or `x` and `actions` for a recurrent rollout. Optional `y` or `targets` arrays add overall and per-step MSE to the output report. Use the dataset-specific flattening, channel order, normalization, and action scaling described for the experiment.
|
| 25 |
+
|
| 26 |
+
For a quick checkpoint and tensor-shape check, omit `--input`:
|
| 27 |
+
|
| 28 |
+
```bash
|
| 29 |
+
python ../../test/infer_dynamics.py \
|
| 30 |
+
--checkpoint checkpoints/CausalWorld_Pushing_Long_Horizon_final_checkpoint.pt \
|
| 31 |
+
--output /tmp/jepa_inference_smoke/predictions.npz
|
| 32 |
+
```
|
| 33 |
+
|
| 34 |
+
This command reports the predicted tensor shape and finite-output status. Dataset runs also report MSE when targets are supplied.
|
| 35 |
+
|
| 36 |
+
[Back to the project overview](<../../README.md>)
|
05_Ten_Task_Dynamics/CausalWorld_Pushing_Long_Horizon/checkpoints/CausalWorld_Pushing_Long_Horizon_final_checkpoint.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:cd7f2a595e2feae5b5827ee7ee6ff4c07c7188fe0bffaf9a398493cf67cd94c3
|
| 3 |
+
size 2039951
|
05_Ten_Task_Dynamics/DMC_Cartpole_Pixels/README.md
ADDED
|
@@ -0,0 +1,36 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# DMC Cartpole from Pixels
|
| 2 |
+
|
| 3 |
+
## Background
|
| 4 |
+
|
| 5 |
+
The cartpole state is observed as pixels rather than low-dimensional coordinates. The model predicts how an action changes the next visual state.
|
| 6 |
+
|
| 7 |
+
## Files
|
| 8 |
+
|
| 9 |
+
- Model implementation: [`../../model`](<../../model>)
|
| 10 |
+
- Checkpoint: [`checkpoints/DMC_Cartpole_Pixels_final_checkpoint.pt`](checkpoints/DMC_Cartpole_Pixels_final_checkpoint.pt)
|
| 11 |
+
- Inference command: [`../../test/infer_dynamics.py`](<../../test/infer_dynamics.py>)
|
| 12 |
+
|
| 13 |
+
## Run inference
|
| 14 |
+
|
| 15 |
+
Run from this experiment directory:
|
| 16 |
+
|
| 17 |
+
```bash
|
| 18 |
+
python ../../test/infer_dynamics.py \
|
| 19 |
+
--checkpoint checkpoints/DMC_Cartpole_Pixels_final_checkpoint.pt \
|
| 20 |
+
--input <preprocessed_transitions.npz> \
|
| 21 |
+
--output <new_output_dir>/predictions.npz
|
| 22 |
+
```
|
| 23 |
+
|
| 24 |
+
Provide `x` and `action` for one-step prediction, or `x` and `actions` for a recurrent rollout. Optional `y` or `targets` arrays add overall and per-step MSE to the output report. Use the dataset-specific flattening, channel order, normalization, and action scaling described for the experiment.
|
| 25 |
+
|
| 26 |
+
For a quick checkpoint and tensor-shape check, omit `--input`:
|
| 27 |
+
|
| 28 |
+
```bash
|
| 29 |
+
python ../../test/infer_dynamics.py \
|
| 30 |
+
--checkpoint checkpoints/DMC_Cartpole_Pixels_final_checkpoint.pt \
|
| 31 |
+
--output /tmp/jepa_inference_smoke/predictions.npz
|
| 32 |
+
```
|
| 33 |
+
|
| 34 |
+
This command reports the predicted tensor shape and finite-output status. Dataset runs also report MSE when targets are supplied.
|
| 35 |
+
|
| 36 |
+
[Back to the project overview](<../../README.md>)
|
05_Ten_Task_Dynamics/DMC_Cartpole_Pixels/checkpoints/DMC_Cartpole_Pixels_final_checkpoint.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:0d0ec41d94f43125aadf995efea8f73a9ec0173a7e0b1afb86e51ed0cda49d12
|
| 3 |
+
size 39657359
|
05_Ten_Task_Dynamics/DMC_VB_Cheetah_Run/README.md
ADDED
|
@@ -0,0 +1,36 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# DMC-VB Cheetah Run
|
| 2 |
+
|
| 3 |
+
## Background
|
| 4 |
+
|
| 5 |
+
This task combines locomotion dynamics with changing visual backgrounds. It studies whether the predictor follows body motion instead of appearance alone.
|
| 6 |
+
|
| 7 |
+
## Files
|
| 8 |
+
|
| 9 |
+
- Model implementation: [`../../model`](<../../model>)
|
| 10 |
+
- Checkpoint: [`checkpoints/DMC_VB_Cheetah_Run_final_checkpoint.pt`](checkpoints/DMC_VB_Cheetah_Run_final_checkpoint.pt)
|
| 11 |
+
- Inference command: [`../../test/infer_dynamics.py`](<../../test/infer_dynamics.py>)
|
| 12 |
+
|
| 13 |
+
## Run inference
|
| 14 |
+
|
| 15 |
+
Run from this experiment directory:
|
| 16 |
+
|
| 17 |
+
```bash
|
| 18 |
+
python ../../test/infer_dynamics.py \
|
| 19 |
+
--checkpoint checkpoints/DMC_VB_Cheetah_Run_final_checkpoint.pt \
|
| 20 |
+
--input <preprocessed_transitions.npz> \
|
| 21 |
+
--output <new_output_dir>/predictions.npz
|
| 22 |
+
```
|
| 23 |
+
|
| 24 |
+
Provide `x` and `action` for one-step prediction, or `x` and `actions` for a recurrent rollout. Optional `y` or `targets` arrays add overall and per-step MSE to the output report. Use the dataset-specific flattening, channel order, normalization, and action scaling described for the experiment.
|
| 25 |
+
|
| 26 |
+
For a quick checkpoint and tensor-shape check, omit `--input`:
|
| 27 |
+
|
| 28 |
+
```bash
|
| 29 |
+
python ../../test/infer_dynamics.py \
|
| 30 |
+
--checkpoint checkpoints/DMC_VB_Cheetah_Run_final_checkpoint.pt \
|
| 31 |
+
--output /tmp/jepa_inference_smoke/predictions.npz
|
| 32 |
+
```
|
| 33 |
+
|
| 34 |
+
This command reports the predicted tensor shape and finite-output status. Dataset runs also report MSE when targets are supplied.
|
| 35 |
+
|
| 36 |
+
[Back to the project overview](<../../README.md>)
|
05_Ten_Task_Dynamics/DMC_VB_Cheetah_Run/checkpoints/DMC_VB_Cheetah_Run_final_checkpoint.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:62fc08813c234b473bb5ab4b535e0fa6e81fe6b7f9302540d539c85ba7d1baa1
|
| 3 |
+
size 39662479
|
05_Ten_Task_Dynamics/Distracting_Control_Cartpole/README.md
ADDED
|
@@ -0,0 +1,36 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Distracting Control Cartpole
|
| 2 |
+
|
| 3 |
+
## Background
|
| 4 |
+
|
| 5 |
+
The physical cartpole evolves while backgrounds and visual distractors vary. The model is used to separate predictable motion from nuisance appearance.
|
| 6 |
+
|
| 7 |
+
## Files
|
| 8 |
+
|
| 9 |
+
- Model implementation: [`../../model`](<../../model>)
|
| 10 |
+
- Checkpoint: [`checkpoints/Distracting_Control_Cartpole_final_checkpoint.pt`](checkpoints/Distracting_Control_Cartpole_final_checkpoint.pt)
|
| 11 |
+
- Inference command: [`../../test/infer_dynamics.py`](<../../test/infer_dynamics.py>)
|
| 12 |
+
|
| 13 |
+
## Run inference
|
| 14 |
+
|
| 15 |
+
Run from this experiment directory:
|
| 16 |
+
|
| 17 |
+
```bash
|
| 18 |
+
python ../../test/infer_dynamics.py \
|
| 19 |
+
--checkpoint checkpoints/Distracting_Control_Cartpole_final_checkpoint.pt \
|
| 20 |
+
--input <preprocessed_transitions.npz> \
|
| 21 |
+
--output <new_output_dir>/predictions.npz
|
| 22 |
+
```
|
| 23 |
+
|
| 24 |
+
Provide `x` and `action` for one-step prediction, or `x` and `actions` for a recurrent rollout. Optional `y` or `targets` arrays add overall and per-step MSE to the output report. Use the dataset-specific flattening, channel order, normalization, and action scaling described for the experiment.
|
| 25 |
+
|
| 26 |
+
For a quick checkpoint and tensor-shape check, omit `--input`:
|
| 27 |
+
|
| 28 |
+
```bash
|
| 29 |
+
python ../../test/infer_dynamics.py \
|
| 30 |
+
--checkpoint checkpoints/Distracting_Control_Cartpole_final_checkpoint.pt \
|
| 31 |
+
--output /tmp/jepa_inference_smoke/predictions.npz
|
| 32 |
+
```
|
| 33 |
+
|
| 34 |
+
This command reports the predicted tensor shape and finite-output status. Dataset runs also report MSE when targets are supplied.
|
| 35 |
+
|
| 36 |
+
[Back to the project overview](<../../README.md>)
|
05_Ten_Task_Dynamics/Distracting_Control_Cartpole/checkpoints/Distracting_Control_Cartpole_final_checkpoint.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:d19c4126fb73ac9bc43e7ee43c6990a33bbc4f927c9f29a2e362998e3e7da4c8
|
| 3 |
+
size 39657423
|
05_Ten_Task_Dynamics/PDEBench_Burgers/README.md
ADDED
|
@@ -0,0 +1,36 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# PDEBench Burgers
|
| 2 |
+
|
| 3 |
+
## Background
|
| 4 |
+
|
| 5 |
+
Burgers dynamics combine nonlinear transport and dissipation. The checkpoint predicts the next discretized field and supports recurrent rollout.
|
| 6 |
+
|
| 7 |
+
## Files
|
| 8 |
+
|
| 9 |
+
- Model implementation: [`../../model`](<../../model>)
|
| 10 |
+
- Checkpoint: [`checkpoints/PDEBench_Burgers_final_checkpoint.pt`](checkpoints/PDEBench_Burgers_final_checkpoint.pt)
|
| 11 |
+
- Inference command: [`../../test/infer_dynamics.py`](<../../test/infer_dynamics.py>)
|
| 12 |
+
|
| 13 |
+
## Run inference
|
| 14 |
+
|
| 15 |
+
Run from this experiment directory:
|
| 16 |
+
|
| 17 |
+
```bash
|
| 18 |
+
python ../../test/infer_dynamics.py \
|
| 19 |
+
--checkpoint checkpoints/PDEBench_Burgers_final_checkpoint.pt \
|
| 20 |
+
--input <preprocessed_transitions.npz> \
|
| 21 |
+
--output <new_output_dir>/predictions.npz
|
| 22 |
+
```
|
| 23 |
+
|
| 24 |
+
Provide `x` and `action` for one-step prediction, or `x` and `actions` for a recurrent rollout. Optional `y` or `targets` arrays add overall and per-step MSE to the output report. Use the dataset-specific flattening, channel order, normalization, and action scaling described for the experiment.
|
| 25 |
+
|
| 26 |
+
For a quick checkpoint and tensor-shape check, omit `--input`:
|
| 27 |
+
|
| 28 |
+
```bash
|
| 29 |
+
python ../../test/infer_dynamics.py \
|
| 30 |
+
--checkpoint checkpoints/PDEBench_Burgers_final_checkpoint.pt \
|
| 31 |
+
--output /tmp/jepa_inference_smoke/predictions.npz
|
| 32 |
+
```
|
| 33 |
+
|
| 34 |
+
This command reports the predicted tensor shape and finite-output status. Dataset runs also report MSE when targets are supplied.
|
| 35 |
+
|
| 36 |
+
[Back to the project overview](<../../README.md>)
|
05_Ten_Task_Dynamics/PDEBench_Burgers/checkpoints/PDEBench_Burgers_final_checkpoint.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:83bbc2cacc9ff7769521a76c3c973c4b3895dd92362a49f2700cd80f8b334826
|
| 3 |
+
size 5009359
|
05_Ten_Task_Dynamics/PDEBench_Navier_Stokes/README.md
ADDED
|
@@ -0,0 +1,36 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# PDEBench Navier–Stokes
|
| 2 |
+
|
| 3 |
+
## Background
|
| 4 |
+
|
| 5 |
+
This task predicts the evolution of a two-dimensional fluid field. It tests whether the latent state retains spatial structure during repeated prediction.
|
| 6 |
+
|
| 7 |
+
## Files
|
| 8 |
+
|
| 9 |
+
- Model implementation: [`../../model`](<../../model>)
|
| 10 |
+
- Checkpoint: [`checkpoints/PDEBench_Navier_Stokes_final_checkpoint.pt`](checkpoints/PDEBench_Navier_Stokes_final_checkpoint.pt)
|
| 11 |
+
- Inference command: [`../../test/infer_dynamics.py`](<../../test/infer_dynamics.py>)
|
| 12 |
+
|
| 13 |
+
## Run inference
|
| 14 |
+
|
| 15 |
+
Run from this experiment directory:
|
| 16 |
+
|
| 17 |
+
```bash
|
| 18 |
+
python ../../test/infer_dynamics.py \
|
| 19 |
+
--checkpoint checkpoints/PDEBench_Navier_Stokes_final_checkpoint.pt \
|
| 20 |
+
--input <preprocessed_transitions.npz> \
|
| 21 |
+
--output <new_output_dir>/predictions.npz
|
| 22 |
+
```
|
| 23 |
+
|
| 24 |
+
Provide `x` and `action` for one-step prediction, or `x` and `actions` for a recurrent rollout. Optional `y` or `targets` arrays add overall and per-step MSE to the output report. Use the dataset-specific flattening, channel order, normalization, and action scaling described for the experiment.
|
| 25 |
+
|
| 26 |
+
For a quick checkpoint and tensor-shape check, omit `--input`:
|
| 27 |
+
|
| 28 |
+
```bash
|
| 29 |
+
python ../../test/infer_dynamics.py \
|
| 30 |
+
--checkpoint checkpoints/PDEBench_Navier_Stokes_final_checkpoint.pt \
|
| 31 |
+
--output /tmp/jepa_inference_smoke/predictions.npz
|
| 32 |
+
```
|
| 33 |
+
|
| 34 |
+
This command reports the predicted tensor shape and finite-output status. Dataset runs also report MSE when targets are supplied.
|
| 35 |
+
|
| 36 |
+
[Back to the project overview](<../../README.md>)
|
05_Ten_Task_Dynamics/PDEBench_Navier_Stokes/checkpoints/PDEBench_Navier_Stokes_final_checkpoint.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:c63e159cef39d2611aec13a1c967d5f506cd3650ea537cffb062bf8f95ca0b18
|
| 3 |
+
size 8159247
|
05_Ten_Task_Dynamics/PDEBench_Reaction_Diffusion_1D/README.md
ADDED
|
@@ -0,0 +1,36 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# PDEBench Reaction–Diffusion 1D
|
| 2 |
+
|
| 3 |
+
## Background
|
| 4 |
+
|
| 5 |
+
Reaction and spatial diffusion operate at different scales. The model predicts the next one-dimensional field from the current field.
|
| 6 |
+
|
| 7 |
+
## Files
|
| 8 |
+
|
| 9 |
+
- Model implementation: [`../../model`](<../../model>)
|
| 10 |
+
- Checkpoint: [`checkpoints/PDEBench_Reaction_Diffusion_1D_final_checkpoint.pt`](checkpoints/PDEBench_Reaction_Diffusion_1D_final_checkpoint.pt)
|
| 11 |
+
- Inference command: [`../../test/infer_dynamics.py`](<../../test/infer_dynamics.py>)
|
| 12 |
+
|
| 13 |
+
## Run inference
|
| 14 |
+
|
| 15 |
+
Run from this experiment directory:
|
| 16 |
+
|
| 17 |
+
```bash
|
| 18 |
+
python ../../test/infer_dynamics.py \
|
| 19 |
+
--checkpoint checkpoints/PDEBench_Reaction_Diffusion_1D_final_checkpoint.pt \
|
| 20 |
+
--input <preprocessed_transitions.npz> \
|
| 21 |
+
--output <new_output_dir>/predictions.npz
|
| 22 |
+
```
|
| 23 |
+
|
| 24 |
+
Provide `x` and `action` for one-step prediction, or `x` and `actions` for a recurrent rollout. Optional `y` or `targets` arrays add overall and per-step MSE to the output report. Use the dataset-specific flattening, channel order, normalization, and action scaling described for the experiment.
|
| 25 |
+
|
| 26 |
+
For a quick checkpoint and tensor-shape check, omit `--input`:
|
| 27 |
+
|
| 28 |
+
```bash
|
| 29 |
+
python ../../test/infer_dynamics.py \
|
| 30 |
+
--checkpoint checkpoints/PDEBench_Reaction_Diffusion_1D_final_checkpoint.pt \
|
| 31 |
+
--output /tmp/jepa_inference_smoke/predictions.npz
|
| 32 |
+
```
|
| 33 |
+
|
| 34 |
+
This command reports the predicted tensor shape and finite-output status. Dataset runs also report MSE when targets are supplied.
|
| 35 |
+
|
| 36 |
+
[Back to the project overview](<../../README.md>)
|
05_Ten_Task_Dynamics/PDEBench_Reaction_Diffusion_1D/checkpoints/PDEBench_Reaction_Diffusion_1D_final_checkpoint.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:7db8dbc4ce2f77f525b34a2b69bb03754aa27d6358e3a8f1a1bd6a704db81535
|
| 3 |
+
size 5009359
|
05_Ten_Task_Dynamics/PDEBench_Reaction_Diffusion_2D/README.md
ADDED
|
@@ -0,0 +1,36 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# PDEBench Reaction–Diffusion 2D
|
| 2 |
+
|
| 3 |
+
## Background
|
| 4 |
+
|
| 5 |
+
This task extends reaction–diffusion forecasting to two-dimensional spatial fields and recurrent prediction.
|
| 6 |
+
|
| 7 |
+
## Files
|
| 8 |
+
|
| 9 |
+
- Model implementation: [`../../model`](<../../model>)
|
| 10 |
+
- Checkpoint: [`checkpoints/PDEBench_Reaction_Diffusion_2D_final_checkpoint.pt`](checkpoints/PDEBench_Reaction_Diffusion_2D_final_checkpoint.pt)
|
| 11 |
+
- Inference command: [`../../test/infer_dynamics.py`](<../../test/infer_dynamics.py>)
|
| 12 |
+
|
| 13 |
+
## Run inference
|
| 14 |
+
|
| 15 |
+
Run from this experiment directory:
|
| 16 |
+
|
| 17 |
+
```bash
|
| 18 |
+
python ../../test/infer_dynamics.py \
|
| 19 |
+
--checkpoint checkpoints/PDEBench_Reaction_Diffusion_2D_final_checkpoint.pt \
|
| 20 |
+
--input <preprocessed_transitions.npz> \
|
| 21 |
+
--output <new_output_dir>/predictions.npz
|
| 22 |
+
```
|
| 23 |
+
|
| 24 |
+
Provide `x` and `action` for one-step prediction, or `x` and `actions` for a recurrent rollout. Optional `y` or `targets` arrays add overall and per-step MSE to the output report. Use the dataset-specific flattening, channel order, normalization, and action scaling described for the experiment.
|
| 25 |
+
|
| 26 |
+
For a quick checkpoint and tensor-shape check, omit `--input`:
|
| 27 |
+
|
| 28 |
+
```bash
|
| 29 |
+
python ../../test/infer_dynamics.py \
|
| 30 |
+
--checkpoint checkpoints/PDEBench_Reaction_Diffusion_2D_final_checkpoint.pt \
|
| 31 |
+
--output /tmp/jepa_inference_smoke/predictions.npz
|
| 32 |
+
```
|
| 33 |
+
|
| 34 |
+
This command reports the predicted tensor shape and finite-output status. Dataset runs also report MSE when targets are supplied.
|
| 35 |
+
|
| 36 |
+
[Back to the project overview](<../../README.md>)
|
05_Ten_Task_Dynamics/PDEBench_Reaction_Diffusion_2D/checkpoints/PDEBench_Reaction_Diffusion_2D_final_checkpoint.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:bcaba73788e20dcae2f273a803765d1f6abec6e280d8064ab3c4a159e9d25f9e
|
| 3 |
+
size 8159183
|