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` | one channel |
|
| 27 |
| Mask temporal length `L` | 16 patches |
|
| 28 |
| Masker parameter `pct_unmasked` | 0.45 |
|
| 29 |
-
| Checkpoint | epoch 10 of 10 (
|
| 30 |
| Encoder parameters | 12.69 M |
|
| 31 |
| Training run | [`46gmjq13`](https://wandb.ai/pierregtch/chan-inv-clf/runs/46gmjq13) |
|
| 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` | one channel |
|
| 27 |
| Mask temporal length `L` | 16 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 | [`46gmjq13`](https://wandb.ai/pierregtch/chan-inv-clf/runs/46gmjq13) |
|
| 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:b5d0ebf384c1e4754934d3f82f0b6270488a91a6d1b6ac07ffecb19cc228924e
|
| 3 |
+
size 50771544
|
epoch_02/model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:97f43e0e479d2756cea3c6956639a530c3f0b5c996d60b48ef144a24769de626
|
| 3 |
+
size 50771544
|
epoch_03/model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:456f7474e2e5955edab6e83b4a613e8de562c6b93e50d613c897450a855d9425
|
| 3 |
+
size 50771544
|
epoch_04/model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:a8701b7d30880f2aef78f633c3f60a06c1937167a6c1115e4a684a46b72f6235
|
| 3 |
+
size 50771544
|
epoch_05/model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:ef86363c6e973263df4679556b0659ca432ca03d2da5bb41074be1c3c956701b
|
| 3 |
+
size 50771544
|
epoch_06/model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:746b6893384b2118c310d7460dd97f014256ec9b7740c6cabcf1ccf1c6f4478f
|
| 3 |
+
size 50771544
|
epoch_07/model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:accb609d7dcc02c8eeb976c2b0e547b340662690b5def0b4e376f234fb077306
|
| 3 |
+
size 50771544
|
epoch_08/model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:625ed12c923adf716f15a74beb383b1dc9934fa042a200881cb6ab237f8d2421
|
| 3 |
+
size 50771544
|
epoch_09/model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:79a7c48884845437b226536d95e471d4de5a5256579928d20abdb9503e081dae
|
| 3 |
+
size 50771544
|