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#

ChannelOcclusionWrapper

Wraps a model to apply channel occlusion (zeroing) at input level.

Functions#

build_dataset(dataset_name, channels, pipeline, seq_len)

Build dataset, handling MASS multi-cohort case.

compute_metrics(all_proba, all_targets[, ignore_index])

Compute metrics from aggregated probabilities and targets.

evaluate_subject(model, inputs, L, device)

Sliding-window voting evaluation for a single subject.

load_model(model_dir, device)

Load model from config.json + model.pt.

main()

Module Contents#

class protosleepnet.baselines.test_occlusion.ChannelOcclusionWrapper(model, mode=None, p=0.0, channels_to_occlude=None)#

Bases: torch.nn.Module

Wraps 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#