Add intermediate checkpoints (epochs 1-9) and update model card
Browse files- README.md +14 -1
- epoch_01/model.safetensors +3 -0
- epoch_02/model.safetensors +3 -0
- epoch_03/model.safetensors +3 -0
- epoch_04/model.safetensors +3 -0
- epoch_05/model.safetensors +3 -0
- epoch_06/model.safetensors +3 -0
- epoch_07/model.safetensors +3 -0
- epoch_08/model.safetensors +3 -0
- epoch_09/model.safetensors +3 -0
README.md
CHANGED
|
@@ -26,7 +26,7 @@ This model is a **masked autoencoder (MAE)**: the encoder sees the unmasked patc
|
|
| 26 |
| Mask spatial radius `r` | 12 cm |
|
| 27 |
| Mask temporal length `L` | 4 patches |
|
| 28 |
| Masker parameter `pct_unmasked` | 0.45 |
|
| 29 |
-
| Checkpoint | epoch 10 of 10 (
|
| 30 |
| Encoder parameters | 12.69 M |
|
| 31 |
| Training run | [`4stxb7s2`](https://wandb.ai/pierregtch/chan-inv-clf/runs/4stxb7s2) |
|
| 32 |
|
|
@@ -38,6 +38,8 @@ The paper recommends r = 9 cm, L = 2: see [`eeg-fm-masking_mae_r9cm_L2`](https:/
|
|
| 38 |
`model.*`), i.e. exactly the tensors loaded for the downstream evaluation of the paper.
|
| 39 |
The MAE decoder is
|
| 40 |
not included; for JEPA the published weights are the student encoder, as evaluated in the paper.
|
|
|
|
|
|
|
| 41 |
* `config.json`: the keyword arguments of `ContextualEncoderBenchmarkWrapper` (architecture +
|
| 42 |
input scaling). Pass it unchanged as `model_kwargs`.
|
| 43 |
* `metadata.json`: masking parameters, training-run id, checkpoint epoch/step.
|
|
@@ -91,6 +93,17 @@ backbone = PretrainedBackbone(
|
|
| 91 |
)
|
| 92 |
```
|
| 93 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 94 |
## Downstream results (OpenEEGBench, frozen encoder + ridge probe)
|
| 95 |
|
| 96 |
Frozen encoder, ridge regression/classification on the flattened contextual features,
|
|
|
|
| 26 |
| Mask spatial radius `r` | 12 cm |
|
| 27 |
| Mask temporal length `L` | 4 patches |
|
| 28 |
| Masker parameter `pct_unmasked` | 0.45 |
|
| 29 |
+
| Checkpoint | end of epoch 10 of 10 (the one evaluated in the paper); epochs 1–9 in `epoch_01/` … `epoch_09/` |
|
| 30 |
| Encoder parameters | 12.69 M |
|
| 31 |
| Training run | [`4stxb7s2`](https://wandb.ai/pierregtch/chan-inv-clf/runs/4stxb7s2) |
|
| 32 |
|
|
|
|
| 38 |
`model.*`), i.e. exactly the tensors loaded for the downstream evaluation of the paper.
|
| 39 |
The MAE decoder is
|
| 40 |
not included; for JEPA the published weights are the student encoder, as evaluated in the paper.
|
| 41 |
+
* `epoch_01/` … `epoch_09/`: `model.safetensors` of the intermediate checkpoints (end of epochs
|
| 42 |
+
1 to 9 of the same run), with the same tensors and key names; they use the same `config.json`.
|
| 43 |
* `config.json`: the keyword arguments of `ContextualEncoderBenchmarkWrapper` (architecture +
|
| 44 |
input scaling). Pass it unchanged as `model_kwargs`.
|
| 45 |
* `metadata.json`: masking parameters, training-run id, checkpoint epoch/step.
|
|
|
|
| 93 |
)
|
| 94 |
```
|
| 95 |
|
| 96 |
+
### Intermediate checkpoints (epochs 1–9)
|
| 97 |
+
|
| 98 |
+
The end-of-epoch checkpoints of the same run are in the subfolders `epoch_01/` … `epoch_09/`
|
| 99 |
+
(the final epoch 10 is the `model.safetensors` at the root). Build `model` as above, then:
|
| 100 |
+
|
| 101 |
+
```python
|
| 102 |
+
epoch = 5 # 1 to 9
|
| 103 |
+
weights = load_file(hf_hub_download(repo, "model.safetensors", subfolder=f"epoch_{epoch:02d}"))
|
| 104 |
+
model.load_state_dict(weights, strict=False)
|
| 105 |
+
```
|
| 106 |
+
|
| 107 |
## Downstream results (OpenEEGBench, frozen encoder + ridge probe)
|
| 108 |
|
| 109 |
Frozen encoder, ridge regression/classification on the flattened contextual features,
|
epoch_01/model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:d7a058e2589263791fc4920a5c8152762d13126925cb9adfd0cfa32448c1a7d0
|
| 3 |
+
size 50771544
|
epoch_02/model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:83b6467c3aca4de9cfe39a7d344b7f3771a89751ac989398ffc5c35f6743acb3
|
| 3 |
+
size 50771544
|
epoch_03/model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:e3ed8555f911bc36103550323fe727c583b50de27e91608c9c983c827cd841ad
|
| 3 |
+
size 50771544
|
epoch_04/model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:628e266bf8d2fc2c410d4bc2cc81311205de863858920254fd66e5470bc2d7f3
|
| 3 |
+
size 50771544
|
epoch_05/model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:baa47aea2e0c7985c593a824e2911aca3e86c4b8d8f9e52290533420de12d5b5
|
| 3 |
+
size 50771544
|
epoch_06/model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:1e06579b4a6f4aa081fe5d2e41ecbe4507119506aa3cbbf05f5fbaa19feb2011
|
| 3 |
+
size 50771544
|
epoch_07/model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:d5d8cb63c917709ef0f6a7615a87eb03da243c897ca85e4bdf2d9aa4fc0b2e16
|
| 3 |
+
size 50771544
|
epoch_08/model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:30b9a7cc98b96b4bbeaaf1dfeee3f3c0ab50fcea52a5859f3859e035c21fceb9
|
| 3 |
+
size 50771544
|
epoch_09/model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:835805044746077ed2250c8f3cf653b5940a50ad6703b7da1713103e7a8e3be8
|
| 3 |
+
size 50771544
|