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