Skip to content

EEGRegressor computes its loss on a (batch, batch) broadcast when each trial has one target, and the model doesn't learn #1180

Description

@raghav-rathi

When a trialwise EEGRegressor is fit on a braindecode dataset with one target per trial (from create_from_X_y, or windows with a metadata target such as age), the model output is (batch, 1) and the target batch is (batch,). The default torch.nn.MSELoss broadcasts them to (batch, batch) and averages every prediction against every target in the batch.

Steps to reproduce (master 4e332c4):

import numpy as np
import torch
from braindecode import EEGRegressor
from braindecode.datasets import create_from_X_y
from braindecode.models import ShallowFBCSPNet

torch.manual_seed(0)
rng = np.random.RandomState(0)
X = rng.randn(16, 4, 250).astype("float32")
y = rng.randn(16).astype("float32")  # one number per trial
dataset = create_from_X_y(X, y, drop_last_window=False, sfreq=100)

net = EEGRegressor(
    ShallowFBCSPNet,
    module__final_conv_length="auto",
    train_split=None,
    batch_size=16,
    max_epochs=1,
    verbose=0,
)
net.fit(dataset, y=None)

X_batch, y_batch = next(iter(net.get_iterator(net.get_dataset(dataset))))
y_pred = net.infer(X_batch).detach()
print("shapes:", tuple(y_pred.shape), tuple(y_batch.shape))
print("loss used by EEGRegressor:", round(net.get_loss(y_pred, y_batch).item(), 4))
print("mean squared error per trial:", round(((y_pred.squeeze(1) - y_batch) ** 2).mean().item(), 4))
shapes: (16, 1) (16,)
loss used by EEGRegressor: 1.3908
mean squared error per trial: 1.8189

PyTorch warns during fit: Using a target size (torch.Size([16])) that is different to the input size (torch.Size([16, 1])). This will likely lead to incorrect results due to broadcasting.

What it does to training: each prediction is pulled toward the mean target of its batch rather than its own target. Same data, model and seed, 30 epochs, target = log amplitude of one channel (script below):

how fit is called correlation with held-out targets test MSE
net.fit(dataset) 0.018 0.1533
net.fit(X, y) with numpy arrays 0.931 0.0204
predict the mean for every trial n/a 0.1490

fit(X, y) works because EEGRegressor.fit reshapes a 1-D numpy y to (n, 1); targets that come from a dataset are not reshaped. None of the shipped examples go through this path: plot_regression.py and bcic_iv_4_ecog_cropped.py train in cropped mode, and bcic_iv_4_ecog_trial.py keeps a length-1 target with target_transform = lambda x: x[0:1].

Also noticed while reproducing: EEGRegressor.fit doesn't return self (regressor.py line 231 calls super().fit(...) without return), so EEGRegressor(...).fit(X, y).predict(X) raises AttributeError: 'NoneType' object has no attribute 'predict'. EEGClassifier.fit returns self.

Lemme know if i can work on this on my end.

Training comparison script
import numpy as np
import torch
from braindecode import EEGRegressor
from braindecode.datasets import create_from_X_y
from braindecode.models import ShallowFBCSPNet

rng = np.random.RandomState(0)
n, n_chans, n_times, sfreq = 600, 4, 250, 100
X = rng.randn(n, n_chans, n_times).astype("float32")
scale = rng.uniform(0.5, 2.0, n).astype("float32")
X[:, 0, :] *= scale[:, None]
y = np.log(scale).astype("float32")
train, test = slice(0, 500), slice(500, None)


def make_net():
    torch.manual_seed(0)
    return EEGRegressor(
        ShallowFBCSPNet,
        module__final_conv_length="auto",
        optimizer=torch.optim.AdamW,
        lr=1e-3,
        train_split=None,
        batch_size=32,
        max_epochs=30,
        verbose=0,
    )


ds_train = create_from_X_y(X[train], y[train], drop_last_window=False, sfreq=sfreq)
ds_test = create_from_X_y(X[test], y[test], drop_last_window=False, sfreq=sfreq)
net_dataset = make_net()
net_dataset.fit(ds_train, y=None)
pred_dataset = net_dataset.predict(ds_test).ravel()
net_numpy = make_net()
net_numpy.fit(X[train], y[train])
pred_numpy = net_numpy.predict(X[test]).ravel()

for name, pred in [("fit(dataset)", pred_dataset), ("fit(X, y)", pred_numpy)]:
    r = np.corrcoef(pred, y[test])[0, 1]
    print(name, round(r, 3), round(((pred - y[test]) ** 2).mean(), 4))
print("mean baseline", round(y[test].var(), 4))
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