English
tinymyo
emg
bio-signals
foundation-model
MatteoFasulo commited on
Commit
bc14906
·
unverified ·
1 Parent(s): adcb5a8

update db5 script with optional majority voting

Browse files
Files changed (1) hide show
  1. scripts/db5.py +91 -15
scripts/db5.py CHANGED
@@ -121,6 +121,49 @@ def augment_train_data(
121
  return new_data, new_labels
122
 
123
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
124
  def notch_filter(
125
  data: np.ndarray, notch_freq: float = 50.0, Q: float = 30.0, fs: float = 200.0
126
  ) -> np.ndarray:
@@ -178,6 +221,7 @@ def process_emg_features(
178
  rerep: np.ndarray,
179
  window_size: int = 1024,
180
  stride: int = 512,
 
181
  ) -> Tuple[np.ndarray, np.ndarray, np.ndarray]:
182
  """Segments raw EMG signals into overlapping windows.
183
 
@@ -196,18 +240,24 @@ def process_emg_features(
196
  """
197
  segs, lbls, reps = [], [], []
198
  N = len(label)
199
- for start in range(0, N, stride):
200
  end = start + window_size
201
- if end > N:
202
- cut = emg[start:N]
203
- pad = np.zeros((end - N, emg.shape[1]))
204
- win = np.vstack([cut, pad])
 
 
 
 
205
  else:
206
- win = emg[start:end]
 
 
207
 
208
  segs.append(win)
209
- lbls.append(label[start])
210
- reps.append(rerep[start])
211
  return np.array(segs), np.array(lbls), np.array(reps)
212
 
213
 
@@ -242,6 +292,18 @@ def main():
242
  default=3,
243
  help="Number of augmented versions to create for each training sample.",
244
  )
 
 
 
 
 
 
 
 
 
 
 
 
245
  args = args.parse_args()
246
 
247
  data_dir = args.data_dir
@@ -261,11 +323,17 @@ def main():
261
  os.system(f"rm {data_dir}/s{i}.zip")
262
  print(f"Downloaded and unzipped subject {i}\n{data_dir}/s{i}.zip")
263
 
264
- fs = 200.0 # original sampling rate
 
265
  window_size, stride = args.seq_len, args.stride
266
 
267
- window_seconds = sequence_to_seconds(window_size, fs)
 
 
 
 
268
  print(f"Window size: {window_size} samples ({window_seconds:.2f} seconds)")
 
269
 
270
  train_reps = [1, 3, 4, 6]
271
  val_reps = [2]
@@ -313,16 +381,24 @@ def main():
313
  elif "E3" in mat:
314
  label = np.where(label != 0, label + 29, 0)
315
 
316
- # filtering at original 200 Hz
317
- emg_filt = bandpass_filter_emg(emg, 20, 90, fs=fs)
318
- emg_filt = notch_filter(emg_filt, 50, 30, fs=fs)
 
 
 
 
 
 
319
 
320
  # z-score
321
- emg_z = (emg_filt - emg_filt.mean(axis=0)) / emg_filt.std(axis=0, ddof=1)
 
 
322
 
323
  # segment
324
  segs, lbls, reps = process_emg_features(
325
- emg_z, label, rerep, window_size, stride
326
  )
327
 
328
  # split by repetition index
 
121
  return new_data, new_labels
122
 
123
 
124
+ def resample_to_rate(
125
+ emg: np.ndarray,
126
+ label: np.ndarray,
127
+ rerep: np.ndarray,
128
+ original_fs: float,
129
+ target_fs: float,
130
+ ) -> Tuple[np.ndarray, np.ndarray, np.ndarray]:
131
+ """Resamples EMG and aligns sample-wise labels/repetition indices.
132
+
133
+ EMG is resampled with a polyphase anti-imaging filter. Labels and repetition
134
+ indices are mapped with nearest-neighbor sampling so they remain discrete.
135
+
136
+ Args:
137
+ emg (np.ndarray): EMG array of shape (T, D).
138
+ label (np.ndarray): Sample-wise labels of shape (T,).
139
+ rerep (np.ndarray): Sample-wise repetition indices of shape (T,).
140
+ original_fs (float): Original sampling frequency in Hz.
141
+ target_fs (float): Target sampling frequency in Hz.
142
+
143
+ Returns:
144
+ Tuple[np.ndarray, np.ndarray, np.ndarray]: Resampled EMG, labels, and
145
+ repetition indices.
146
+ """
147
+ if target_fs <= 0:
148
+ raise ValueError("target_fs must be positive")
149
+ if np.isclose(original_fs, target_fs):
150
+ return emg, label, rerep
151
+
152
+ from fractions import Fraction
153
+
154
+ ratio = Fraction(target_fs / original_fs).limit_denominator(1000)
155
+ up, down = ratio.numerator, ratio.denominator
156
+ emg_resampled = signal.resample_poly(emg, up, down, axis=0)
157
+
158
+ new_len = emg_resampled.shape[0]
159
+ source_idx = np.rint(
160
+ np.arange(new_len, dtype=np.float64) * original_fs / target_fs
161
+ ).astype(np.int64)
162
+ source_idx = np.clip(source_idx, 0, len(label) - 1)
163
+
164
+ return emg_resampled, label[source_idx], rerep[source_idx]
165
+
166
+
167
  def notch_filter(
168
  data: np.ndarray, notch_freq: float = 50.0, Q: float = 30.0, fs: float = 200.0
169
  ) -> np.ndarray:
 
