Add BrainOmni weights converted from OpenTSLab/BrainOmni@9a4d3c70 (braindecode PR w41/brainomni-rehost @ d6d9804d)
d111525 verified Download convert_brainomni_checkpoints.py from braindecode/brainomni-tiny-pretrained: direct link, hf CLI and curl.
- Browser
- Download file 4.44 kB
-
https://huggingface.co/braindecode/brainomni-tiny-pretrained/resolve/main/convert_brainomni_checkpoints.py
- Command line
-
hf download hf://braindecode/brainomni-tiny-pretrained/convert_brainomni_checkpoints.py
-
curl -L -o convert_brainomni_checkpoints.py https://huggingface.co/braindecode/brainomni-tiny-pretrained/resolve/main/convert_brainomni_checkpoints.py
4.44 kB
| """Convert the released BrainOmni checkpoints to braindecode-native files. | |
| Source: https://huggingface.co/OpenTSLab/BrainOmni at revision | |
| 9a4d3c70495370397ccfbfd6d2496f25647545a5 (MIT), files | |
| ``braintokenizer/BrainTokenizer.pt``, ``tiny/BrainOmni.pt``, ``base/BrainOmni.pt`` | |
| and their ``model_cfg.json``. | |
| Usage:: | |
| python convert_brainomni_checkpoints.py OUT_DIR [SOURCE_DIR] | |
| writes ``OUT_DIR/{braintokenizer,brainomni-tiny,brainomni-base}-pretrained`` | |
| with ``config.json``, ``model.safetensors`` and ``pytorch_model.bin`` | |
| (``save_pretrained``). Without ``SOURCE_DIR`` the files are downloaded. | |
| Key changes: ``quantizer.rvq.`` -> ``quantizer.``, ``decoder.`` -> | |
| ``final_layer.``, the doubled ``conv``/``convtr`` wrappers and SEANet's | |
| ``.model`` are flattened, ``weight_g``/``weight_v`` -> weight-norm | |
| parametrizations, the feed-forward linears are renumbered as | |
| ``FeedForwardBlock``. The pretraining-only ``mask_token`` and ``predict_head`` | |
| are dropped. The RoPE cache ``rotate`` is stored by the release as cosines only | |
| (the released code copies it into a complex cache, so the sine is zero); it is | |
| kept as ``(cos, 0)`` pairs. The BrainOmni classification head is not | |
| pretrained: it is a seeded random ``nn.Linear`` default init. | |
| The stored ``chs_info`` (19 EEG channels, 10-20) is only a default; pass | |
| your own ``chs_info`` to ``from_pretrained``. | |
| """ | |
| import json | |
| import sys | |
| from pathlib import Path | |
| import mne | |
| import torch | |
| from braindecode.models import BrainOmni, BrainTokenizer | |
| REPO, REVISION = "OpenTSLab/BrainOmni", "9a4d3c70495370397ccfbfd6d2496f25647545a5" | |
| RENAMES = { | |
| "n_dim": "emb_dim", | |
| "n_head": "tokenizer_num_heads", | |
| "lm_head": "num_heads", | |
| "lm_depth": "depth", | |
| "lm_dropout": "drop_prob", | |
| } | |
| UNUSED = {"mask_ratio", "num_quantizers_used", "quantize_optimize_method"} | |
| CH_NAMES = ( | |
| "Fp1 Fp2 F7 F3 Fz F4 F8 T7 C3 Cz C4 T8 P7 P3 Pz P4 P8 O1 O2".split() | |
| ) | |
| def fetch(source, path): | |
| if source is not None: | |
| return Path(source) / path | |
| from huggingface_hub import hf_hub_download | |
| return Path(hf_hub_download(REPO, path, revision=REVISION)) | |
| def rename(key): | |
| key = key.replace("quantizer.rvq.", "quantizer.") | |
| if key.startswith("decoder."): | |
| key = "final_layer." + key.removeprefix("decoder.") | |
| key = key.replace("tokenizer.decoder.", "tokenizer.final_layer.") | |
| key = key.replace(".convtr.convtr.", ".convtr.").replace(".conv.conv.", ".conv.") | |
| key = key.replace("seanet_encoder.model.", "seanet_encoder.") | |
| key = key.replace("seanet_decoder.model.", "seanet_decoder.") | |
| key = key.replace(".weight_g", ".parametrizations.weight.original0") | |
| key = key.replace(".weight_v", ".parametrizations.weight.original1") | |
| for name in ("ff", "aggregate_mlp"): | |
| key = key.replace(f"{name}.layer.0.", f"{name}.0.") | |
| key = key.replace(f"{name}.layer.2.", f"{name}.3.") | |
| return key | |
| def convert(source, folder, cls, out, **kwargs): | |
| config = json.loads(fetch(source, f"{folder}/model_cfg.json").read_text()) | |
| config = {RENAMES.get(k, k): v for k, v in config.items() if k not in UNUSED} | |
| dropout = config.pop("dropout") | |
| config["drop_prob" if cls is BrainTokenizer else "tokenizer_drop_prob"] = dropout | |
| info = mne.create_info(CH_NAMES, 256.0, "eeg") | |
| info.set_montage("standard_1020") | |
| torch.manual_seed(0) # the BrainOmni head is a seeded random init | |
| model = cls(chs_info=info["chs"], n_times=512, sfreq=256.0, **config, **kwargs) | |
| official = torch.load( | |
| fetch(source, f"{folder}/{cls.__name__}.pt"), weights_only=True | |
| ) | |
| state = {} | |
| for key, value in official.items(): | |
| if key == "mask_token" or key.startswith("predict_head."): | |
| continue | |
| if key.endswith("rope_embedding_layer.rotate"): | |
| value = torch.stack((value, torch.zeros_like(value)), dim=-1) | |
| state[rename(key)] = value | |
| for key, value in model.state_dict().items(): | |
| if key.startswith("final_layer.") and cls is BrainOmni: | |
| state[key] = value | |
| model.load_state_dict(state, strict=True) | |
| model.save_pretrained(out) | |
| return model | |
| if __name__ == "__main__": | |
| out = Path(sys.argv[1]) | |
| source = sys.argv[2] if len(sys.argv) > 2 else None | |
| convert(source, "braintokenizer", BrainTokenizer, out / "braintokenizer-pretrained") | |
| for size in ("tiny", "base"): | |
| convert(source, size, BrainOmni, out / f"brainomni-{size}-pretrained", n_outputs=2) | |