import numpy as np from scipy.signal import butter, filtfilt from statsmodels.tsa.stl._stl import STL def highpass_filter(data, cutoff_freq, fs=288): nyquist = 0.5 * fs normal_cutoff = cutoff_freq / nyquist b, a = butter(N=3, Wn=normal_cutoff, btype="high", analog=False) return filtfilt(b, a, data) def mirror_extend(series, extend_len): """Mirrors the beginning and end of the time series to stabilize smoothing.""" # Mirror extension start_extension = series[:extend_len][::-1] # Reverse first part end_extension = series[-extend_len:][::-1] # Reverse last part extended_series = np.concatenate([start_extension, series, end_extension]) return extended_series def get_trend(input_curve: np.ndarray | list, measurements_per_day: int = 288) -> np.ndarray: extension_len = 3 extended_input_curve = mirror_extend(input_curve, extension_len * measurements_per_day) stl = STL(extended_input_curve, period=measurements_per_day, robust=False, trend=measurements_per_day * 14 + 1) trend = stl.fit().trend return trend[extension_len * measurements_per_day:-extension_len * measurements_per_day] def get_curve_composition(input_curve: np.ndarray | list, measurements_per_day: int = 288) -> tuple: """ Decomposes the input curve into trend, seasonal, residual and smoothed components. :param input_curve: raw input curve :param measurements_per_day: seasonal period, here: measurements per day -> 288 :return: composition of curve as tuple (trend, seasonal, residual, smoothed) """ extension_len = 3 extended_input_curve = mirror_extend(input_curve, extension_len * measurements_per_day) stl_results = STL(extended_input_curve, period=measurements_per_day, robust=False).fit() long_term = (extended_input_curve - stl_results.seasonal) wiggles = highpass_filter(long_term, 0.1) smoothed = long_term - wiggles return ( stl_results.trend[extension_len * measurements_per_day:-extension_len * measurements_per_day], stl_results.seasonal[extension_len * measurements_per_day:-extension_len * measurements_per_day], stl_results.resid[extension_len * measurements_per_day:-extension_len * measurements_per_day], smoothed[extension_len * measurements_per_day:-extension_len * measurements_per_day], )