"""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()