Skip to content

reset_head() leaves the saved config at the old n_outputs, so a fine-tuned model can't be loaded back #1179

Description

@raghav-rathi

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.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions