braindecode.training.predict_trials#

braindecode.training.predict_trials(module, dataset, return_targets=True, batch_size=1, num_workers=0)[source]#

Create trialwise predictions and optionally also return trialwise targets from a cropped dataset given a module.

Parameters:
  • module (torch.nn.Module) – A pytorch model implementing forward.

  • dataset (braindecode.datasets.BaseConcatDataset) – A braindecode dataset to be predicted.

  • return_targets (bool) – If True, additionally returns the trial targets.

  • batch_size (int) – The batch size used to iterate the dataset.

  • num_workers (int) – Number of workers used in DataLoader to iterate the dataset.

Returns:

trial_predictions: np.ndarray | list of np.ndarray

3-dimensional array (n_trials x n_classes x n_predictions), where the number of predictions depend on the chosen window size and the receptive field of the network. If trials have different lengths, a list with one (n_classes x n_predictions) array per trial.

trial_targetsnp.ndarray | list of np.ndarray

Ground-truth targets from the dataset. Only returned when return_targets=True. An np.ndarray with a leading trial dimension: (n_trials,) for scalar targets, with further dimensions for multi-value or equal-length sequence targets. If trial lengths differ, a list with one target array per trial.