Skip to content

MSVTNet.forward returns branch predictions only (drops main head) when return_features=True #1069

Description

@bruAristimunha

MSVTNet.forward (braindecode/models/msvtnet.py) drops the main classification head when return_features=True:

x = self.final_layer(x)
if self.return_features:
    return torch.stack(branch_preds)   # branch predictions ONLY; main head `x` discarded
return x

The authors' MSVTNet (MSVTNet.py) returns both — return x, bx (main logits + branch logits) — which is required for the paper's joint deep-supervision loss (JointCrossEntoryLoss = lamd*CE(main) + (1-lamd)*sum(CE(branch))). As implemented, the joint loss is unreproducible: with return_features=True you get only the branches (forward(...).shape == [n_branches, B, n_classes]), and the main head is never supervised jointly.

Fix

Return (x, branch_preds) (or (x, torch.stack(branch_preds))) when return_features=True, matching the source, so callers can apply the joint loss and score on the main head. (Same class of issue as the EEGTCNet/ATCNet fidelity fixes in #1060/#1061.) Happy to PR.

Activity

  1. added a commit that references this issue on Jun 24, 2026
    1c05b41
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