File size: 2,190 Bytes
56c68f7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
"""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")