fixes
This commit is contained in:
@@ -0,0 +1,50 @@
|
||||
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],
|
||||
)
|
||||
Reference in New Issue
Block a user