sensor-modeling 0.2.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (114) hide show
  1. sensor_modeling/__init__.py +45 -0
  2. sensor_modeling/alerts/__init__.py +26 -0
  3. sensor_modeling/alerts/alert.py +532 -0
  4. sensor_modeling/analysis/__init__.py +43 -0
  5. sensor_modeling/analysis/_frame.py +19 -0
  6. sensor_modeling/analysis/behavioral_analysis.py +57 -0
  7. sensor_modeling/analysis/behavioral_metrics.py +66 -0
  8. sensor_modeling/analysis/comparison.py +164 -0
  9. sensor_modeling/analysis/dependency_network.py +408 -0
  10. sensor_modeling/analysis/granger_causality.py +314 -0
  11. sensor_modeling/analysis/pipeline.py +168 -0
  12. sensor_modeling/analysis/reporting.py +109 -0
  13. sensor_modeling/baseline/__init__.py +30 -0
  14. sensor_modeling/baseline/adaptive.py +520 -0
  15. sensor_modeling/baseline/features.py +224 -0
  16. sensor_modeling/change_point/__init__.py +13 -0
  17. sensor_modeling/change_point/_validation.py +31 -0
  18. sensor_modeling/change_point/adaptive_normalization.py +55 -0
  19. sensor_modeling/change_point/embedding_cpd.py +60 -0
  20. sensor_modeling/change_point/energy_efficient.py +57 -0
  21. sensor_modeling/change_point/genetic_optimization.py +65 -0
  22. sensor_modeling/cli.py +416 -0
  23. sensor_modeling/context/__init__.py +33 -0
  24. sensor_modeling/context/occupancy.py +529 -0
  25. sensor_modeling/data/__init__.py +5 -0
  26. sensor_modeling/data/loaders.py +146 -0
  27. sensor_modeling/data/preprocessing.py +83 -0
  28. sensor_modeling/data/synthetic.py +121 -0
  29. sensor_modeling/data/validation.py +81 -0
  30. sensor_modeling/evaluation/__init__.py +92 -0
  31. sensor_modeling/evaluation/ablation.py +303 -0
  32. sensor_modeling/evaluation/attribution.py +474 -0
  33. sensor_modeling/evaluation/detection.py +297 -0
  34. sensor_modeling/evaluation/metrics.py +541 -0
  35. sensor_modeling/evaluation/provenance.py +309 -0
  36. sensor_modeling/examples/__init__.py +1 -0
  37. sensor_modeling/examples/demos/__init__.py +1 -0
  38. sensor_modeling/examples/demos/ambient_pipeline_demo.py +418 -0
  39. sensor_modeling/examples/demos/bernoulli_ar_demo.py +356 -0
  40. sensor_modeling/examples/demos/cpd_ar_demo.py +25 -0
  41. sensor_modeling/examples/demos/cpd_benchmark.py +42 -0
  42. sensor_modeling/examples/demos/hmm_granger_demo.py +30 -0
  43. sensor_modeling/examples/demos/nhpp_pelt_demo.py +80 -0
  44. sensor_modeling/examples/tutorials/__init__.py +1 -0
  45. sensor_modeling/fusion/__init__.py +46 -0
  46. sensor_modeling/fusion/defaults.py +296 -0
  47. sensor_modeling/fusion/emissions.py +339 -0
  48. sensor_modeling/fusion/estimate.py +375 -0
  49. sensor_modeling/fusion/filter.py +323 -0
  50. sensor_modeling/health/__init__.py +31 -0
  51. sensor_modeling/health/monitor.py +590 -0
  52. sensor_modeling/health/status.py +74 -0
  53. sensor_modeling/hmm/__init__.py +15 -0
  54. sensor_modeling/hmm/adaptive_hmm.py +22 -0
  55. sensor_modeling/hmm/base.py +134 -0
  56. sensor_modeling/hmm/circadian_hmm.py +22 -0
  57. sensor_modeling/hmm/heterogeneous_hmm.py +22 -0
  58. sensor_modeling/hmm/hierarchical_hmm.py +35 -0
  59. sensor_modeling/hmm/scaled_dirichlet_hmm.py +23 -0
  60. sensor_modeling/interop/__init__.py +57 -0
  61. sensor_modeling/interop/fhir.py +418 -0
  62. sensor_modeling/interop/privacy.py +308 -0
  63. sensor_modeling/models/__init__.py +12 -0
  64. sensor_modeling/models/bernoulli_ar/__init__.py +6 -0
  65. sensor_modeling/models/bernoulli_ar/base_model.py +569 -0
  66. sensor_modeling/models/bernoulli_ar/multivariate_model.py +411 -0
  67. sensor_modeling/models/change_point_detection/__init__.py +10 -0
  68. sensor_modeling/models/change_point_detection/deep.py +65 -0
  69. sensor_modeling/models/change_point_detection/pelt.py +159 -0
  70. sensor_modeling/models/nhpp_pelt/__init__.py +5 -0
  71. sensor_modeling/models/nhpp_pelt/bspline.py +96 -0
  72. sensor_modeling/models/nhpp_pelt/cli.py +243 -0
  73. sensor_modeling/models/nhpp_pelt/diagnostics.py +234 -0
  74. sensor_modeling/models/nhpp_pelt/io.py +58 -0
  75. sensor_modeling/models/nhpp_pelt/model.py +408 -0
  76. sensor_modeling/models/nhpp_pelt/optimizer.py +142 -0
  77. sensor_modeling/models/nhpp_pelt/plotting.py +218 -0
  78. sensor_modeling/models/nhpp_pelt/quad.py +72 -0
  79. sensor_modeling/models/nhpp_pelt/regularization.py +121 -0
  80. sensor_modeling/models/nhpp_pelt/utils.py +174 -0
  81. sensor_modeling/observations/__init__.py +59 -0
  82. sensor_modeling/observations/adapters.py +195 -0
  83. sensor_modeling/observations/ingest.py +269 -0
  84. sensor_modeling/observations/observation.py +270 -0
  85. sensor_modeling/observations/registry.py +262 -0
  86. sensor_modeling/observations/stream.py +342 -0
  87. sensor_modeling/observations/types.py +107 -0
  88. sensor_modeling/observations/units.py +117 -0
  89. sensor_modeling/online/__init__.py +36 -0
  90. sensor_modeling/online/benchmarks.py +242 -0
  91. sensor_modeling/online/pipeline.py +485 -0
  92. sensor_modeling/simulation/__init__.py +54 -0
  93. sensor_modeling/simulation/faults.py +191 -0
  94. sensor_modeling/simulation/household.py +862 -0
  95. sensor_modeling/states/__init__.py +23 -0
  96. sensor_modeling/states/markov.py +105 -0
  97. sensor_modeling/states/ontology.py +238 -0
  98. sensor_modeling/utils/__init__.py +41 -0
  99. sensor_modeling/utils/data_io.py +199 -0
  100. sensor_modeling/utils/logging_config.py +10 -0
  101. sensor_modeling/utils/missing.py +188 -0
  102. sensor_modeling/utils/plotting.py +98 -0
  103. sensor_modeling/utils/validation.py +117 -0
  104. sensor_modeling/visualization/__init__.py +3 -0
  105. sensor_modeling/visualization/clinical.py +67 -0
  106. sensor_modeling/visualization/interactive.py +208 -0
  107. sensor_modeling/visualization/research.py +60 -0
  108. sensor_modeling/visualization/web_app.py +137 -0
  109. sensor_modeling-0.2.0.dist-info/METADATA +683 -0
  110. sensor_modeling-0.2.0.dist-info/RECORD +114 -0
  111. sensor_modeling-0.2.0.dist-info/WHEEL +5 -0
  112. sensor_modeling-0.2.0.dist-info/entry_points.txt +18 -0
  113. sensor_modeling-0.2.0.dist-info/licenses/LICENSE +21 -0
  114. sensor_modeling-0.2.0.dist-info/top_level.txt +1 -0
