diff --git a/sleep_utils/usleep_utils.py b/sleep_utils/usleep_utils.py index 3c09677..ac2bfff 100644 --- a/sleep_utils/usleep_utils.py +++ b/sleep_utils/usleep_utils.py @@ -13,18 +13,9 @@ import numpy as np import warnings import pandas as pd +import requests +from io import BytesIO -def tempfile_wrapper(func): - @wraps(func) - def wrapped(*args, **kwargs): - try: - tempfile_name = tempfile.NamedTemporaryFile().name + '.edf' - res = func(*args, **kwargs, tmp_edf=tempfile_name) - finally: - if os.path.isfile(tempfile_name): - os.remove(tempfile_name) - return res - return wrapped def disable_ssl_verify(): """will monkey-patch requests made by usleep-api to veryify=False""" @@ -43,10 +34,65 @@ def patched_request(self, endpoint, method, as_json=False, USleepAPI._request = patched_request warnings.warn('patched SSL to accept insecure connections') -@tempfile_wrapper -def predict_usleep_raw(raw, api_token, eeg_chs=None, eog_chs=None, - ch_groups=None, model='U-Sleep v2.0', saveto=None, - seconds_per_label=30, tmp_edf=None, return_proba=False): +def score_sleep(raw = None, + edf_file = None, + api_token = None, + backend = 'sleepyland', + backend_url = None, + eeg_chs=None, + eog_chs=None, + ch_groups=None, + model=None, + saveto=None, + tmp_edf=None, + seconds_per_label=30, + return_proba=False): + + if (raw is None) == (edf_file is None): + raise ValueError('either raw or edf_file has to be provided, not both or neither') + + elif raw is not None: + + return _score_sleep_raw(raw, + api_token = api_token, + backend = backend, + backend_url = backend_url, + eeg_chs = eeg_chs, + eog_chs = eog_chs, + ch_groups = ch_groups, + model = model, + saveto = saveto, + seconds_per_label = seconds_per_label, + return_proba = return_proba) + elif edf_file is not None: + + return _score_sleep_file(edf_file, + api_token = api_token, + backend = backend, + backend_url = backend_url, + eeg_chs = eeg_chs, + eog_chs = eog_chs, + ch_groups = ch_groups, + model = model, + saveto = saveto, + seconds_per_label = seconds_per_label, + return_proba = return_proba) + + +#@tempfile_wrapper +def _score_sleep_raw(raw, + api_token = None, + backend = 'sleepyland', + backend_url = None, + eeg_chs=None, + eog_chs=None, + ch_groups=None, + model='U-Sleep v2.0', + saveto=None, + seconds_per_label=30, + tmp_edf=None, + return_proba=False): + """ Run U-Sleep prediction on an mne.io.Raw object. @@ -60,6 +106,10 @@ def predict_usleep_raw(raw, api_token, eeg_chs=None, eog_chs=None, The raw EEG recording. api_token : str U-Sleep API token (https://sleep.ai.ku.dk). + backend : str + Which backend should be used for scoring. + backend_url : str + URL for different backends. eeg_chs : list of str, optional EEG channel names for prediction. eog_chs : list of str, optional @@ -88,22 +138,37 @@ def predict_usleep_raw(raw, api_token, eeg_chs=None, eog_chs=None, raw = raw.copy() # work on copy as we resample data etc. # convert to EDF file if not print('converting file to EDF') - chs = list(set(eeg_chs + eog_chs)) + if eeg_chs and eog_chs: + chs = list(set(eeg_chs + eog_chs)) + elif ch_groups: + chs = {channel for channel_value in ch_groups for channel in channel_value} # chs_idx = [i for i, ch in enumerate(raw.ch_names) if ch in chs] # only keep channels that are actually requested - if any([ch not in raw.ch_names for ch in chs]): - raw.drop_channels([ch for ch in raw.ch_names if not ch in chs]) +# if any([ch not in raw.ch_names for ch in chs]): +# raw.drop_channels([ch for ch in raw.ch_names if not ch in chs]) # is resampled anyway internally, reduce data size if raw.info['sfreq']>128: print('downsampling to 128 hz') raw.resample(128, n_jobs=-2) - mne.export.export_raw(tmp_edf, raw, fmt='edf', overwrite=True) - return predict_usleep(tmp_edf, api_token, eeg_chs=eeg_chs, eog_chs=eog_chs, - ch_groups=None, model=model, saveto=saveto, - seconds_per_label=seconds_per_label, - return_proba=return_proba) + try: + tmp_edf = tempfile.NamedTemporaryFile().name + '.edf' + + mne.export.export_raw(tmp_edf, raw, fmt='edf', overwrite=True) + return _score_sleep_file(tmp_edf, + api_token = api_token, + backend = backend, + backend_url = backend_url, + eeg_chs=eeg_chs, + eog_chs=eog_chs, + ch_groups=ch_groups, + model=model, + saveto=saveto, + seconds_per_label=seconds_per_label, + return_proba=return_proba) + finally: + os.remove(tmp_edf) def delete_all_sessions(api_token): """convenience function to delete all sessions and data""" @@ -111,9 +176,17 @@ def delete_all_sessions(api_token): api = USleepAPI(api_token=api_token) api.delete_all_sessions() -def predict_usleep(edf_file, api_token, eeg_chs=None, eog_chs=None, - ch_groups=None, model='U-Sleep v2.0', saveto=None, - seconds_per_label=30, return_proba=False): +def _score_sleep_file(edf_file, + api_token = None, + backend = 'sleepyland', + backend_url=None, + eeg_chs=None, + eog_chs=None, + ch_groups=None, + model='U-Sleep v2.0', + saveto=None, + seconds_per_label=30, + return_proba=False): """ Run U-Sleep prediction on an EDF file via the U-Sleep API. @@ -124,6 +197,10 @@ class probabilities. ---------- edf_file : str Path to a local EDF file. + backend : str + Which backend should be used for scoring. + backend_url : str + URL for different backends. api_token : str U-Sleep API token (https://sleep.ai.ku.dk). eeg_chs : list of str, optional @@ -149,18 +226,79 @@ class probabilities. Label probabilities (if return_proba is True). """ from sleep_utils import write_hypno - try: - from usleep_api import USleepAPI - except ModuleNotFoundError as e: - raise(ModuleNotFoundError(f"{e}\n If missing, please install via 'pip install usleep_api --no-deps'")) # parameter checks +# from tools import write_hypno - if len(eeg_chs)==0 or len(eog_chs)==0: - raise ValueError('One element missing: {len(eeg_chs)=}, {len(eog_chs)=}') +# if len(eeg_chs)==0 or len(eog_chs)==0: +# raise ValueError('One element missing: {len(eeg_chs)=}, {len(eog_chs)=}') assert 0