|
1 | | -# %% |
2 | | -from tensorflow.keras import layers, models |
3 | | -from tqdm import tqdm |
| 1 | +import warnings |
4 | 2 |
|
5 | 3 | from sleepecg import ( |
6 | 4 | evaluate, |
|
13 | 11 | save_classifier, |
14 | 12 | set_nsrr_token, |
15 | 13 | ) |
| 14 | +from tensorflow.keras import layers, models |
| 15 | +from tqdm import tqdm |
16 | 16 |
|
17 | | -# %% Read data and extract features |
18 | 17 | set_nsrr_token("your-token-here") |
19 | | -records = list(read_mesa()) |
20 | | - |
21 | | -feature_extraction_params = { |
22 | | - "lookback": 120, |
23 | | - "lookforward": 150, |
24 | | - "feature_selection": [ |
25 | | - "hrv-time", |
26 | | - "hrv-frequency", |
27 | | - "recording_start_time", |
28 | | - "age", |
29 | | - "gender", |
30 | | - ], |
31 | | - "min_rri": 0.3, |
32 | | - "max_rri": 2, |
33 | | - "max_nans": 0.5, |
34 | | -} |
35 | | - |
36 | | -features_train, stages_train, feature_ids = extract_features( |
37 | | - tqdm(records), |
38 | | - **feature_extraction_params, |
39 | | - n_jobs=-2, |
40 | | -) |
41 | 18 |
|
42 | | -# %% Merge sleep stages, pad and mask data as preparation for keras NN |
43 | | -stages_mode = "wake-rem-nrem" |
| 19 | +TRAIN = True # set to False to skip training and load classifier from disk |
44 | 20 |
|
45 | | -features_train_pad, stages_train_pad, _ = prepare_data_keras( |
46 | | - features_train, |
47 | | - stages_train, |
48 | | - stages_mode, |
49 | | -) |
50 | | -print_class_balance(stages_train_pad, stages_mode) |
51 | | - |
52 | | -# %% Define and train model |
53 | | -model = models.Sequential( |
54 | | - [ |
55 | | - layers.Input((None, features_train_pad.shape[2])), |
56 | | - layers.Masking(-1), |
57 | | - layers.BatchNormalization(), |
58 | | - layers.Dense(64), |
59 | | - layers.ReLU(), |
60 | | - layers.Bidirectional(layers.GRU(8, return_sequences=True)), |
61 | | - layers.Bidirectional(layers.GRU(8, return_sequences=True)), |
62 | | - layers.Dense(stages_train_pad.shape[-1], activation="softmax"), |
63 | | - ] |
| 21 | +# silence warnings (which might pop up during feature extraction) |
| 22 | +warnings.filterwarnings( |
| 23 | + "ignore", category=RuntimeWarning, message="HR analysis window too short" |
64 | 24 | ) |
65 | 25 |
|
66 | | -model.compile( |
67 | | - optimizer="rmsprop", |
68 | | - loss="categorical_crossentropy", |
69 | | - metrics=["accuracy"], |
70 | | -) |
71 | | -model.build() |
72 | | -model.summary() |
73 | | - |
74 | | -# %% Train model |
75 | | -model.fit( |
76 | | - features_train_pad, |
77 | | - stages_train_pad, |
78 | | - epochs=25, |
79 | | -) |
| 26 | +if TRAIN: |
| 27 | + print("‣ Starting training...") |
| 28 | + print("‣‣ Extracting features...") |
| 29 | + records = list(read_mesa(offline=False)) |
80 | 30 |
|
81 | | -# %% Store classifier |
82 | | -save_classifier( |
83 | | - name="wrn-gru-mesa", |
84 | | - model=model, |
85 | | - stages_mode=stages_mode, |
86 | | - feature_extraction_params=feature_extraction_params, |
87 | | - mask_value=-1, |
88 | | - classifiers_dir="./classifiers", |
89 | | -) |
| 31 | + feature_extraction_params = { |
| 32 | + "lookback": 120, |
| 33 | + "lookforward": 150, |
| 34 | + "feature_selection": [ |
| 35 | + "hrv-time", |
| 36 | + "hrv-frequency", |
| 37 | + "recording_start_time", |
| 38 | + "age", |
| 39 | + "gender", |
| 40 | + ], |
| 41 | + "min_rri": 0.3, |
| 42 | + "max_rri": 2, |
| 43 | + "max_nans": 0.5, |
| 44 | + } |
| 45 | + |
| 46 | + features_train, stages_train, feature_ids = extract_features( |
| 47 | + tqdm(records), |
| 48 | + **feature_extraction_params, |
| 49 | + n_jobs=-1, |
| 50 | + ) |
| 51 | + |
| 52 | + print("‣‣ Preparing data for Keras...") |
| 53 | + stages_mode = "wake-rem-nrem" |
| 54 | + |
| 55 | + features_train_pad, stages_train_pad, _ = prepare_data_keras( |
| 56 | + features_train, |
| 57 | + stages_train, |
| 58 | + stages_mode, |
| 59 | + ) |
| 60 | + print_class_balance(stages_train_pad, stages_mode) |
| 61 | + |
| 62 | + print("‣‣ Defining model...") |
| 63 | + model = models.Sequential( |
| 64 | + [ |
| 65 | + layers.Input((None, features_train_pad.shape[2])), |
| 66 | + layers.Masking(-1), |
| 67 | + layers.BatchNormalization(), |
| 68 | + layers.Dense(64), |
| 69 | + layers.ReLU(), |
| 70 | + layers.Bidirectional(layers.GRU(8, return_sequences=True)), |
| 71 | + layers.Bidirectional(layers.GRU(8, return_sequences=True)), |
| 72 | + layers.Dense(stages_train_pad.shape[-1], activation="softmax"), |
| 73 | + ] |
| 74 | + ) |
| 75 | + |
| 76 | + model.compile( |
| 77 | + optimizer="rmsprop", |
| 78 | + loss="categorical_crossentropy", |
| 79 | + metrics=["accuracy"], |
| 80 | + ) |
| 81 | + model.build() |
| 82 | + model.summary() |
| 83 | + |
| 84 | + print("‣‣ Training model...") |
| 85 | + model.fit( |
| 86 | + features_train_pad, |
| 87 | + stages_train_pad, |
| 88 | + epochs=25, |
| 89 | + ) |
| 90 | + |
| 91 | + print("‣‣ Saving classifier...") |
| 92 | + save_classifier( |
| 93 | + name="wrn-gru-mesa", |
| 94 | + model=model, |
| 95 | + stages_mode=stages_mode, |
| 96 | + feature_extraction_params=feature_extraction_params, |
| 97 | + mask_value=-1, |
| 98 | + classifiers_dir="./classifiers", |
| 99 | + ) |
90 | 100 |
|
91 | | -# %% Load classifier from disk for validation |
| 101 | +print("‣ Starting testing...") |
| 102 | +print("‣‣ Loading classifier...") |
92 | 103 | clf = load_classifier("wrn-gru-mesa", "./classifiers") |
93 | 104 |
|
94 | | -# %% Read data and extract features |
95 | | -shhs = list(read_shhs()) |
| 105 | +print("‣‣ Extracting features...") |
| 106 | +shhs = list(read_shhs(offline=False)) |
96 | 107 |
|
97 | 108 | features_test, stages_test, feature_ids = extract_features( |
98 | 109 | tqdm(shhs), |
99 | 110 | **clf.feature_extraction_params, |
100 | 111 | n_jobs=-2, |
101 | 112 | ) |
102 | 113 |
|
103 | | -# %% Predict & evaluate |
| 114 | +print("‣‣ Evaluating classifier...") |
104 | 115 | features_test_pad, stages_test_pad, _ = prepare_data_keras( |
105 | 116 | features_test, |
106 | 117 | stages_test, |
|
0 commit comments