brainomni-tiny-pretrained / convert_brainomni_checkpoints.py
bruAristimunha's picture
Add BrainOmni weights converted from OpenTSLab/BrainOmni@9a4d3c70 (braindecode PR w41/brainomni-rehost @ d6d9804d)
d111525 verified
Raw History Blame Contribute Delete
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)