The Hub notes on every model page show replacing the head with reset_head(n_outputs=...) or from_pretrained(..., n_outputs=...) before saving. For 14 of the 19 models that implement reset_head, the saved model then can't be loaded again.
Steps to reproduce (master 4e332c4):
import tempfile
from braindecode.models import EEGPT
model = EEGPT(n_chans=8, n_outputs=4, n_times=1024, sfreq=256)
model.reset_head(10)
print(model.n_outputs, model.get_config()["n_outputs"]) # 10 4
with tempfile.TemporaryDirectory() as d:
model.save_pretrained(d)
EEGPT.from_pretrained(d)
RuntimeError: ... size mismatch for final_layer.probe2.parametrizations.weight.original:
copying a param with shape torch.Size([10, 496]) from checkpoint,
the shape in current model is torch.Size([4, 496]).
EEGPT.from_pretrained(repo, n_outputs=10) has the same effect: the model reports n_outputs == 10 while get_config()["n_outputs"] is still 4, so pushing the fine-tuned model to the Hub publishes a config that doesn't match its weights.
I ran reset_head followed by a save/load round trip on every registered model:
- fails: BENDR, BIOT, CBraMod, EEGDINO, EEGPT, Labram, MetaNeuromotorHand, MVPFormer, REVE, STEEGFormer, SignalJEPA_Contextual, SignalJEPA_PostLocal, SignalJEPA_PreLocal, ZUNA
- works: DANCE, EMG2QwertyNet, NeuroPose, SensingDynamics, VEMG2Pose
A second problem on the same path: the new head is built in train mode even when the model is in eval. EEGPT and EEGDINO have dropout in the head, so predictions stop being repeatable while model.training reports False:
import torch
from braindecode.models import EEGPT
x = torch.randn(2, 8, 1024)
model = EEGPT(n_chans=8, n_outputs=4, n_times=1024, sfreq=256).eval()
model.reset_head(10)
with torch.no_grad():
print(model.training, torch.equal(model(x), model(x))) # False False
from_pretrained(..., n_outputs=10) returns a model in the same state.
test_reset_head in test/unit_tests/models/test_return_features.py checks model.n_outputs but not the saved config, and calls model.eval() after reset_head, so neither problem shows up in CI.
I'd like to work on this, including a regression test, if that's welcome.
The Hub notes on every model page show replacing the head with
reset_head(n_outputs=...)orfrom_pretrained(..., n_outputs=...)before saving. For 14 of the 19 models that implementreset_head, the saved model then can't be loaded again.Steps to reproduce (master 4e332c4):
EEGPT.from_pretrained(repo, n_outputs=10)has the same effect: the model reportsn_outputs == 10whileget_config()["n_outputs"]is still 4, so pushing the fine-tuned model to the Hub publishes a config that doesn't match its weights.I ran
reset_headfollowed by a save/load round trip on every registered model:A second problem on the same path: the new head is built in train mode even when the model is in eval. EEGPT and EEGDINO have dropout in the head, so predictions stop being repeatable while
model.trainingreportsFalse:from_pretrained(..., n_outputs=10)returns a model in the same state.test_reset_headintest/unit_tests/models/test_return_features.pychecksmodel.n_outputsbut not the saved config, and callsmodel.eval()afterreset_head, so neither problem shows up in CI.I'd like to work on this, including a regression test, if that's welcome.