Download convert_mirepnet_checkpoint.py from braindecode/mirepnet-pretrained: direct link, hf CLI and curl.
- Browser
- Download file 1.3 kB
-
https://huggingface.co/braindecode/mirepnet-pretrained/resolve/857f1e3642976be9f2ca6883e5508c3f7c91d86f/convert_mirepnet_checkpoint.py
- Command line
-
hf download hf://braindecode/mirepnet-pretrained@857f1e3642976be9f2ca6883e5508c3f7c91d86f/convert_mirepnet_checkpoint.py
-
curl -L -o convert_mirepnet_checkpoint.py https://huggingface.co/braindecode/mirepnet-pretrained/resolve/857f1e3642976be9f2ca6883e5508c3f7c91d86f/convert_mirepnet_checkpoint.py
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() | |