mapa-pretrained / convert_mapa_checkpoint.py
bruAristimunha's picture
Add MAPA mapa_vits384 converted from bentang18/MAPA@988efbf3 (Apache-2.0)
56c68f7 verified
Raw History Blame Contribute Delete
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")