dariuscty commited on
Commit
7904126
·
verified ·
1 Parent(s): 8c53c54

Initial upload: JEPA-Anything full project + domain checkpoints

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. .DS_Store +0 -0
  2. 01_Vision_Controlled_Binding/README.md +12 -0
  3. 01_Vision_Controlled_Binding/Shapes3D_JEPA_Pretraining/README.md +36 -0
  4. 01_Vision_Controlled_Binding/Shapes3D_JEPA_Pretraining/checkpoints/Shapes3D_JEPA_Pretraining_final_checkpoint.pt +3 -0
  5. 01_Vision_Controlled_Binding/Shapes3D_Main_Factor_Prediction/README.md +36 -0
  6. 01_Vision_Controlled_Binding/Shapes3D_Main_Factor_Prediction/checkpoints/Shapes3D_Main_Factor_Prediction_final_checkpoint.pt +3 -0
  7. 02_Single_Cell/Norman_Single_Cell_Perturbation_Prediction/README.md +44 -0
  8. 02_Single_Cell/Norman_Single_Cell_Perturbation_Prediction/checkpoints/Norman_Single_Cell_Perturbation_Prediction_final_checkpoint.pt +3 -0
  9. 02_Single_Cell/Norman_Single_Cell_Perturbation_Prediction/model/Norman_Single_Cell_Perturbation_Prediction_vocab.json +0 -0
  10. 02_Single_Cell/Norman_Single_Cell_Perturbation_Prediction/model/cell_jepa/__init__.py +12 -0
  11. 02_Single_Cell/Norman_Single_Cell_Perturbation_Prediction/model/cell_jepa/objectives.py +220 -0
  12. 02_Single_Cell/Norman_Single_Cell_Perturbation_Prediction/model/cell_jepa/objectives_anything.py +558 -0
  13. 02_Single_Cell/Norman_Single_Cell_Perturbation_Prediction/test/Norman_Single_Cell_Perturbation_Prediction_test.py +143 -0
  14. 02_Single_Cell/Norman_Single_Cell_Perturbation_Prediction/test/Norman_Single_Cell_Perturbation_Prediction_test_utils.py +67 -0
  15. 02_Single_Cell/README.md +11 -0
  16. 04_CITRIS_Intervention/CITRIS_Interventional_Pong/README.md +36 -0
  17. 04_CITRIS_Intervention/CITRIS_Interventional_Pong/checkpoints/CITRIS_Interventional_Pong_final_checkpoint.pt +3 -0
  18. 04_CITRIS_Intervention/README.md +12 -0
  19. 04_CITRIS_Intervention/six_step_one_step/README.md +35 -0
  20. 04_CITRIS_Intervention/six_step_one_step/checkpoints/seed11/h6_dense.pt +3 -0
  21. 04_CITRIS_Intervention/six_step_one_step/checkpoints/seed11/h6_opf.pt +3 -0
  22. 04_CITRIS_Intervention/six_step_one_step/checkpoints/seed23/h6_dense.pt +3 -0
  23. 04_CITRIS_Intervention/six_step_one_step/checkpoints/seed23/h6_opf.pt +3 -0
  24. 04_CITRIS_Intervention/six_step_one_step/checkpoints/seed37/h6_dense.pt +3 -0
  25. 04_CITRIS_Intervention/six_step_one_step/checkpoints/seed37/h6_opf.pt +3 -0
  26. 04_CITRIS_Intervention/six_step_one_step/checkpoints/seed53/h6_dense.pt +3 -0
  27. 04_CITRIS_Intervention/six_step_one_step/checkpoints/seed53/h6_opf.pt +3 -0
  28. 04_CITRIS_Intervention/six_step_one_step/checkpoints/seed71/h6_dense.pt +3 -0
  29. 04_CITRIS_Intervention/six_step_one_step/checkpoints/seed71/h6_opf.pt +3 -0
  30. 04_CITRIS_Intervention/six_step_one_step/test/evaluate.py +60 -0
  31. 05_Ten_Task_Dynamics/CLEVRER_Visual_Dynamics/README.md +36 -0
  32. 05_Ten_Task_Dynamics/CLEVRER_Visual_Dynamics/checkpoints/CLEVRER_Visual_Dynamics_final_checkpoint.pt +3 -0
  33. 05_Ten_Task_Dynamics/CausalWorld_Pushing/README.md +36 -0
  34. 05_Ten_Task_Dynamics/CausalWorld_Pushing/checkpoints/CausalWorld_Pushing_final_checkpoint.pt +3 -0
  35. 05_Ten_Task_Dynamics/CausalWorld_Pushing_Long_Horizon/README.md +36 -0
  36. 05_Ten_Task_Dynamics/CausalWorld_Pushing_Long_Horizon/checkpoints/CausalWorld_Pushing_Long_Horizon_final_checkpoint.pt +3 -0
  37. 05_Ten_Task_Dynamics/DMC_Cartpole_Pixels/README.md +36 -0
  38. 05_Ten_Task_Dynamics/DMC_Cartpole_Pixels/checkpoints/DMC_Cartpole_Pixels_final_checkpoint.pt +3 -0
  39. 05_Ten_Task_Dynamics/DMC_VB_Cheetah_Run/README.md +36 -0
  40. 05_Ten_Task_Dynamics/DMC_VB_Cheetah_Run/checkpoints/DMC_VB_Cheetah_Run_final_checkpoint.pt +3 -0
  41. 05_Ten_Task_Dynamics/Distracting_Control_Cartpole/README.md +36 -0
  42. 05_Ten_Task_Dynamics/Distracting_Control_Cartpole/checkpoints/Distracting_Control_Cartpole_final_checkpoint.pt +3 -0
  43. 05_Ten_Task_Dynamics/PDEBench_Burgers/README.md +36 -0
  44. 05_Ten_Task_Dynamics/PDEBench_Burgers/checkpoints/PDEBench_Burgers_final_checkpoint.pt +3 -0
  45. 05_Ten_Task_Dynamics/PDEBench_Navier_Stokes/README.md +36 -0
  46. 05_Ten_Task_Dynamics/PDEBench_Navier_Stokes/checkpoints/PDEBench_Navier_Stokes_final_checkpoint.pt +3 -0
  47. 05_Ten_Task_Dynamics/PDEBench_Reaction_Diffusion_1D/README.md +36 -0
  48. 05_Ten_Task_Dynamics/PDEBench_Reaction_Diffusion_1D/checkpoints/PDEBench_Reaction_Diffusion_1D_final_checkpoint.pt +3 -0
  49. 05_Ten_Task_Dynamics/PDEBench_Reaction_Diffusion_2D/README.md +36 -0
  50. 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