Download convert_mapa_checkpoint.py from braindecode/mapa-pretrained: direct link, hf CLI and curl.
- Browser
- Download file 2.19 kB
-
https://huggingface.co/braindecode/mapa-pretrained/resolve/main/convert_mapa_checkpoint.py
- Command line
-
hf download hf://braindecode/mapa-pretrained/convert_mapa_checkpoint.py
-
curl -L -o convert_mapa_checkpoint.py https://huggingface.co/braindecode/mapa-pretrained/resolve/main/convert_mapa_checkpoint.py
2.19 kB
| """Convert the released MAPA checkpoint to a braindecode-native file. | |
| Source: https://huggingface.co/bentang18/MAPA at revision | |
| 988efbf31a7d1f38533b848c993a719d6f900b1f (Apache-2.0), file ``mapa_vits384.pt`` | |
| (sha256 2d236089a2f1a3cc2827e3f150c4a2ba14c51bbfaf0ce0888f84b92a6eb25a7a). | |
| Usage:: | |
| python convert_mapa_checkpoint.py OUT_DIR [SOURCE_FILE] | |
| writes ``OUT_DIR/mapa-pretrained`` with ``config.json``, ``model.safetensors`` | |
| and ``pytorch_model.bin`` (``save_pretrained``). Without ``SOURCE_FILE`` the | |
| file is downloaded. | |
| Key changes: the feed-forward ``encoder.blocks.{i}.mlp.fc1``/``fc2`` become the | |
| ``FeedForwardBlock`` children ``mlp.0``/``mlp.3``; every other key is kept. The | |
| classification head ``final_layer`` is not pretrained: it is a seeded random | |
| ``nn.Linear`` default init. The stored montage (4 channels, no labels or | |
| regions) is only a default; pass ``chs_info`` or ``n_chans``, and | |
| ``contact_labels`` and ``regions``, to ``from_pretrained``. | |
| """ | |
| import hashlib | |
| import sys | |
| from pathlib import Path | |
| import torch | |
| from braindecode.models import MAPA | |
| REPO, REVISION = "bentang18/MAPA", "988efbf31a7d1f38533b848c993a719d6f900b1f" | |
| FILENAME = "mapa_vits384.pt" | |
| SHA256 = "2d236089a2f1a3cc2827e3f150c4a2ba14c51bbfaf0ce0888f84b92a6eb25a7a" | |
| def convert(source, out): | |
| if source is None: | |
| from huggingface_hub import hf_hub_download | |
| source = hf_hub_download(REPO, FILENAME, revision=REVISION) | |
| assert hashlib.sha256(Path(source).read_bytes()).hexdigest() == SHA256 | |
| released = torch.load(source, map_location="cpu", weights_only=True)["model"] | |
| state = { | |
| key.replace(".mlp.fc1.", ".mlp.0.").replace(".mlp.fc2.", ".mlp.3."): value | |
| for key, value in released.items() | |
| } | |
| torch.manual_seed(0) # the head is a seeded random init | |
| model = MAPA(n_outputs=2, n_chans=4, n_times=2048, sfreq=2048) | |
| state.update( | |
| {k: v for k, v in model.state_dict().items() if k.startswith("final_layer.")} | |
| ) | |
| model.load_state_dict(state, strict=True) | |
| model.save_pretrained(out) | |
| return model | |
| if __name__ == "__main__": | |
| convert(sys.argv[2] if len(sys.argv) > 2 else None, Path(sys.argv[1]) / "mapa-pretrained") | |