Sensory-guided human-machine joint learning accelerates the acquisition of motor imagery brain computer interface control.
The 3 matches · 1 of them tie a paragraph to a whole file, not to given lines: a weak match, whose lines are not tinted
- [1] § Methods › Joint learning framework and sample reweighting algorithm ↔ onlineFeedback.py, the whole file · a weak match · score 0.63 · 4–40 Hz, bandpass filtered, downsampled, online, window, EEG
- [2] § Methods › Joint learning framework and sample reweighting algorithm ↔ updateModel.py, lines 701–735 · score 0.55 · linear discriminant, LDA, CSP, score, probability, filters
- [3] § Methods › Joint learning framework and sample reweighting algorithm ↔ updateModel.py, lines 995–1089 · score 0.50 · bandpass filtered, 4–40 Hz, window, training
Paper
Loaded from Europe PMC by your browser, not stored by OSCR: doi.org · Europe PMC
The paper is loaded when this pane is shown.
The authors' code
Python · 1,138 lines · 35 KB · MIT · 2 matches
- """
- Update Weighted EEGNet from public-release MATLAB runData files.
- Expected MATLAB structure:
- runData.trialSignal : samples x channels x trials
- runData.trialTargetClass : samples x trials, or one label per trial
- This script replaces the old updateWeightedEEGNet.py workflow that loaded .dat files directly.
- It does not use BCI2kReader and does not read raw .dat files.
- The user edits the parameters in the __main__ section and runs this script directly.
- """
- import re
- from pathlib import Path
- from typing import Any, Dict, List, Optional, Tuple
- import numpy as np
- import scipy.io as sio
- import scipy.signal as signal
- import torch
- from scipy.stats import rankdata
- from sklearn.discriminant_analysis import LinearDiscriminantAnalysis
- from sklearn.metrics import accuracy_score
- from torch import nn
- from torch.utils.data import Dataset, DataLoader
- from torch.utils.data.dataloader import default_collate
- from resources.CSP import CSP
- from resources.eegnet import EEGNetv4
- # -------------------------
- # Dataset wrappers
- # -------------------------
- class CustomDataset(Dataset):
- def __init__(self, data: np.ndarray, labels: np.ndarray):
- self.data = data
- self.labels = labels
- def __len__(self) -> int:
- return len(self.data)
- def __getitem__(self, index: int):
- return self.data[index], self.labels[index]
- class WeightedDataset(Dataset):
- def __init__(self, data: np.ndarray, labels: np.ndarray, sample_weights: np.ndarray):
- self.data = data
- self.labels = labels
- self.sample_weights = sample_weights
- def __len__(self) -> int:
- return len(self.data)
- def __getitem__(self, index: int):
- return self.data[index], self.labels[index], self.sample_weights[index]
- # -------------------------
- # MATLAB loading utilities
- # -------------------------
- def _mat_struct_to_dict(obj: Any) -> Any:
- """
- Recursively convert scipy.io MATLAB mat_struct objects to Python dictionaries.
- Important:
- This function preserves numpy structured arrays because hdf5storage can load
- MATLAB structs as numpy arrays with dtype.names.
- """
- if isinstance(obj, np.ndarray) and obj.dtype.names is not None:
- return obj
- if isinstance(obj, np.void) and obj.dtype.names is not None:
- return obj
- if hasattr(obj, "_fieldnames"):
- out = {}
- for name in obj._fieldnames:
- out[name] = _mat_struct_to_dict(getattr(obj, name))
- return out
- if isinstance(obj, np.ndarray) and obj.dtype == object:
- if obj.size == 1:
- return _mat_struct_to_dict(obj.item())
- return np.array([_mat_struct_to_dict(x) for x in obj.flat], dtype=object).reshape(obj.shape)
- return obj
- def load_mat_file(path: str) -> Dict[str, Any]:
- """
- Load MATLAB .mat file.
- Tries:
- 1. scipy.io.loadmat
- 2. mat73
- 3. hdf5storage
- Public release files saved as MATLAB -v7.3 may require mat73 or hdf5storage.
- """
- path = str(path)
- try:
- mat = sio.loadmat(path, squeeze_me=True, struct_as_record=False)
- return {k: _mat_struct_to_dict(v) for k, v in mat.items() if not k.startswith("__")}
- except NotImplementedError:
- pass
- except ValueError as exc:
- if "Unknown mat file type" not in str(exc) and "Please use HDF reader" not in str(exc):
- raise
- try:
- import mat73 # type: ignore
- return mat73.loadmat(path)
- except Exception:
- pass
- try:
- import hdf5storage # type: ignore
- return hdf5storage.loadmat(path)
- except Exception as exc:
- raise RuntimeError(
- "Could not load this .mat file. If it was saved with MATLAB -v7.3, "
- "install one of: pip install mat73, or pip install hdf5storage.\n"
- f"File: {path}\nOriginal error: {exc}"
- )
- def squeeze_item(x: Any) -> Any:
- """
- Squeeze simple arrays, but preserve structured arrays.
- Do not call .item() on structured arrays, otherwise dtype field names may be lost.
- """
- if isinstance(x, np.ndarray):
- if x.dtype.names is not None:
- return np.squeeze(x)
- x = np.squeeze(x)
- if x.shape == ():
- try:
- return x.item()
- except Exception:
- return x
- return x
- def list_mat_fields(obj: Any) -> List[str]:
- """
- List fields from MATLAB struct loaded as dict, mat_struct,
- numpy structured array, numpy void, or nested object array.
- """
- obj = squeeze_item(obj)
- if isinstance(obj, dict):
- return list(obj.keys())
- if hasattr(obj, "__dict__"):
- return [k for k in obj.__dict__.keys() if not k.startswith("_")]
- if isinstance(obj, np.void) and obj.dtype.names is not None:
- return list(obj.dtype.names)
- if isinstance(obj, np.ndarray):
- obj = np.squeeze(obj)
- if obj.dtype.names is not None:
- return list(obj.dtype.names)
- if obj.size == 1:
- try:
- return list_mat_fields(obj.item())
- except Exception:
- return []
- return []
- def get_mat_field(obj: Any, field_name: str) -> Any:
- """
- Robust field getter for MATLAB structs loaded as:
- - dict
- - scipy mat_struct
- - numpy structured array
- - numpy void
- - numpy object array containing one struct
- """
- obj = squeeze_item(obj)
- if isinstance(obj, dict):
- if field_name in obj:
- return squeeze_item(obj[field_name])
- raise KeyError(f"Field '{field_name}' not found. Available fields: {list(obj.keys())}")
- if hasattr(obj, field_name):
- return squeeze_item(getattr(obj, field_name))
- if hasattr(obj, "__dict__") and field_name in obj.__dict__:
- return squeeze_item(obj.__dict__[field_name])
- if isinstance(obj, np.void) and obj.dtype.names is not None:
- if field_name in obj.dtype.names:
- return squeeze_item(obj[field_name])
- raise KeyError(f"Field '{field_name}' not found. Available fields: {obj.dtype.names}")
- if isinstance(obj, np.ndarray):
- obj = np.squeeze(obj)
- if obj.dtype.names is not None:
- if field_name in obj.dtype.names:
- return squeeze_item(obj[field_name])
- raise KeyError(f"Field '{field_name}' not found. Available fields: {obj.dtype.names}")
- if obj.size == 1:
- return get_mat_field(obj.item(), field_name)
- raise KeyError(
- f"Field '{field_name}' not found. "
- f"Object type: {type(obj)}. "
- f"Available fields: {list_mat_fields(obj)}"
- )
- def has_mat_field(obj: Any, field_name: str) -> bool:
- try:
- get_mat_field(obj, field_name)
- return True
- except Exception:
- return False
- def get_nested(d: Dict[str, Any], keys: List[str], default: Any = None) -> Any:
- cur = d
- for key in keys:
- if isinstance(cur, dict) and key in cur:
- cur = cur[key]
- else:
- return default
- return cur
- def to_string_list(x: Any) -> List[str]:
- """
- Convert MATLAB string / cellstr / char-array-ish objects to list[str].
- If selected_channels is not parsed cleanly, training is not affected.
- """
- if x is None:
- return []
- try:
- x = squeeze_item(x)
- if isinstance(x, str):
- return [x]
- if isinstance(x, bytes):
- return [x.decode()]
- arr = np.asarray(x)
- if arr.dtype.kind in {"U", "S"}:
- return [str(v) for v in arr.ravel().tolist()]
- if arr.dtype == object:
- out = []
- for v in arr.ravel():
- if isinstance(v, bytes):
- out.append(v.decode())
- elif isinstance(v, str):
- out.append(v)
- elif isinstance(v, np.ndarray):
- out.append("".join(v.astype(str).ravel().tolist()))
- else:
- out.append(str(v))
- return out
- return [str(v) for v in arr.ravel().tolist()]
- except Exception:
- return []
- # -------------------------
- # Public runData extraction
- # -------------------------
- def _as_numeric_array(x: Any) -> np.ndarray:
- arr = np.asarray(x)
- if arr.dtype == object:
- raise ValueError("Expected numeric array but got object array/cell array.")
- return arr.astype(np.float64)
- def _cell_to_trials(cell_obj: Any) -> np.ndarray:
- """
- Convert MATLAB cell array of trial matrices into trials x channels x samples.
- This is not recommended for the update pipeline unless all trials have the same length.
- Public EEGNet-style update files should usually be numeric arrays, not cells.
- """
- arr = np.asarray(cell_obj, dtype=object)
- trials = []
- for item in arr.ravel():
- trial = np.asarray(item, dtype=np.float64)
- if trial.ndim != 2:
- trial = np.squeeze(trial)
- if trial.ndim != 2:
- raise ValueError("Each trialSignal cell must be a 2D matrix: samples x channels or channels x samples.")
- # Public data stores cells as trialLength x channels.
- if trial.shape[0] >= trial.shape[1]:
- trial = trial.T # channels x samples
- trials.append(trial)
- lengths = {t.shape[1] for t in trials}
- channels = {t.shape[0] for t in trials}
- if len(lengths) != 1 or len(channels) != 1:
- raise ValueError(
- "trialSignal is a variable-length cell array. Weighted update needs fixed-length EEGNet-style trials. "
- "Use fixed-length runData files or add a padding/cropping rule."
- )
- return np.stack(trials, axis=0) # trials x channels x samples
- def trial_signal_to_numpy(trial_signal: Any) -> np.ndarray:
- """
- Convert runData.trialSignal to trials x channels x samples.
- Expected fixed-length public format:
- samples x channels x trials
- """
- arr = np.asarray(trial_signal)
- if arr.dtype == object:
- return _cell_to_trials(arr)
- arr = arr.astype(np.float64)
- arr = np.squeeze(arr)
- if arr.ndim != 3:
- raise ValueError(f"trialSignal must be 3D numeric array or cell array. Got shape {arr.shape}.")
- # Most public EEGNet-style data:
- # samples x channels x trials, e.g. 5000 x 62 x 30.
- # Convert to trials x channels x samples.
- if arr.shape[0] > arr.shape[1] and arr.shape[2] < arr.shape[0]:
- return np.transpose(arr, (2, 1, 0))
- # If already trials x channels x samples.
- if arr.shape[2] > arr.shape[1] and arr.shape[0] < arr.shape[2]:
- return arr
- # Conservative fallback: assume samples x channels x trials.
- return np.transpose(arr, (2, 1, 0))
- def _mode_or_mean_label(v: np.ndarray) -> int:
- v = v[np.isfinite(v)]
- if v.size == 0:
- raise ValueError("Empty target label vector.")
- rounded = np.rint(v).astype(int)
- vals, counts = np.unique(rounded, return_counts=True)
- return int(vals[np.argmax(counts)])
- def target_class_to_labels(trial_target_class: Any, n_trials: int) -> np.ndarray:
- """
- Convert runData.trialTargetClass into one label per trial.
- Common fixed-length format:
- samples x trials
- """
- arr = np.asarray(trial_target_class)
- if arr.dtype == object:
- labels = []
- for item in arr.ravel():
- v = np.asarray(item, dtype=np.float64).ravel()
- labels.append(_mode_or_mean_label(v))
- return np.asarray(labels, dtype=int)
- arr = np.squeeze(arr.astype(np.float64))
- if arr.ndim == 1:
- if arr.size == n_trials:
- return np.rint(arr).astype(int)
- raise ValueError(f"1D trialTargetClass length {arr.size} does not match n_trials {n_trials}.")
- if arr.ndim != 2:
- raise ValueError(f"trialTargetClass must be 1D/2D/cell. Got shape {arr.shape}.")
- # samples x trials
- if arr.shape[1] == n_trials:
- return np.asarray([_mode_or_mean_label(arr[:, i]) for i in range(n_trials)], dtype=int)
- # trials x samples
- if arr.shape[0] == n_trials:
- return np.asarray([_mode_or_mean_label(arr[i, :]) for i in range(n_trials)], dtype=int)
- raise ValueError(f"Cannot align trialTargetClass shape {arr.shape} with n_trials {n_trials}.")
- def extract_structured_runData_field(run_data: np.ndarray, field_name: str) -> Any:
- """
- Extract a field from structured ndarray runData.
- Example current format:
- runData shape: (1,)
- runData dtype.names includes trialSignal
- runData["trialSignal"][0] has shape 5000 x 62 x trials
- """
- if not isinstance(run_data, np.ndarray) or run_data.dtype.names is None:
- raise TypeError("run_data is not a numpy structured array.")
- if field_name not in run_data.dtype.names:
- raise KeyError(f"runData.{field_name} not found. Available fields: {run_data.dtype.names}")
- out = run_data[field_name]
- if out.shape[0] == 1:
- out = out[0]
- return out
- def parse_meta_from_structured_runData(run_data: np.ndarray) -> Tuple[float, List[str]]:
- """
- Extract sampling_rate_hz and selected_channels from structured runData.meta if possible.
- If parsing fails, return fs=1000 and empty channel list.
- """
- fs = 1000.0
- selected_channels: List[str] = []
- if not isinstance(run_data, np.ndarray) or run_data.dtype.names is None:
- return fs, selected_channels
- if "meta" not in run_data.dtype.names:
- return fs, selected_channels
- try:
- meta = run_data["meta"][0]
- if isinstance(meta, np.ndarray) and meta.dtype.names is not None:
- if "sampling_rate_hz" in meta.dtype.names:
- fs_raw = meta["sampling_rate_hz"]
- if isinstance(fs_raw, np.ndarray):
- fs_raw = np.squeeze(fs_raw)
- if fs_raw.shape == ():
- fs = float(fs_raw.item())
- else:
- fs = float(np.asarray(fs_raw).ravel()[0])
- else:
- fs = float(fs_raw)
- if "selected_channels" in meta.dtype.names:
- try:
- selected_channels = to_string_list(meta["selected_channels"])
- except Exception:
- selected_channels = []
- elif isinstance(meta, np.void) and meta.dtype.names is not None:
- if "sampling_rate_hz" in meta.dtype.names:
- fs = float(np.asarray(meta["sampling_rate_hz"]).squeeze())
- if "selected_channels" in meta.dtype.names:
- try:
- selected_channels = to_string_list(meta["selected_channels"])
- except Exception:
- selected_channels = []
- except Exception as exc:
- print(f"Warning: could not parse meta cleanly: {exc}")
- fs = 1000.0
- selected_channels = []
- if not np.isfinite(fs) or fs <= 0:
- fs = 1000.0
- return fs, selected_channels
- def load_runData_trials(mat_path: str) -> Tuple[np.ndarray, np.ndarray, List[str], float, Dict[int, int]]:
- """
- Load already segmented runData file.
- Returns:
- sig: trials x channels x samples
- labels_zero_based: trials
- selected_channels: list[str]
- sampling_rate_hz: float
- label_mapping: dict from original label -> zero-based label
- """
- mat = load_mat_file(mat_path)
- if "runData" not in mat:
- raise KeyError(f"File does not contain runData: {mat_path}. Available keys: {list(mat.keys())}")
- run_data = mat["runData"]
- print("Top-level mat keys:", list(mat.keys()))
- print("runData type:", type(run_data))
- # if isinstance(run_data, np.ndarray):
- # print("runData shape:", run_data.shape)
- # print("runData dtype:", run_data.dtype)
- # print("runData dtype.names:", run_data.dtype.names)
- # print("runData ndim:", run_data.ndim)
- # Case 1: hdf5storage loaded MATLAB struct as structured ndarray.
- if isinstance(run_data, np.ndarray) and run_data.dtype.names is not None:
- trial_signal = extract_structured_runData_field(run_data, "trialSignal")
- trial_target_class = extract_structured_runData_field(run_data, "trialTargetClass")
- sampling_rate_hz, selected_channels = parse_meta_from_structured_runData(run_data)
- # Case 2: runData loaded as dict / mat_struct.
- else:
- if not isinstance(run_data, dict):
- run_data = _mat_struct_to_dict(run_data)
- if isinstance(run_data, dict):
- if "trialSignal" not in run_data:
- raise KeyError(f"runData.trialSignal not found. Available fields: {list(run_data.keys())}")
- if "trialTargetClass" not in run_data:
- raise KeyError(
- f"runData.trialTargetClass not found. Available fields: {list(run_data.keys())}. "
- "This update script is for EEGNet-style runData files."
- )
- trial_signal = run_data["trialSignal"]
- trial_target_class = run_data["trialTargetClass"]
- meta = run_data.get("meta", {}) if isinstance(run_data.get("meta", {}), dict) else {}
- selected_channels = to_string_list(meta.get("selected_channels", []))
- sampling_rate_hz = float(np.asarray(meta.get("sampling_rate_hz", 1000)).squeeze())
- else:
- trial_signal = get_mat_field(run_data, "trialSignal")
- trial_target_class = get_mat_field(run_data, "trialTargetClass")
- selected_channels = []
- sampling_rate_hz = 1000.0
- if has_mat_field(run_data, "meta"):
- meta = get_mat_field(run_data, "meta")
- if has_mat_field(meta, "sampling_rate_hz"):
- sampling_rate_hz = float(np.asarray(get_mat_field(meta, "sampling_rate_hz")).squeeze())
- if has_mat_field(meta, "selected_channels"):
- selected_channels = to_string_list(get_mat_field(meta, "selected_channels"))
- if not np.isfinite(sampling_rate_hz) or sampling_rate_hz <= 0:
- sampling_rate_hz = 1000.0
- sig = trial_signal_to_numpy(trial_signal)
- labels = target_class_to_labels(trial_target_class, sig.shape[0])
- # Convert arbitrary integer labels to contiguous 0..C-1 for PyTorch NLLLoss.
- unique_labels = sorted([int(x) for x in np.unique(labels)])
- label_mapping = {old: new for new, old in enumerate(unique_labels)}
- labels_zero_based = np.asarray([label_mapping[int(x)] for x in labels], dtype=np.int64)
- return sig, labels_zero_based, selected_channels, sampling_rate_hz, label_mapping
- # -------------------------
- # Signal preprocessing
- # -------------------------
- def bandpass_filter_trials(sig: np.ndarray, fs: float, low: float = 4.0, high: float = 40.0) -> np.ndarray:
- """Butterworth bandpass on trials x channels x samples."""
- nyq = fs / 2.0
- if high >= nyq:
- high = nyq - 1.0
- if low <= 0 or high <= low:
- raise ValueError(f"Invalid bandpass range: low={low}, high={high}, fs={fs}")
- b, a = signal.butter(4, [low, high], btype="bandpass", fs=fs)
- return signal.filtfilt(b, a, sig, axis=-1)
- def slice_trials(
- sig: np.ndarray,
- labels: np.ndarray,
- fs: float,
- window_sec: float = 1.0,
- step_sec: float = 0.04,
- start_sec: float = 0.5,
- end_sec: float = 4.5,
- ) -> Tuple[np.ndarray, np.ndarray]:
- """
- Slice trials into sliding 1-second windows from 0.5 to 4.5 s with 40 ms step.
- Input:
- sig: trials x channels x samples
- Output:
- windows x channels x window_samples
- """
- n_trials, n_ch, n_samples = sig.shape
- start = int(round(start_sec * fs))
- stop = int(round(end_sec * fs))
- win = int(round(window_sec * fs))
- step = int(round(step_sec * fs))
- if stop > n_samples:
- stop = n_samples
- if start < 0 or start + win > stop:
- raise ValueError(
- f"Invalid slicing window. n_samples={n_samples}, fs={fs}, "
- f"start={start}, stop={stop}, win={win}."
- )
- n_slices = int(np.floor((stop - start - win) / step) + 1)
- out = np.zeros((n_trials * n_slices, n_ch, win), dtype=np.float64)
- out_labels = np.zeros(n_trials * n_slices, dtype=np.int64)
- k = 0
- for i in range(n_trials):
- for j in range(n_slices):
- s = start + j * step
- out[k] = sig[i, :, s:s + win]
- out_labels[k] = labels[i]
- k += 1
- out = out - np.mean(out, axis=2, keepdims=True)
- return out, out_labels
- def resample_trials(sig: np.ndarray, fs_in: float, fs_out: float = 100.0) -> Tuple[np.ndarray, float]:
- if abs(fs_in - fs_out) < 1e-6:
- return sig, fs_in
- from fractions import Fraction
- frac = Fraction(fs_out / fs_in).limit_denominator(1000)
- resampled = signal.resample_poly(sig, up=frac.numerator, down=frac.denominator, axis=-1)
- return resampled, fs_out
- # -------------------------
- # Weighting utilities
- # -------------------------
- def cal_init_weight(train_data: np.ndarray, train_label: np.ndarray) -> np.ndarray:
- """CSP + LDA initial sample weights for binary labels."""
- train_label = np.ravel(train_label.astype(int))
- csp = CSP(m_filters=4)
- _, eig_vec = csp.fit(train_data, train_label)
- features = np.zeros((train_data.shape[0], 2 * csp.m_filters))
- for i in range(train_data.shape[0]):
- features[i, :] = csp.transform(train_data[i, :, :], eig_vec)
- lda = LinearDiscriminantAnalysis()
- lda.fit(features, train_label)
- pred_prob = lda.predict_proba(features)
- pred_label = lda.predict(features)
- acc = accuracy_score(pred_label, train_label)
- print(f"CSP TrainAcc: {acc:.4f}")
- weights = np.zeros(train_data.shape[0], dtype=np.float64)
- for i in range(train_data.shape[0]):
- if pred_label[i] == train_label[i]:
- weights[i] = np.max(pred_prob[i, :])
- else:
- weights[i] = np.min(pred_prob[i, :])
- return weights
- def cal_init_weight_2d(train_data: np.ndarray, train_label: np.ndarray) -> np.ndarray:
- """One-vs-rest CSP + LDA initial weights for multiclass labels."""
- train_label = np.ravel(train_label.astype(int))
- weights = np.zeros(train_data.shape[0], dtype=np.float64)
- classes = np.unique(train_label)
- for cls in classes:
- tmp_label = np.ones_like(train_label)
- cls_idx = np.where(train_label == cls)[0]
- tmp_label[cls_idx] = 0
- csp = CSP(m_filters=4)
- _, eig_vec = csp.fit(train_data, tmp_label)
- features = np.zeros((train_data.shape[0], 2 * csp.m_filters))
- for i in range(train_data.shape[0]):
- features[i, :] = csp.transform(train_data[i, :, :], eig_vec)
- lda = LinearDiscriminantAnalysis()
- lda.fit(features, tmp_label)
- pred_prob = lda.predict_proba(features)
- pred_tmp = lda.predict(features)
- acc = accuracy_score(pred_tmp, tmp_label)
- print(f"CSP one-vs-rest class {cls} TrainAcc: {acc:.4f}")
- for idx in cls_idx:
- if pred_tmp[idx] == tmp_label[idx]:
- weights[idx] = np.max(pred_prob[idx, :])
- else:
- weights[idx] = np.min(pred_prob[idx, :])
- return weights
- def select_init_data(data: np.ndarray, labels: np.ndarray, weights: np.ndarray, init_size: float):
- n = max(1, int(np.floor(init_size * labels.shape[0])))
- chosen = np.argsort(weights)[-n:][::-1]
- chosen_weight = np.zeros_like(weights)
- chosen_weight[chosen] = weights[chosen]
- return data[chosen], labels[chosen], weights[chosen], chosen_weight
- def select_init_data_2d(data: np.ndarray, labels: np.ndarray, weights: np.ndarray, init_size: float):
- chosen_weight = np.zeros_like(weights)
- classes = np.unique(labels)
- n_per_class = max(1, int(np.floor(init_size * (labels.shape[0] / len(classes)))))
- for cls in classes:
- cls_idx = np.where(labels == cls)[0]
- local = np.argsort(weights[cls_idx])[-n_per_class:][::-1]
- chosen_weight[cls_idx[local]] = weights[cls_idx[local]]
- mask = chosen_weight > 0
- return data[mask], labels[mask], weights[mask], chosen_weight
- def get_data_weight_round(
- model,
- loss_fn,
- dataset,
- train_data,
- train_label,
- round_size,
- device,
- ):
- all_loss = []
- all_weight = np.zeros(train_label.shape[0], dtype=np.float64)
- model.eval()
- loader = DataLoader(dataset, batch_size=len(dataset), shuffle=False)
- with torch.no_grad():
- for x, y in loader:
- x = x.to(device=device, dtype=torch.float32)
- y = torch.squeeze(y).to(device=device).long()
- pred = model(x)
- sample_loss = loss_fn(pred, y)
- all_loss.extend(sample_loss.detach().cpu().numpy().tolist())
- norm_loss = rankdata(all_loss, nan_policy="omit") / len(all_loss)
- for i, norm_val in enumerate(norm_loss):
- if norm_val <= round_size:
- if round_size >= 1:
- epsilon = 1e-2
- all_weight[i] = np.log(norm_val + epsilon) / np.log(epsilon)
- else:
- numerator = norm_val + 1 - round_size
- denominator = np.log(1 - round_size)
- if numerator > 0 and denominator != 0:
- all_weight[i] = np.log(numerator) / denominator
- else:
- all_weight[i] = 0.0
- all_weight = np.nan_to_num(all_weight, nan=0.0, posinf=0.0, neginf=0.0)
- all_weight[all_weight < 0] = 0.0
- wmax = np.nanmax(all_weight)
- if not np.isfinite(wmax) or wmax == 0:
- print("Warning: round weights are all zero; using constant 0.1 weights.")
- all_weight[:] = 0.1
- else:
- all_weight /= wmax
- mask = all_weight > 0
- return train_data[mask], train_label[mask], all_weight[mask], all_weight
- # -------------------------
- # Training
- # -------------------------
- def train_weighted(
- model,
- loss_fn,
- max_epoch,
- train_dataset,
- device,
- lr=0.005,
- wd=1e-4,
- batch_size=128,
- ):
- model.to(device)
- model = model.to(torch.float32)
- optimizer = torch.optim.Adam(model.parameters(), lr=lr, weight_decay=wd)
- scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, max(max_epoch - 1, 1))
- last_loss = torch.tensor(float("nan"), device=device)
- for epoch in range(max_epoch):
- model.train()
- loader = DataLoader(
- train_dataset,
- batch_size=batch_size,
- shuffle=True,
- drop_last=False,
- collate_fn=lambda x: [y.to(device) for y in default_collate(x)],
- )
- for data, label, sample_weight in loader:
- optimizer.zero_grad()
- data = data.to(torch.float32)
- label = torch.squeeze(label).long()
- sample_weight = sample_weight.to(torch.float32)
- pred = model(data)
- loss_per_sample = loss_fn(pred, label)
- denom = torch.sum(sample_weight).clamp_min(1e-8)
- loss = torch.sum(loss_per_sample * sample_weight) / denom
- loss.backward()
- optimizer.step()
- last_loss = loss.detach()
- scheduler.step()
- print(f"Epoch {epoch + 1:03d}/{max_epoch}, Train Loss: {last_loss.item():.6f}")
- return model
- def fit_weighted_eegnet(
- train_data: np.ndarray,
- train_label: np.ndarray,
- previous_model_path: Optional[str] = None,
- init_size: float = 0.2,
- rounds: int = 8,
- max_epoch_per_round: int = 30,
- device: Optional[str] = None,
- ):
- train_label = np.asarray(train_label, dtype=np.int64).ravel()
- classes = np.unique(train_label)
- if classes.min() != 0 or classes.max() != len(classes) - 1:
- raise ValueError(f"Labels must be contiguous 0..C-1. Got {classes}.")
- if device is None:
- device = "cuda" if torch.cuda.is_available() else "cpu"
- device_obj = torch.device(device)
- print(f"Using device: {device_obj}")
- print(f"Training data shape: {train_data.shape}")
- print(f"Training labels: {classes}")
- all_weight_log = []
- if len(classes) <= 2:
- init_weight = cal_init_weight(train_data, train_label)
- round_data, round_label, round_weight, all_weight = select_init_data(
- train_data,
- train_label,
- init_weight,
- init_size,
- )
- else:
- init_weight = cal_init_weight_2d(train_data, train_label)
- round_data, round_label, round_weight, all_weight = select_init_data_2d(
- train_data,
- train_label,
- init_weight,
- init_size,
- )
- all_dataset = CustomDataset(train_data, train_label)
- model = EEGNetv4(
- train_data.shape[1],
- len(classes),
- n_times=train_data.shape[2],
- add_log_softmax=True,
- kernel_length=min(50, train_data.shape[2]),
- )
- if previous_model_path:
- print(f"Loading previous model: {previous_model_path}")
- state = torch.load(previous_model_path, map_location=device_obj)
- model.load_state_dict(state)
- loss_fn = nn.NLLLoss(reduction="none")
- for round_idx in range(rounds):
- if round_idx == 0:
- round_size = init_size
- print(f"Init with {init_size}")
- else:
- round_size = (1 - init_size) / (rounds - 1) * round_idx + init_size
- if round_size > 1 - 1e-2:
- round_size = 1.0
- print(f"Round {round_idx + 1}/{rounds}, round size: {round_size:.4f}")
- round_data, round_label, round_weight, all_weight = get_data_weight_round(
- model,
- loss_fn,
- all_dataset,
- train_data,
- train_label,
- round_size,
- device_obj,
- )
- all_weight_log.append(all_weight.copy())
- train_dataset = WeightedDataset(
- round_data,
- round_label,
- round_weight.astype(np.float32),
- )
- model = train_weighted(
- model,
- loss_fn,
- max_epoch_per_round,
- train_dataset,
- device_obj,
- )
- return model, all_weight_log
- # -------------------------
- # Main
- # -------------------------
- def derive_run_name(mat_path: str) -> str:
- base = Path(mat_path).stem
- # S001_sess01_run02 -> sess01_run02
- m = re.match(r"S\d{3}_(.+)$", base)
- return m.group(1) if m else base
- def main(
- mat_path,
- save_dir=r".\Parameter",
- previous_model_path="",
- output_prefix="",
- apply_bandpass=True,
- target_fs=100.0,
- init_size=0.2,
- rounds=8,
- epochs=30,
- device=None,
- ):
- save_dir = Path(save_dir)
- save_dir.mkdir(parents=True, exist_ok=True)
- print(f"MAT file: {mat_path}")
- print(f"Previous model: {previous_model_path if previous_model_path else '[none]'}")
- print(f"Save directory: {save_dir}")
- sig, label, channel_names, fs, label_mapping = load_runData_trials(mat_path)
- print(f"Loaded trialSignal: {sig.shape} [trials x channels x samples]")
- print(f"Original sampling rate: {fs}")
- print(f"Label mapping original->zero-based: {label_mapping}")
- if channel_names:
- print(f"Channels: {len(channel_names)}")
- if apply_bandpass:
- print("Applying 4-40 Hz bandpass filter.")
- sig = bandpass_filter_trials(sig, fs=fs, low=4.0, high=40.0)
- else:
- print("Bandpass filtering disabled.")
- sig, label = slice_trials(
- sig,
- label,
- fs=fs,
- window_sec=1.0,
- step_sec=0.04,
- start_sec=0.5,
- end_sec=4.5,
- )
- print(f"After sliding-window slicing: {sig.shape}")
- sig, fs_out = resample_trials(sig, fs_in=fs, fs_out=target_fs)
- print(f"After resampling to {fs_out:g} Hz: {sig.shape}")
- model, all_weight_log = fit_weighted_eegnet(
- train_data=sig,
- train_label=label,
- previous_model_path=previous_model_path if previous_model_path else None,
- init_size=init_size,
- rounds=rounds,
- max_epoch_per_round=epochs,
- device=device,
- )
- prefix = output_prefix if output_prefix else derive_run_name(mat_path)
- model_path = save_dir / f"{prefix}Model.pth"
- weight_path = save_dir / f"{prefix}WeightLog.mat"
- mapping_path = save_dir / f"{prefix}LabelMapping.mat"
- torch.save(model.state_dict(), model_path)
- sio.savemat(
- weight_path,
- {
- "AllWeightLog": np.asarray(all_weight_log, dtype=object)
- }
- )
- sio.savemat(
- mapping_path,
- {
- "original_labels": np.array(list(label_mapping.keys())),
- "zero_based_labels": np.array(list(label_mapping.values())),
- }
- )
- print(f"Saved model: {model_path}")
- print(f"Saved weight log: {weight_path}")
- print(f"Saved label mapping: {mapping_path}")
- return {
- "model_path": str(model_path),
- "weight_path": str(weight_path),
- "mapping_path": str(mapping_path),
- "label_mapping": label_mapping,
- "fs_out": fs_out,
- "data_shape": sig.shape,
- }
- if __name__ == "__main__":
- # ============================================================
- # User parameters: edit these directly
- # ============================================================
- mat_path = r"D:\Tactile\PublicRelease_Final\JointLearning\S001\S001_sess04_run03.mat"
- save_dir = r".\Parameter"
- # First run: use ""
- # Later run: use previous model, for example:
- # previous_model_path = r".\Parameter\sess01_run01Model.pth"
- previous_model_path = ""
- # If empty, output prefix is derived from mat filename.
- output_prefix = ""
- apply_bandpass = True
- target_fs = 100.0
- init_size = 0.2
- rounds = 9
- epochs = 30
- # None = cuda if available, otherwise cpu.
- # Or set "cuda" / "cpu".
- device = None
- # ============================================================
- # Run update
- # ============================================================
- result = main(
- mat_path=mat_path,
- save_dir=save_dir,
- previous_model_path=previous_model_path,
- output_prefix=output_prefix,
- apply_bandpass=apply_bandpass,
- target_fs=target_fs,
- init_size=init_size,
- rounds=rounds,
- epochs=epochs,
- device=device,
- )
- print("Finished.")
- print(result)
updateModel.py at commit 57710f7, under MIT · at the source
Overview
- Department of Biomedical Engineering, Carnegie Mellon University,Pittsburgh, PA USA
- Department of Electrical and Computer Engineering, Carnegie Mellon University,Pittsburgh, PA USA
- Neuroscience Institute, Carnegie Mellon University,Pittsburgh, PA USA
Abstract
The abstract is not reproduced here: the paper's license (CC BY-NC-ND) does not allow it. Read it in the paper, at the publisher or on Europe PMC.
Repositories
Its files are read in the Code ↔ Paper reader above, with 3 matches between paragraphs and lines of code.
bfinl/SensoryGuidedJointLearning
57710f76a5ddd6f9a3aad28566ca55776a424400, 2 June 2026Availability: 1 check, the latest on 27 September 2026: the link answers
- 27 September 2026: the link answers
8 files
- onlineFeedback.py, Python, 156 lines, 1 match
- resources/
CSP.py , Python, 38 lines - resources/
eegnet.py , Python, 420 lines - resources/
functions.py , Python, 411 lines - resources/
modules.py , Python, 573 lines - updateModel.py, Python, 1,138 lines, 2 matches
- LICENSE, License, 21 lines
- README.md, Text, 42 lines
Zenodo 20480141
Availability: 1 check, the latest on 27 September 2026: the link answers (HTTP 200)
- 27 September 2026: the link answers (HTTP 200)
8 files
- onlineFeedback.py, Python, 156 lines
- resources/
CSP.py , Python, 38 lines - resources/
eegnet.py , Python, 420 lines - resources/
functions.py , Python, 411 lines - resources/
modules.py , Python, 573 lines - updateModel.py, Python, 1,138 lines
- LICENSE, License, 21 lines
- README.md, Text, 42 lines
Code availability statement
The paper has a code availability statement. Its license (CC BY-NC-ND) does not allow reproducing it here; in short, from what the harvester recognized in it:
- it points to the authors' code: bfinl/
SensoryGuidedJointLearni , Zenodo 20480141ng
Read it in the paper: doi.org/10.1038/s41467-026-75435-5.
Tracing map
Proposed by the machine: these links were found in the paper and verified at the source, without human review. The map will receive a Zenodo DOI once one of the paper's authors has validated it with their ORCID.
What the map holds:
- 2 repositories of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
- 12 scripts, each with its path and the digest of its content;
- 3 matches between paragraphs of the paper and lines of the code (method lexical-v1);
- neither the text of the paper nor the code itself.
Its JSON (tracing-map.json) is deposited on Zenodo with its DOI once the map is validated.
Data
No dataset and no data link were found in the paper.
Data availability statement
The paper has a data availability statement. Its license (CC BY-NC-ND) does not allow reproducing it here; in short, from what the harvester recognized in it:
- no repository, dataset or request procedure was recognized in it
Read it in the paper: doi.org/10.1038/s41467-026-75435-5.
Versions
The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.
Version 1, 27 September 2026: the first record
Recorded: type, language, journal, volume, issue, pages, dates, 5 authors, 5 keywords, 14 MeSH terms, 2 funders, 74 references.
Cite
This paper
Wang, H., Zhang, Y., Karrenbach, M., Ding, Y., & He, B. (2026). Sensory-guided human-machine joint learning accelerates the acquisition of motor imagery brain computer interface control. Nature communications, 17(1), 6177. https://
BibTeX
@article{wang2026sensory
author = {Wang, Hanwen and Zhang, Yisha and Karrenbach, Maxim and Ding, Yidan and He, Bin},
title = {{Sensory-guided human-machine joint learning accelerates the acquisition of motor imagery brain computer interface control}},
journal = {Nature communications},
year = {2026},
month = jul,
volume = {17},
number = {1},
pages = {6177},
publisher = {Nature Publishing Group},
issn = {2041-1723},
doi = {10.1038/
url = {https://
pmid = {42457692},
pmcid = {PMC13373212}
}
RIS
TY - JOUR
AU - Wang, Hanwen
AU - Zhang, Yisha
AU - Karrenbach, Maxim
AU - Ding, Yidan
AU - He, Bin
TI - Sensory-guided human-machine joint learning accelerates the acquisition of motor imagery brain computer interface control
T2 - Nature communications
J2 - Nat Commun
PY - 2026
DA - 2026/
VL - 17
IS - 1
SP - 6177
SN - 2041-1723
PB - Nature Publishing Group
DO - 10.1038/
UR - https://
LA - en
ER -
CSL-JSON
{
"id": "10.1038/
"type": "article-journal",
"title": "Sensory-guided human-machine joint learning accelerates the acquisition of motor imagery brain computer interface control",
"container-title": "Nature communications",
"author": [
{
"family": "Wang",
"given": "Hanwen"
},
{
"family": "Zhang",
"given": "Yisha"
},
{
"family": "Karrenbach",
"given": "Maxim"
},
{
"family": "Ding",
"given": "Yidan"
},
{
"family": "He",
"given": "Bin"
}
],
"container-title-short":
"volume": "17",
"issue": "1",
"page": "6177",
"DOI": "10.1038/
"PMID": "42457692",
"PMCID": "PMC13373212",
"ISSN": "2041-1723",
"publisher": "Nature Publishing Group",
"URL": "https://
"language": "en",
"issued": {
"date-parts": [
[
2026,
7,
15
]
]
}
}
The tracing map gets a citation of its own once an author has validated it and it has a DOI.
Similar papers
The papers with a page that share the most with this one: the tools found in their code, their categories, datasets, cited references and authors, the rarest counting most.
- [1] doi:10.1038/s41467-026-69853-8 [code]
- Transcranial focused ultrasound induces source localizable cortical activation in resting state humans when applied concurrently with transcranial electric stimulation.Journal: Nature communicationsIn common: MNE-Python, SciPy, NumPy, EEG, 2 authors
- [2] doi:10.1186/s12984-026-02041-3 [code]
- Mental tasks induce common modulations of oscillations in cortex and spinal cord.Journal: Journal of neuroengineering and rehabilitationIn common: EEG, 6 references
- [3] doi:10.1371/journal.pcbi.1014112 [code]
- Opposing cortical forces: Alpha slowing and sensorimotor mu acceleration during motor-related BCI training.Journal: PLoS computational biologyIn common: EEG, 6 references
- [4] doi:10.1073/pnas.2510015122 [code]
- Mapping epileptogenic brain using a unified spatial–temporal–spectra
l source imaging framework Journal: n/aIn common: EEG, 2 references, author Bin He - [5] doi:10.1038/s44385-026-00098-2 [code]
- Cross-region neural signal reconstruction to lift electrode placement constraints in SSVEP brain-computer interfaces.Journal: npj biomedical innovationsIn common: EEG, 5 references
- [6] doi:10.3390/s26134045
- Brain Signal for Secure EEG Biometric Authentication: A Comprehensive Survey.Journal: Sensors (Basel, Switzerland)In common: EEG, 6 references
- [7] doi:10.1162/imag.a.1259 [code]
- Neuronal avalanches as a predictive biomarker for guiding tailored BCI training programs.Journal: Imaging neuroscience (Cambridge, Mass.)In common: MNE-Python, scikit-learn, SciPy, 1 other tool, EEG, 2 references
- [8] doi:10.1007/s10916-026-02374-5 [code]
- Attention-Enhanced U-Net for Sensor-Efficient High-Density EEG Reconstruction in Wearable Brain Monitoring Systems.Journal: Journal of medical systemsIn common: MNE-Python, PyTorch, scikit-learn, 2 other tools, EEG, 2 references
- [9] doi:10.1038/s41593-026-02258-4 [code]
- Laminar organization of cellular microcircuits modulating human interictal epileptiform discharges.Journal: Nature neuroscienceIn common: MNE-Python, scikit-learn, SciPy, 1 other tool, EEG, 2 references
- [10] doi:10.1186/s13634-026-01330-2 [code]
- Leednet: a lightweight network for event detection in EEG signals.Journal: Journal on advances in signal processingIn common: MNE-Python, PyTorch, scikit-learn, 2 other tools, EEG, 2 references
Contribute
The authors of this paper can claim it, correct its record and validate its tracing map, and the maintainers of its code (its owner, or a public member of its organization) correct what it says of their repository; anyone signed in can ask for its removal. Every request goes to OSCR's own machine, which answers it; your account page follows them.
Sign in with ORCID to claim this paper as one of its authors, correct its record or validate its tracing map: when the paper's metadata lists your ORCID iD, you are recognized at once. Maintainers of its code: sign in with GitHub, then claim the repository on your account page.
Claim this paper
Correct its record
Say what each link of this record is, remove the ones that are not the paper's, add the ones that are missing. The correction becomes a new version of the record, in its Versions section.
Validate its tracing map
You validate the map as this page shows it: 2 repositories of the authors' code, each at its verified commit and with its license, 12 scripts, and 3 matches between paragraphs and code (see the Code and Map sections). It then receives a DOI on Zenodo, with you (your ORCID iD) and OSCR as its creators; the code itself is not deposited.
The map's fingerprint: sha256:584eaa3e2b38c158…
Add the badge to its README
The badge links the code to this page. Copy one of these into the README of the paper's code: only you decide where it goes, and nothing is changed for you.
Markdown
[, paste the snippet at the top, then “Commit changes…” and, to review it first, “Create a new branch and start a pull request”. You open the pull request; OSCR asks for no permission.
Request its removal
To ask OSCR to remove this record, the copies of its authors' scripts or its tracing map, use the removal request page: signed in, you say who you are, what to remove and why, then review and confirm the request. Published rules decide every request (how).
Discussion, reproductions, activity
Discussion: questions and error reports about this paper and its code, from signed-in readers and its authors. It opens with sign-in.
Reproductions: reports from readers who ran the authors' code: what they reproduced, with which environment, commit and data. It opens with sign-in.
Activity: what happens around this paper: new versions of its record, its map's validation, discussions and reproductions. It opens with sign-in.
