Skip to content

FBMSNet: expose source features and document center-loss training details #1082

Description

@bruAristimunha

Problem

Braindecode's FBMSNet.forward() currently returns only logits. The released FBMSNet source returns both the class output and the flattened pre-classifier feature:

f = torch.flatten(x, start_dim=1)
c = self.fc(f)
return c, f

That feature output is not incidental: the FBMSNet paper/source trains with cross-entropy plus center loss on the flattened feature vector. Without a native way to get the feature tensor, downstream source-faithful reproductions need a local subclass that copies most of FBMSNet.forward() only to return (logits, features).

This issue is separate from the filter-bank phase issue opened in #1081. #1081 is about causal filter-bank preprocessing. This issue is about FBMSNet model/training details.

Current Braindecode behavior

In braindecode.models.FBMSNet (1.6.1dev0 checked locally), the constructor has the source-relevant architecture knobs:

FBMSNet(
    n_bands=9,
    n_filters_spat=36,
    temporal_layer="LogVarLayer",
    stride_factor=4,
    dilatability=8,
    kernels_weights=(15, 31, 63, 125),
    cnn_max_norm=2,
    linear_max_norm=0.5,
    filter_parameters=None,
)

but forward() flattens and classifies internally, then returns only x:

x = self.flatten_layer(x)
x = self.final_layer(x)
return x

Source details to align with

The released FBMSNet model source uses these defaults/details:

class FBMSNet(nn.Module):
    def __init__(..., temporalLayer='LogVarLayer', num_Feat=36,
                 dilatability=8, dropoutP=0.5, ...):
        self.strideFactor = 4
        self.mixConv2d = MixedConv2d(
            in_channels=9,
            out_channels=num_Feat,
            kernel_size=[(1,15),(1,31),(1,63),(1,125)],
        )
        self.scb = self.SCB(..., out_chan=num_Feat*dilatability, ...)
        self.temporalLayer = LogVarLayer(dim=3)
        self.fc = LinearWithConstraint(..., max_norm=0.5) + LogSoftmax

    def forward(self, x):
        ...
        f = torch.flatten(x, start_dim=1)
        c = self.fc(f)
        return c, f

The source training path uses the returned feature for center loss:

centerloss = CenterLoss(num_classes=4, feat_dim=1152)
optimzer4center = optim.SGD(centerloss.parameters(), lr=0.1)
...
output, feature = self.net(...)
closs = centerloss(labels, feature)
loss = lossFn(output, labels) / batch_size
total_loss = loss + 0.0005 * closs
total_loss.backward()
optimizer.step()
optimzer4center.step()

The source CenterLoss implementation initializes centers from NumPy seed 19981127, computes

(feature - centers[label]).pow(2).sum() / 2 / batch_size

and applies a count-normalized gradient update to the centers.

The source also runs a second training stage that reloads/continues after stage 1 and trains with the validation data merged back into the training data. In our reproduction config this corresponds to stage2_max_epochs=600.

Why this matters

The FBMSNet paper explicitly frames the method as CE + center loss, not CE-only. The source hard-codes a 1152-dimensional center-loss feature for the standard 4-class BCIC-IV-2a setup (n_filters_spat=36, dilatability=8, stride_factor=4 => 36*8*4 = 1152).

Downstream, the only remaining local Braindecode model override in our reproducibility pipeline is:

class SourceFBMSNet(FBMSNet):
    def forward(self, x):
        ...
        feats = self.flatten_layer(x)
        logits = self.final_layer(feats)
        return logits, feats

Everything else we previously carried for Braindecode model fidelity has moved upstream, but this one still requires a subclass.

Requested enhancement

A small API addition would remove the need for downstream subclasses, for example:

  1. Add return_features: bool = False to FBMSNet, matching the pattern already used in some other Braindecode models.
  2. When return_features=True, return (logits, features) where features is the output of flatten_layer immediately before final_layer.
  3. Optionally document/source-test the FBMSNet center-loss training details:
    • feature dim out_channels_spatial * stride_factor (1152 for the default source setup);
    • center-loss lambda 0.0005;
    • source center optimizer SGD(lr=0.1);
    • center init seed 19981127;
    • second-stage training on train+validation for up to 600 epochs.

A minimal test could assert that FBMSNet(..., return_features=True)(x) returns (logits, features) and that features.shape[1] == model.out_channels_spatial * model.stride_factor for the default FBMSNet setup.

Activity

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

Metadata

Metadata

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