221
  rerep: np.ndarray,
222
  window_size: int = 1024,
223
  stride: int = 512,
224
+ majority: bool = False,
225
  ) -> Tuple[np.ndarray, np.ndarray, np.ndarray]:
226
  """Segments raw EMG signals into overlapping windows.
227
 
 
240
  """
241
  segs, lbls, reps = [], [], []
242
  N = len(label)
243
+ for start in range(0, N - window_size + 1, stride):
244
  end = start + window_size
245
+ win = emg[start:end]
246
+ window_labels = label[start:end]
247
+ window_reps = rerep[start:end]
248
+
249
+ if majority:
250
+ # Majority voting over gestures
251
+ labels_u, labels_c = np.unique(window_labels, return_counts=True)
252
+ assigned_label = int(labels_u[np.argmax(labels_c)])
253
  else:
254
+ # Pick index 0 for labe
255
+ assigned_label = window_labels[0]
256
+ assigned_rep = window_reps[0]
257
 
258
  segs.append(win)
259
+ lbls.append(assigned_label)
260
+ reps.append(assigned_rep)
261
  return np.array(segs), np.array(lbls), np.array(reps)
262
 
263
 
 
292
  default=3,
293
  help="Number of augmented versions to create for each training sample.",
294
  )
295
+ args.add_argument(
296
+ "--resample-2khz",
297
+ action="store_true",
298
+ help=(
299
+ "If set, resample EMG from 200 Hz to 2000 Hz before segmentation. "
300
+ "Labels and repetition indices are aligned with nearest-neighbor mapping."
301
+ ),
302
+ )
303
+ args.add_argument(
304
+ "--majority",
305
+ action="store_true",
306
+ )
307
  args = args.parse_args()
308
 
309
  data_dir = args.data_dir
 
323
  os.system(f"rm {data_dir}/s{i}.zip")
324
  print(f"Downloaded and unzipped subject {i}\n{data_dir}/s{i}.zip")
325
 
326
+ original_fs = 200.0
327
+ target_fs = 2000.0 if args.resample_2khz else original_fs
328
  window_size, stride = args.seq_len, args.stride
329
 
330
+ window_seconds = sequence_to_seconds(window_size, target_fs)
331
+ print(
332
+ f"Sampling rate: {target_fs:.0f} Hz"
333
+ + (f" (resampled from {original_fs:.0f} Hz)" if args.resample_2khz else "")
334
+ )
335
  print(f"Window size: {window_size} samples ({window_seconds:.2f} seconds)")
336
+ print(f"{args.majority=}")
337
 
338
  train_reps = [1, 3, 4, 6]
339
  val_reps = [2]
 
381
  elif "E3" in mat:
382
  label = np.where(label != 0, label + 29, 0)
383
 
384
+ # Filter at the original acquisition rate. Upsampling afterward does not
385
+ # create new frequency content, but provides a 2 kHz sample grid when requested.
386
+ emg_filt = bandpass_filter_emg(emg, 20, 90, fs=original_fs)
387
+ emg_filt = notch_filter(emg_filt, 50, 30, fs=original_fs)
388
+
389
+ if args.resample_2khz:
390
+ emg_filt, label, rerep = resample_to_rate(
391
+ emg_filt, label, rerep, original_fs, target_fs
392
+ )
393
 
394
  # z-score
395
+ channel_std = emg_filt.std(axis=0, ddof=1)
396
+ channel_std[channel_std == 0] = 1.0
397
+ emg_z = (emg_filt - emg_filt.mean(axis=0)) / channel_std
398
 
399
  # segment
400
  segs, lbls, reps = process_emg_features(
401
+ emg_z, label, rerep, window_size, stride, args.majority
402
  )
403
 
404
  # split by repetition index