Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 11 additions & 0 deletions apps/predbat/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -166,6 +166,17 @@
"icon": "mdi:multiplication",
"default": 1.2,
},
{
"name": "ml_retrain_interval_hours",
"friendly_name": "ML Retrain Interval",
"type": "input_number",
"min": 1,
"max": 48,
"step": 1,
"unit": "hours",
"icon": "mdi:clock-end",
"default": 2,
},
{
"name": "battery_rate_max_scaling",
"friendly_name": "Battery rate max scaling charge",
Expand Down
13 changes: 7 additions & 6 deletions apps/predbat/load_ml_component.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@

ComponentBase wrapper that manages the LoadPredictor lifecycle including
data fetching, periodic training/fine-tuning, prediction generation, and
status sensor publishing. Retrains every 2 hours and updates predictions
status sensor publishing. Retrains every 2 hours by default and updates predictions
every 30 minutes.
"""

Expand All @@ -30,8 +30,7 @@
import traceback
import numpy as np

# Training intervals
RETRAIN_INTERVAL_SECONDS = 2 * 60 * 60 # 2 hours between training cycles
# Prediction interval
PREDICTION_INTERVAL_SECONDS = 30 * 60 # 30 minutes between predictions

# Database schema version - increment when the saved format changes to force a clean rebuild
Expand Down Expand Up @@ -721,8 +720,10 @@ async def run(self, seconds, first):
is_initial = not self.initial_training_done

# Retrain if the model is older than the retrain interval (rather than on a fixed tick)
retrain_age_seconds = (self.now_utc - self.last_train_time).total_seconds() if self.last_train_time else RETRAIN_INTERVAL_SECONDS
should_train = not first and (retrain_age_seconds >= RETRAIN_INTERVAL_SECONDS)
retrain_interval_hours = self.get_arg("ml_retrain_interval_hours", 2)
retrain_interval_seconds = retrain_interval_hours * 60 * 60
retrain_age_seconds = (self.now_utc - self.last_train_time).total_seconds() if self.last_train_time else retrain_interval_seconds
should_train = not first and (retrain_age_seconds >= retrain_interval_seconds)

# Fetch fresh load data periodically (every N minutes)
should_fetch = first or should_train or ((seconds % PREDICTION_INTERVAL_SECONDS) == 0)
Expand Down Expand Up @@ -756,7 +757,7 @@ async def run(self, seconds, first):
self.log("ML Component: Initial training is required, delaying until component has started")
return True
elif should_train:
self.log("ML Component: Starting fine-tune training (2h interval), model age is {} hours".format(retrain_age_seconds / 3600.0))
self.log("ML Component: Starting fine-tune training ({}h interval), model age is {} hours".format(retrain_interval_hours, retrain_age_seconds / 3600.0))
elif should_fetch:
# If not training either then no need to print anything
self.log("ML Component: No training needed, model age is {} hours".format(dp2(retrain_age_seconds / 3600.0)))
Expand Down
85 changes: 83 additions & 2 deletions apps/predbat/tests/test_load_ml.py
Original file line number Diff line number Diff line change
Expand Up @@ -64,6 +64,7 @@ def test_load_ml(my_predbat=None):
("component_stale_midnight_baseline", _test_component_stale_midnight_baseline, "LoadMLComponent baseline handling when publishing crosses midnight"),
("car_subtraction_direct", _test_car_subtraction_direct, "Direct car_subtraction method with interpolation and smoothing"),
("component_run_data_merge", _test_component_run_data_merge, "LoadMLComponent run() data fetch, save and merge across two runs"),
("component_retrain_interval", _test_component_retrain_interval, "Configurable ML retraining interval and live changes"),
("component_init_predictor_last_train_time", _test_component_init_predictor_sets_last_train_time, "LoadMLComponent _init_predictor sets last_train_time from embedded training_timestamp"),
("nan_inf_robustness", _test_nan_inf_robustness, "LoadPredictor handles NaN, Inf, and None across input channels without NaN loss"),
("database_zero_preservation", _test_database_zero_preservation, "Database save/load preserves 0.0 values across roundtrip"),
Expand Down Expand Up @@ -981,6 +982,86 @@ def _test_prediction_with_temp():
assert max_minute >= 2800, f"Predictions should span ~48h (2880 min), got {max_minute} min"


def _test_component_retrain_interval():
"""Check retraining boundaries, live configuration, startup and prediction cadence."""
import asyncio
from types import SimpleNamespace
from unittest.mock import AsyncMock, Mock
from config import CONFIG_ITEMS
from load_ml_component import LoadMLComponent, PREDICTION_INTERVAL_SECONDS

async def run_test():
"""Exercise the real run loop without performing expensive ML work."""
settings = {}
base = SimpleNamespace(
log=Mock(),
local_tz=timezone.utc,
prefix="predbat",
args={},
config_root=None,
now_utc=datetime(2026, 1, 3, 12, tzinfo=timezone.utc),
midnight_utc=datetime(2026, 1, 3, tzinfo=timezone.utc),
prediction_started=False,
get_arg=lambda key, default=None, **kwargs: settings.get(key, default),
)
# Start disabled to avoid requiring real load sensors in this scheduling test.
component = LoadMLComponent(base, load_ml_enable=False, load_ml_database_days=0)
component.ml_enable = True
component.data_ready = True
component.load_data_age_days = 7
component.initial_training_done = True
component._do_fetch = AsyncMock()
component._do_training = AsyncMock()
component._update_model_status = Mock()
component._get_predictions = Mock()
component._publish_entity = Mock()

# A missing setting retains the two-hour default; longer intervals must not
# train at the old threshold. Both sides of each boundary catch unit mistakes.
for interval in (None, 1, 24, 48):
if interval is not None:
settings["ml_retrain_interval_hours"] = interval
hours = 2 if interval is None else interval
for age, expected in ((hours * 3600 - 1, 0), (hours * 3600, 1)):
component.last_train_time = base.now_utc - timedelta(seconds=age)
component._do_training.reset_mock()
assert await component.run(seconds=30, first=False)
assert component._do_training.await_count == expected, f"Interval {interval}, age {age}: unexpected training"
if expected:
assert any(f"({hours}h interval)" in call.args[0] for call in base.log.call_args_list), "Training log must show the configured interval"

# Changing the setting on the same component takes effect without a restart.
settings["ml_retrain_interval_hours"] = 24
component.last_train_time = base.now_utc - timedelta(hours=3)
component._do_training.reset_mock()
component._get_predictions.reset_mock()
assert await component.run(seconds=PREDICTION_INTERVAL_SECONDS, first=False)
component._do_training.assert_not_awaited()
component._get_predictions.assert_called_once()
settings["ml_retrain_interval_hours"] = 2
assert await component.run(seconds=30, first=False)
component._do_training.assert_awaited_once_with(False)

# A long configured interval must not delay the first-ever training, but
# the startup pass must still defer it until the component has started.
settings["ml_retrain_interval_hours"] = 24
component.last_train_time = None
component.initial_training_done = False
component._do_training.reset_mock()
assert await component.run(seconds=0, first=True)
component._do_training.assert_not_awaited()
assert await component.run(seconds=30, first=False)
component._do_training.assert_awaited_once_with(True)

item = next(item for item in CONFIG_ITEMS if item["name"] == "ml_retrain_interval_hours")
assert item["type"] == "input_number"
assert item["unit"] == "hours"
assert (item["default"], item["min"], item["max"], item["step"]) == (2, 1, 48, 1)
assert item["max"] <= component.ml_max_model_age_hours, "Interval must not exceed the model staleness limit"

asyncio.run(run_test())


def _test_component_run_data_merge():
"""Test LoadMLComponent.run() - mocks _fetch_load_data and verifies:
1. Run 1 (first=True): data is fetched and stored, but training is deferred (no save).
Expand Down Expand Up @@ -1104,7 +1185,7 @@ async def mock_save_database_history():
mock_base.now_utc = mock_base.now_utc + timedelta(minutes=ELAPSED_MINUTES)

# ── Run 2 (first=False, seconds=30) ─────────────────────────────────────
# last_train_time is None -> retrain_age_seconds = RETRAIN_INTERVAL_SECONDS -> should_train=True
# last_train_time is None -> retrain_age_seconds = retrain_interval_seconds -> should_train=True
# Expect: shift old keys, merge fresh data, run initial training, save once.
component._fetch_load_data = AsyncMock(return_value=(fetch_data_2, 7, 3.0, None, None, None, None))

Expand Down Expand Up @@ -1133,7 +1214,7 @@ async def mock_save_database_history():

# ── Run 3 (first=False, seconds=PREDICTION_INTERVAL_SECONDS) ─────────────
# last_train_time = 30 min ago (set in mock_do_training to component.now_utc of Run 2).
# retrain_age_seconds = 30*60 = 1800 < RETRAIN_INTERVAL_SECONDS (7200) -> should_train=False.
# retrain_age_seconds = 30*60 = 1800 < retrain_interval_seconds (7200 by default) -> should_train=False.
# seconds % PREDICTION_INTERVAL_SECONDS == 0 -> should_fetch=True.
# Expect: fetch+predict+save only, no training.
component._fetch_load_data = AsyncMock(return_value=(fetch_data_3, 7, 3.0, None, None, None, None))
Expand Down
2 changes: 1 addition & 1 deletion docs/components.md
Original file line number Diff line number Diff line change
Expand Up @@ -1689,7 +1689,7 @@ This provides more accurate load predictions than simple averaging, especially f
- Optionally incorporates PV generation and temperature forecast data
- Trains a multi-layer neural network on your historical patterns
- Makes autoregressive predictions for 48 hours ahead in 5-minute intervals
- Fine-tunes periodically (every 2 hours) to adapt to changing patterns
- Fine-tunes periodically (every 2 hours by default) to adapt to changing patterns; adjust the [ML Retrain Interval](load-ml.md#retraining-interval) config entity to change this
- Validates predictions and falls back gracefully if accuracy is poor
- Publishes predictions to `sensor.predbat_load_ml_forecast`

Expand Down
14 changes: 12 additions & 2 deletions docs/load-ml.md
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,7 @@ The ML Load Prediction component uses a lightweight multi-layer perceptron (MLP)
- Deep neural network with 3 hidden layers [512, 256, 64 neurons]
- Optimised with He initialisation and AdamW weight decay for robust training
- Automatically trains on historical data (requires at least 1 day, recommended 7+ days; fetches up to `load_ml_max_days_history` days from HA and accumulates up to `load_ml_database_days` days in the on-disk database)
- Fine-tunes periodically (every 2 hours) using full dataset to adapt to changing patterns
- Fine-tunes periodically (every 2 hours by default, configurable with `ml_retrain_interval_hours`) using full dataset to adapt to changing patterns
- Time-weighted training prioritizes recent data while learning from historical patterns
- Model persists across restarts
- Falls back gracefully if predictions are unreliable
Expand Down Expand Up @@ -228,6 +228,16 @@ predbat:
- **When to decrease**: To save disk space, or if you prefer the model to forget older patterns faster
- **Disk usage**: Each day of history uses approximately 5.6 KB (5 channels × 288 steps/day × 4 bytes, plus minimal metadata/format overhead)

### Retraining Interval

Set **ML Retrain Interval** (`input_number.predbat_ml_retrain_interval_hours`) in Home Assistant or `ml_retrain_interval_hours` on Predbat's web configuration page.

- **Default**: 2 hours
- **Range**: 1–48 hours, in 1-hour steps. The maximum matches the model's existing 48-hour staleness limit.
- **Example**: Set to 24 for daily retraining to reduce how often CPU-intensive training runs.
- Changes take effect on the next ML component cycle without restarting. The interval is measured from the last training timestamp, not a fixed time of day.
- Initial training still runs as soon as the component has started and enough data is available. Predictions continue to update every 30 minutes between training sessions.

### History Accumulation and the Database

The ML component maintains two distinct layers of historical data:
Expand Down Expand Up @@ -383,7 +393,7 @@ Good predictions require:
4. **Energy Rate Data**: Automatically included - helps model learn consumption patterns based on time-of-use tariffs
5. **PV Generation Data**: If you have solar panels, include `pv_today` sensor for better correlation
6. **Clean Data**: Avoid gaps or incorrect readings in historical data
7. **Recent Training**: Model retrains every 2 hours using full dataset with time-weighted sampling to adapt to changing patterns
7. **Recent Training**: Model retrains every `ml_retrain_interval_hours` hours (default 2) using full dataset with time-weighted sampling to adapt to changing patterns

### Understanding MAE (Mean Absolute Error)

Expand Down