51 lines
2.3 KiB
Python
51 lines
2.3 KiB
Python
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],
|
|
)
|