protosleepnet.baselines.test_occlusion#
Channel occlusion robustness test for baseline 3ch models.
Tests how the original Phan models (SleepTransformer, SeqSleepNet) with 3 channels handle missing/occluded channels at inference time.
Evaluates on the test split with 5 scenarios (clean + 4 occlusion types). Supports SHHS (SleepTransformer) and MASS (SeqSleepNet) datasets.
- Usage:
# SleepTransformer on SHHS python -m protosleepnet.baselines.test_occlusion –model_dir /path/to/sleeptransformer-phan-3ch –gpu_id 0
# SeqSleepNet on MASS python -m protosleepnet.baselines.test_occlusion –model_dir /path/to/seqsleepnet-phan-3ch –dataset mass –seq_len 20 –gpu_id 0
Attributes#
Classes#
Wraps a model to apply channel occlusion (zeroing) at input level. |
Functions#
|
Build dataset, handling MASS multi-cohort case. |
|
Compute metrics from aggregated probabilities and targets. |
|
Sliding-window voting evaluation for a single subject. |
|
Load model from config.json + model.pt. |
|
Module Contents#
- class protosleepnet.baselines.test_occlusion.ChannelOcclusionWrapper(model, mode=None, p=0.0, channels_to_occlude=None)#
Bases:
torch.nn.ModuleWraps a model to apply channel occlusion (zeroing) at input level.
- Modes:
None: no occlusion (passthrough) “random”: zero each channel independently with probability p per epoch
(at least 1 channel kept per sample per epoch)
“fixed”: zero specific channels for the entire night
- forward(x)#
- channels_to_occlude = []#
- mode = None#
- model#
- p = 0.0#
- protosleepnet.baselines.test_occlusion.build_dataset(dataset_name, channels, pipeline, seq_len)#
Build dataset, handling MASS multi-cohort case.
- protosleepnet.baselines.test_occlusion.compute_metrics(all_proba, all_targets, ignore_index=-1)#
Compute metrics from aggregated probabilities and targets.
- protosleepnet.baselines.test_occlusion.evaluate_subject(model, inputs, L, device)#
Sliding-window voting evaluation for a single subject.
- protosleepnet.baselines.test_occlusion.load_model(model_dir, device)#
Load model from config.json + model.pt.
- protosleepnet.baselines.test_occlusion.main()#
- protosleepnet.baselines.test_occlusion.CHANNELS = ['EEG', 'EOG', 'EMG']#
- protosleepnet.baselines.test_occlusion.CLASS_NAMES = ['W', 'N1', 'N2', 'N3', 'REM']#
- protosleepnet.baselines.test_occlusion.PIPELINE = 'seqsleepnet'#
- protosleepnet.baselines.test_occlusion.SCENARIOS#