@@ -0,0 +1,224 @@
1
+ """Daily behavioural features derived from the state posterior.
2
+
3
+ A baseline needs a small number of stable, interpretable quantities per day:
4
+ how long the resident appeared to be asleep, how often they were active in the
5
+ kitchen, how much the day was spent in states the model could not name.
6
+
7
+ Two decisions here matter for the honesty of everything downstream.
8
+
9
+ *Aggregate the posterior, not the argmax.* Time in a state is accumulated as
10
+ ``sum P(state | t) * dt`` rather than by counting the intervals where a state
11
+ happened to win. A day made of 60%-confident guesses and a day made of
12
+ 99%-confident conclusions are genuinely different, and collapsing to argmax
13
+ first would erase that difference before the baseline ever sees it.
14
+
15
+ *Record how much of the day was actually observed.* Every summary carries the
16
+ fraction of the day covered by estimates and the mean sensor coverage behind
17
+ them. Days where the apparatus was not working are identifiable as such, so
18
+ the baseline can refuse them instead of recording a broken sensor as a quiet
19
+ day.
20
+ """
21
+
22
+ from __future__ import annotations
23
+
24
+ import logging
25
+ from collections.abc import Iterable, Sequence
26
+ from dataclasses import dataclass
27
+ from datetime import date, datetime, timedelta, tzinfo
28
+
29
+ from ..fusion.estimate import StateEstimate
30
+ from ..states.ontology import BehaviouralState
31
+
32
+ logger = logging.getLogger(__name__)
33
+
34
+ SECONDS_PER_DAY = 86400.0
35
+
36
+
37
+ @dataclass(frozen=True)
38
+ class DailySummary:
39
+ """Aggregated behaviour and data quality for one local calendar day.
40
+
41
+ Attributes
42
+ ----------
43
+ day
44
+ The local calendar date summarised.
45
+ hours
46
+ Expected hours spent in each latent state, from the posterior.
47
+ transitions
48
+ Expected number of state changes across the day.
49
+ coverage
50
+ Mean sensor coverage behind the day's estimates, in ``[0, 1]``.
51
+ abstention
52
+ Fraction of the observed day for which the model declined to name a
53
+ state.
54
+ observed
55
+ Fraction of the day's 24 hours covered by estimates at all.
56
+ """
57
+
58
+ day: date
59
+ hours: dict[BehaviouralState, float]
60
+ transitions: float
61
+ coverage: float
62
+ abstention: float
63
+ observed: float
64
+
65
+ def hours_in(self, state: BehaviouralState) -> float:
66
+ """Return expected hours spent in *state*."""
67
+ return self.hours.get(state, 0.0)
68
+
69
+ def is_usable(self, min_coverage: float = 0.5, min_observed: float = 0.6) -> bool:
70
+ """Whether the day was observed well enough to inform a baseline.
71
+
72
+ A day that fails this test is not a quiet day. It is a day the
73
+ apparatus did not watch, and it must be excluded rather than
74
+ recorded as low activity.
75
+ """
76
+ return self.coverage >= min_coverage and self.observed >= min_observed
77
+
78
+ def to_dict(self) -> dict[str, object]:
79
+ """Return a serialisable form of the summary."""
80
+ return {
81
+ "day": self.day.isoformat(),
82
+ "hours": {state.value: value for state, value in self.hours.items()},
83
+ "transitions": self.transitions,
84
+ "coverage": self.coverage,
85
+ "abstention": self.abstention,
86
+ "observed": self.observed,
87
+ }
88
+
89
+
90
+ def _local_day(moment: datetime, zone: tzinfo | None) -> date:
91
+ """Return the local calendar date of *moment*."""
92
+ return (moment.astimezone(zone) if zone is not None else moment).date()
93
+
94
+
95
+ def _split_by_day(
96
+ start: datetime, end: datetime, zone: tzinfo | None
97
+ ) -> list[tuple[date, float]]:
98
+ """Split an interval into per-local-day durations in seconds.
99
+
100
+ Days are bounded by local midnight, so a day containing a DST transition
101
+ is correctly 23 or 25 hours long rather than assumed to be 24.
102
+ """
103
+ if end <= start:
104
+ return []
105
+ pieces: list[tuple[date, float]] = []
106
+ cursor = start
107
+ while cursor < end:
108
+ local = cursor.astimezone(zone) if zone is not None else cursor
109
+ next_midnight_local = datetime.combine(
110
+ local.date() + timedelta(days=1), datetime.min.time(), tzinfo=local.tzinfo
111
+ )
112
+ boundary = min(next_midnight_local.astimezone(cursor.tzinfo), end)
113
+ if boundary <= cursor: # pragma: no cover - defensive against odd zones
114
+ boundary = end
115
+ pieces.append((local.date(), (boundary - cursor).total_seconds()))
116
+ cursor = boundary
117
+ return pieces
118
+
119
+
120
+ def summarise_days(
121
+ estimates: Sequence[StateEstimate],
122
+ *,
123
+ tz: tzinfo | None = None,
124
+ max_interval: timedelta = timedelta(hours=1),
125
+ ) -> list[DailySummary]:
126
+ """Aggregate a run of state estimates into per-day behavioural summaries.
127
+
128
+ Each estimate is taken to describe the interval since the previous one,
129
+ which is the standard filtering approximation: the posterior at ``t``
130
+ conditions on everything up to ``t``.
131
+
132
+ Parameters
133
+ ----------
134
+ estimates
135
+ State estimates in non-decreasing time order.
136
+ tz
137
+ Timezone whose calendar days are used. Defaults to the timezone of
138
+ the first estimate, so days align with the resident's local clock.
139
+ max_interval
140
+ Longest interval a single estimate may be held to describe. Gaps
141
+ beyond this are left uncounted rather than being attributed to
142
+ whatever state was current before the outage.
143
+
144
+ Returns
145
+ -------
146
+ list[DailySummary]
147
+ One summary per local day that has any observed time, in order.
148
+ """
149
+ if not estimates:
150
+ return []
151
+ ordered = list(estimates)
152
+ for previous, current in zip(ordered, ordered[1:]):
153
+ if current.at < previous.at:
154
+ raise ValueError("estimates must be in non-decreasing time order")
155
+
156
+ zone = tz if tz is not None else ordered[0].at.tzinfo
157
+ states = ordered[0].ontology.states
158
+
159
+ seconds: dict[date, dict[BehaviouralState, float]] = {}
160
+ observed: dict[date, float] = {}
161
+ coverage: dict[date, float] = {}
162
+ abstained: dict[date, float] = {}
163
+ transitions: dict[date, float] = {}
164
+
165
+ for previous, current in zip(ordered, ordered[1:]):
166
+ span = current.at - previous.at
167
+ if span <= timedelta(0) or span > max_interval:
168
+ continue
169
+ for day, duration in _split_by_day(previous.at, current.at, zone):
170
+ bucket = seconds.setdefault(day, dict.fromkeys(states, 0.0))
171
+ for index, state in enumerate(states):
172
+ bucket[state] += float(current.belief[index]) * duration
173
+ observed[day] = observed.get(day, 0.0) + duration
174
+ coverage[day] = coverage.get(day, 0.0) + current.completeness * duration
175
+ if current.abstained:
176
+ abstained[day] = abstained.get(day, 0.0) + duration
177
+
178
+ # Expected number of state changes over the interval, from the
179
+ # probability that the two posteriors disagree.
180
+ change = 1.0 - float((previous.belief * current.belief).sum())
181
+ day = _local_day(current.at, zone)
182
+ transitions[day] = transitions.get(day, 0.0) + change
183
+
184
+ summaries = []
185
+ for day in sorted(observed):
186
+ total = observed[day]
187
+ summaries.append(
188
+ DailySummary(
189
+ day=day,
190
+ hours={state: value / 3600.0 for state, value in seconds[day].items()},
191
+ transitions=transitions.get(day, 0.0),
192
+ coverage=coverage[day] / total if total > 0 else 0.0,
193
+ abstention=abstained.get(day, 0.0) / total if total > 0 else 0.0,
194
+ observed=total / SECONDS_PER_DAY,
195
+ )
196
+ )
197
+ return summaries
198
+
199
+
200
+ def feature_series(
201
+ summaries: Iterable[DailySummary],
202
+ state: BehaviouralState,
203
+ *,
204
+ min_coverage: float = 0.5,
205
+ min_observed: float = 0.6,
206
+ ) -> tuple[list[date], list[float]]:
207
+ """Extract a per-day series of hours in *state*, skipping unusable days.
208
+
209
+ Days the apparatus did not observe are dropped rather than recorded as
210
+ zeros. A missing day is missing evidence; treating it as an observation
211
+ of no activity is the single most damaging mistake this codebase exists
212
+ to avoid.
213
+ """
214
+ days: list[date] = []
215
+ values: list[float] = []
216
+ for summary in summaries:
217
+ if not summary.is_usable(min_coverage, min_observed):
218
+ logger.debug(
219
+ "Excluding poorly observed day %s from the series", summary.day
220
+ )
221
+ continue
222
+ days.append(summary.day)
223
+ values.append(summary.hours_in(state))
224
+ return days, values
@@ -0,0 +1,13 @@
1
+ """Change point detection algorithms."""
2
+
3
+ from .adaptive_normalization import AdaptiveNormalizer
4
+ from .embedding_cpd import EmbeddingCPD
5
+ from .energy_efficient import EnergyEfficientCPD
6
+ from .genetic_optimization import GeneticOptimizationCPD
7
+
8
+ __all__ = [
9
+ "EmbeddingCPD",
10
+ "EnergyEfficientCPD",
11
+ "AdaptiveNormalizer",
12
+ "GeneticOptimizationCPD",
13
+ ]
@@ -0,0 +1,31 @@
1
+ """Validation helpers for lightweight change-point detectors."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import numpy as np
6
+
7
+
8
+ def validate_positive_int(value: int, name: str) -> None:
9
+ """Require a positive integer configuration value."""
10
+ if isinstance(value, bool) or not isinstance(value, int) or value < 1:
11
+ raise ValueError(f"{name} must be a positive integer")
12
+
13
+
14
+ def validate_series(series: np.ndarray) -> np.ndarray:
15
+ """Return a one-dimensional finite numeric series."""
16
+ values = np.asarray(series, dtype=float)
17
+ if values.ndim != 1:
18
+ raise ValueError("series must be one-dimensional")
19
+ if values.size == 0:
20
+ raise ValueError("series must contain at least one observation")
21
+ if not np.isfinite(values).all():
22
+ raise ValueError("series must contain only finite values")
23
+ return values
24
+
25
+
26
+ def validate_positive_threshold(threshold: float) -> float:
27
+ """Return a positive finite detection threshold."""
28
+ value = float(threshold)
29
+ if not np.isfinite(value) or value <= 0:
30
+ raise ValueError("threshold must be positive and finite")
31
+ return value
@@ -0,0 +1,55 @@
1
+ """Adaptive normalization for real-time change point detection.
2
+
3
+ Based on Gupta et al. (2022) real-time preprocessing approach.
4
+ """
5
+
6
+ from __future__ import annotations
7
+
8
+ import logging
9
+ from dataclasses import dataclass
10
+
11
+ import numpy as np
12
+
13
+ from ..utils.plotting import plot_change_points
14
+ from ._validation import (
15
+ validate_positive_int,
16
+ validate_positive_threshold,
17
+ validate_series,
18
+ )
19
+
20
+ logger = logging.getLogger(__name__)
21
+
22
+
23
+ @dataclass
24
+ class AdaptiveNormalizer:
25
+ """Normalize data online before detecting change points."""
26
+
27
+ window: int = 20
28
+
29
+ def __post_init__(self) -> None:
30
+ validate_positive_int(self.window, "window")
31
+
32
+ def fit(self, series: np.ndarray) -> AdaptiveNormalizer:
33
+ self.series = validate_series(series)
34
+ return self
35
+
36
+ def _normalize(self) -> np.ndarray:
37
+ normed = []
38
+ for i in range(len(self.series)):
39
+ start = max(0, i - self.window)
40
+ segment = self.series[start : i + 1]
41
+ normed.append((self.series[i] - segment.mean()) / (segment.std() + 1e-6))
42
+ return np.array(normed)
43
+
44
+ def predict(self, threshold: float = 3.0, plot: bool = False) -> np.ndarray:
45
+ """Detect change points after adaptive normalization."""
46
+ if not hasattr(self, "series"):
47
+ raise ValueError("Model must be fitted before prediction")
48
+ threshold = validate_positive_threshold(threshold)
49
+ normed = self._normalize()
50
+ diffs = np.abs(np.diff(normed))
51
+ cps = np.where(diffs > threshold)[0] + 1
52
+ if plot:
53
+ plot_change_points(self.series, cps, title="Adaptive Normalization CPD")
54
+ logger.info("AdaptiveNormalizer detected %d change points", len(cps))
55
+ return cps
@@ -0,0 +1,60 @@
1
+ """Neural embedding change point detection.
2
+
3
+ Implementation inspired by Dadi et al. (2021) "ADAF".
4
+ """
5
+
6
+ from __future__ import annotations
7
+
8
+ import logging
9
+ from dataclasses import dataclass
10
+
11
+ import numpy as np
12
+
13
+ from ..utils.plotting import plot_change_points
14
+ from ._validation import (
15
+ validate_positive_int,
16
+ validate_positive_threshold,
17
+ validate_series,
18
+ )
19
+
20
+ logger = logging.getLogger(__name__)
21
+
22
+
23
+ @dataclass
24
+ class EmbeddingCPD:
25
+ """Detect change points using simple neural embeddings.
26
+
27
+ This simplified detector follows the spirit of the neural embedding
28
+ approach by Dadi et al. (2021). It computes rolling mean embeddings and
29
+ flags large embedding differences as change points.
30
+ """
31
+
32
+ window: int = 5
33
+
34
+ def __post_init__(self) -> None:
35
+ validate_positive_int(self.window, "window")
36
+
37
+ def fit(self, series: np.ndarray) -> EmbeddingCPD:
38
+ """Learn embeddings from the input series."""
39
+ self.series = validate_series(series)
40
+ kernel = np.ones(self.window) / self.window
41
+ self.embeddings = np.convolve(self.series, kernel, mode="valid")
42
+ logger.debug("Computed embeddings of length %d", len(self.embeddings))
43
+ return self
44
+
45
+ def predict(self, threshold: float = 0.2, plot: bool = False) -> np.ndarray:
46
+ """Return detected change point indices.
47
+
48
+ Args:
49
+ threshold: Difference threshold on successive embeddings.
50
+ plot: Whether to plot results using utility helpers.
51
+ """
52
+ if not hasattr(self, "embeddings"):
53
+ raise ValueError("Model must be fitted before prediction")
54
+ threshold = validate_positive_threshold(threshold)
55
+ diffs = np.abs(np.diff(self.embeddings))
56
+ cps = np.where(diffs > threshold)[0] + 1
57
+ if plot:
58
+ plot_change_points(self.series, cps, title="Embedding CPD")
59
+ logger.info("EmbeddingCPD detected %d change points", len(cps))
60
+ return cps
@@ -0,0 +1,57 @@
1
+ """Energy-efficient change point detection.
2
+
3
+ Simplified implementation of the CPAM algorithm from Cook et al. (2020).
4
+ """
5
+
6
+ from __future__ import annotations
7
+
8
+ import logging
9
+ from dataclasses import dataclass
10
+
11
+ import numpy as np
12
+
13
+ from ..utils.plotting import plot_change_points
14
+ from ._validation import (
15
+ validate_positive_int,
16
+ validate_positive_threshold,
17
+ validate_series,
18
+ )
19
+
20
+ logger = logging.getLogger(__name__)
21
+
22
+
23
+ @dataclass
24
+ class EnergyEfficientCPD:
25
+ """Detect change points using energy-efficient statistics."""
26
+
27
+ window: int = 10
28
+
29
+ def __post_init__(self) -> None:
30
+ validate_positive_int(self.window, "window")
31
+
32
+ def fit(self, series: np.ndarray) -> EnergyEfficientCPD:
33
+ self.series = validate_series(series)
34
+ return self
35
+
36
+ def predict(self, threshold: float = 0.5, plot: bool = False) -> np.ndarray:
37
+ """Detect change points based on energy across windows."""
38
+ if not hasattr(self, "series"):
39
+ raise ValueError("Model must be fitted before prediction")
40
+ threshold = validate_positive_threshold(threshold)
41
+ n = len(self.series)
42
+ energies = [
43
+ np.sum(
44
+ (
45
+ self.series[i : i + self.window]
46
+ - self.series[i : i + self.window].mean()
47
+ )
48
+ ** 2
49
+ )
50
+ for i in range(n - self.window)
51
+ ]
52
+ diffs = np.abs(np.diff(energies))
53
+ cps = np.where(diffs > threshold)[0] + 1
54
+ if plot:
55
+ plot_change_points(self.series, cps, title="Energy Efficient CPD")
56
+ logger.info("EnergyEfficientCPD detected %d change points", len(cps))
57
+ return cps
@@ -0,0 +1,65 @@
1
+ """Genetic algorithm based change point tuning.
2
+
3
+ Simplified GA approach inspired by Awais et al. (2016).
4
+ """
5
+
6
+ from __future__ import annotations
7
+
8
+ import logging
9
+ from dataclasses import dataclass
10
+
11
+ import numpy as np
12
+
13
+ from ..utils.plotting import plot_change_points
14
+ from ._validation import validate_positive_int, validate_series
15
+
16
+ logger = logging.getLogger(__name__)
17
+
18
+
19
+ @dataclass
20
+ class GeneticOptimizationCPD:
21
+ """Use a tiny genetic search to tune the detection threshold."""
22
+
23
+ population: int = 5
24
+ generations: int = 5
25
+
26
+ def __post_init__(self) -> None:
27
+ validate_positive_int(self.population, "population")
28
+ validate_positive_int(self.generations, "generations")
29
+
30
+ def fit(self, series: np.ndarray) -> GeneticOptimizationCPD:
31
+ self.series = validate_series(series)
32
+ self.threshold_ = self._ga_search()
33
+ logger.debug("Optimized threshold to %.3f", self.threshold_)
34
+ return self
35
+
36
+ def _ga_search(self) -> float:
37
+ rng = np.random.default_rng(0)
38
+ candidates = rng.uniform(0.1, 2.0, self.population)
39
+ for _ in range(self.generations):
40
+ scores = [self._fitness(th) for th in candidates]
41
+ best = np.argmin(scores)
42
+ # mutate around best
43
+ candidates = candidates[best] + rng.normal(0, 0.1, self.population)
44
+ candidates = np.clip(candidates, 0.1, 2.0)
45
+ return candidates[0]
46
+
47
+ def _fitness(self, threshold: float) -> float:
48
+ diffs = np.abs(np.diff(self.series))
49
+ cps = np.where(diffs > threshold)[0]
50
+ # simple fitness: prefer at least one cp around middle
51
+ target = len(self.series) // 2
52
+ if len(cps) == 0:
53
+ return target
54
+ return min(abs(cp - target) for cp in cps)
55
+
56
+ def predict(self, plot: bool = False) -> np.ndarray:
57
+ """Detect change points using GA-optimized threshold."""
58
+ if not hasattr(self, "threshold_"):
59
+ raise ValueError("Model must be fitted before prediction")
60
+ diffs = np.abs(np.diff(self.series))
61
+ cps = np.where(diffs > self.threshold_)[0] + 1
62
+ if plot:
63
+ plot_change_points(self.series, cps, title="GA Optimization CPD")
64
+ logger.info("GeneticOptimizationCPD detected %d change points", len(cps))
65
+ return cps