mirepnet-pretrained / convert_mirepnet_checkpoint.py
bruAristimunha's picture
Add Braindecode-format MIRepNet checkpoint
857f1e3 verified
Raw History Blame
1.3 kB
"""Convert the official MIRepNet checkpoint to Braindecode Hub format."""
from __future__ import annotations
import argparse
import hashlib
from pathlib import Path
import torch
from braindecode.models import MIRepNet
SOURCE_SHA256 = "432288958007e344a5a84a9ffe9d0e5e5c0cb616aef86c85522375a3f4da9aaf"
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("source", type=Path)
parser.add_argument("output", type=Path)
args = parser.parse_args()
if hashlib.sha256(args.source.read_bytes()).hexdigest() != SOURCE_SHA256:
raise ValueError("Source checkpoint SHA-256 does not match the official file.")
source = torch.load(args.source, map_location="cpu", weights_only=True)
model = MIRepNet(n_chans=45, n_outputs=3, n_times=1000, sfreq=250)
incompatible = model.load_state_dict(source, strict=False)
expected_unexpected = {
"mask_token",
"embedding.chan_embed.weight",
*(key for key in source if key.startswith("decoder.")),
}
if (
incompatible.missing_keys
or set(incompatible.unexpected_keys) != expected_unexpected
):
raise RuntimeError(f"Unexpected conversion result: {incompatible}")
model.save_pretrained(args.output)
if __name__ == "__main__":
main()