Files
temperature-based-fertility…/code/utils/smoothing.py
T
Alex Blank c6defa2065 fixes
2025-05-19 13:59:16 +02:00

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],
)