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.
MSVTNet.forward(braindecode/models/msvtnet.py) drops the main classification head whenreturn_features=True: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: withreturn_features=Trueyou 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))) whenreturn_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.