goad-toolkit 0.1.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.
- goad_toolkit/__init__.py +0 -0
- goad_toolkit/analytics.py +325 -0
- goad_toolkit/config.py +18 -0
- goad_toolkit/dataprocessor.py +51 -0
- goad_toolkit/datatransforms.py +184 -0
- goad_toolkit/distributions.py +80 -0
- goad_toolkit/filehandler.py +65 -0
- goad_toolkit/main.py +0 -0
- goad_toolkit/models.py +44 -0
- goad_toolkit/visualizer.py +505 -0
- goad_toolkit-0.1.0.dist-info/METADATA +231 -0
- goad_toolkit-0.1.0.dist-info/RECORD +14 -0
- goad_toolkit-0.1.0.dist-info/WHEEL +4 -0
- goad_toolkit-0.1.0.dist-info/entry_points.txt +2 -0
|
@@ -0,0 +1,65 @@
|
|
|
1
|
+
from pathlib import Path
|
|
2
|
+
from typing import Optional
|
|
3
|
+
|
|
4
|
+
import pandas as pd
|
|
5
|
+
import requests
|
|
6
|
+
from loguru import logger
|
|
7
|
+
|
|
8
|
+
from goad_toolkit.config import FileConfig
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
class FileHandler:
|
|
12
|
+
def __init__(self, config: FileConfig):
|
|
13
|
+
self.config = config
|
|
14
|
+
self.check_dirs()
|
|
15
|
+
|
|
16
|
+
def download(self, filename: Optional[Path] = None) -> None:
|
|
17
|
+
if not filename:
|
|
18
|
+
filename = self.config.filename
|
|
19
|
+
filepath = self.config.data_dir / "raw" / filename
|
|
20
|
+
if not filepath.exists():
|
|
21
|
+
req = requests.get(self.config.url, timeout=10)
|
|
22
|
+
with filepath.open("wb") as f:
|
|
23
|
+
f.write(req.content)
|
|
24
|
+
logger.success(f"Downloaded data to {filepath}")
|
|
25
|
+
else:
|
|
26
|
+
logger.info(f"{filepath} already exists\nRemove file to download again")
|
|
27
|
+
|
|
28
|
+
def check_dirs(self) -> None:
|
|
29
|
+
dirs = [
|
|
30
|
+
self.config.data_dir,
|
|
31
|
+
self.config.data_dir / "raw",
|
|
32
|
+
self.config.data_dir / "processed",
|
|
33
|
+
]
|
|
34
|
+
|
|
35
|
+
for dir in dirs:
|
|
36
|
+
if not dir.exists():
|
|
37
|
+
logger.info(f"Creating directory {dir}")
|
|
38
|
+
dir.mkdir(parents=True, exist_ok=True)
|
|
39
|
+
|
|
40
|
+
def load(self, filename: Optional[Path] = None, raw: bool = True) -> pd.DataFrame:
|
|
41
|
+
"""Load data from raw directory"""
|
|
42
|
+
if raw:
|
|
43
|
+
subdir = "raw"
|
|
44
|
+
else:
|
|
45
|
+
subdir = "processed"
|
|
46
|
+
|
|
47
|
+
if not filename:
|
|
48
|
+
filename = self.config.filename
|
|
49
|
+
|
|
50
|
+
file_path = self.config.data_dir / subdir / filename
|
|
51
|
+
if not file_path.exists() and raw:
|
|
52
|
+
self.download(filename=filename)
|
|
53
|
+
|
|
54
|
+
data = pd.read_csv(file_path, parse_dates=["date"], index_col="date")
|
|
55
|
+
logger.info(f"Loaded data from {file_path}")
|
|
56
|
+
return data
|
|
57
|
+
|
|
58
|
+
def save(self, data: pd.DataFrame, filename: Optional[Path]) -> None:
|
|
59
|
+
"""Save data to processed directory"""
|
|
60
|
+
if not filename:
|
|
61
|
+
filename = self.config.filename
|
|
62
|
+
|
|
63
|
+
processed_path = self.config.data_dir / "processed" / filename
|
|
64
|
+
data.to_csv(processed_path)
|
|
65
|
+
logger.success(f"Saved processed data to {processed_path}")
|
goad_toolkit/main.py
ADDED
|
File without changes
|
goad_toolkit/models.py
ADDED
|
@@ -0,0 +1,44 @@
|
|
|
1
|
+
from typing import Callable
|
|
2
|
+
|
|
3
|
+
import numpy as np
|
|
4
|
+
from scipy.optimize import minimize
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
def linear_model(x: np.ndarray, params: tuple | list[float]) -> np.ndarray:
|
|
8
|
+
"""a basic linear model"""
|
|
9
|
+
a, b = params
|
|
10
|
+
yhat = a * x + b
|
|
11
|
+
return yhat
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
def logistic(x: np.ndarray, k, x0, limit=1.0) -> np.ndarray:
|
|
15
|
+
"""
|
|
16
|
+
Parameters:
|
|
17
|
+
x: Independent variable
|
|
18
|
+
k: Growth rate
|
|
19
|
+
x0: Midpoint (inflection point)
|
|
20
|
+
L: Upper limit (default 1)
|
|
21
|
+
"""
|
|
22
|
+
return limit / (1 + np.exp(-k * (x - x0)))
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
def mse(y: np.ndarray, yhat: np.ndarray):
|
|
26
|
+
"""mean squared error loss function"""
|
|
27
|
+
squared_diff = (y - yhat) ** 2
|
|
28
|
+
return np.mean(squared_diff)
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
def train_model(
|
|
32
|
+
X: np.ndarray, # noqa: N803
|
|
33
|
+
y: np.ndarray,
|
|
34
|
+
model_fn: Callable,
|
|
35
|
+
loss_fn: Callable,
|
|
36
|
+
params: list[float],
|
|
37
|
+
bounds=None,
|
|
38
|
+
) -> list[float]:
|
|
39
|
+
def objective(params):
|
|
40
|
+
yhat = model_fn(X, params)
|
|
41
|
+
return loss_fn(y, yhat)
|
|
42
|
+
|
|
43
|
+
result = minimize(fun=objective, x0=params, bounds=bounds)
|
|
44
|
+
return result.x
|
|
@@ -0,0 +1,505 @@
|
|
|
1
|
+
from abc import ABC, abstractmethod
|
|
2
|
+
from dataclasses import dataclass, field
|
|
3
|
+
from typing import Any, List, Optional, Tuple
|
|
4
|
+
|
|
5
|
+
import matplotlib.dates as mdates
|
|
6
|
+
import matplotlib.pyplot as plt
|
|
7
|
+
import numpy as np
|
|
8
|
+
import pandas as pd
|
|
9
|
+
import seaborn as sns
|
|
10
|
+
from loguru import logger
|
|
11
|
+
from matplotlib.axes import Axes
|
|
12
|
+
from matplotlib.figure import Figure
|
|
13
|
+
|
|
14
|
+
from goad_toolkit.analytics import FitResult, Result
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
@dataclass
|
|
18
|
+
class PlotSettings:
|
|
19
|
+
"""Base settings for all plots."""
|
|
20
|
+
|
|
21
|
+
figsize: Tuple[int, int] = (10, 6)
|
|
22
|
+
title: str = "Plot"
|
|
23
|
+
xlabel: str = "X"
|
|
24
|
+
ylabel: str = "Y"
|
|
25
|
+
legend_title: Optional[str] = None
|
|
26
|
+
grid: bool = False # Added grid option
|
|
27
|
+
grid_alpha: Optional[float] = None # Grid transparency
|
|
28
|
+
subplot_titles: List[str] = field(default_factory=list) # Titles for subplots
|
|
29
|
+
|
|
30
|
+
def __repr__(self):
|
|
31
|
+
return (
|
|
32
|
+
f"PlotSettings(figsize={self.figsize},\n title={self.title},\n "
|
|
33
|
+
f"xlabel={self.xlabel},\n ylabel={self.ylabel},\n "
|
|
34
|
+
f"legend_title={self.legend_title},\n grid={self.grid},\n "
|
|
35
|
+
f"grid_alpha={self.grid_alpha}"
|
|
36
|
+
)
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
class BasePlot(ABC):
|
|
40
|
+
"""Base class for creating plots."""
|
|
41
|
+
|
|
42
|
+
def __init__(self, settings: PlotSettings, n_plots: Optional[int] = None):
|
|
43
|
+
self.settings = settings
|
|
44
|
+
self.fig = None
|
|
45
|
+
self.ax = None
|
|
46
|
+
|
|
47
|
+
def plot(self, n_plots: Optional[int] = None, *args, **kwargs):
|
|
48
|
+
"""Abstract method for plotting data."""
|
|
49
|
+
if self.fig is None:
|
|
50
|
+
self.create_figure(n_plots)
|
|
51
|
+
return self.build(*args, **kwargs)
|
|
52
|
+
|
|
53
|
+
def create_figure(self, n_plots: Optional[int] = None):
|
|
54
|
+
"""Create a figure and configure it based on settings.
|
|
55
|
+
Parameters:
|
|
56
|
+
-----------
|
|
57
|
+
n_plots : Optional[int]
|
|
58
|
+
Number of subplots to create. If None, creates a single plot.
|
|
59
|
+
If > 1, creates a grid of subplots.
|
|
60
|
+
"""
|
|
61
|
+
if n_plots is None or n_plots <= 1:
|
|
62
|
+
# Create a single plot
|
|
63
|
+
self.fig, self.ax = plt.subplots(figsize=self.settings.figsize)
|
|
64
|
+
axes = [self.ax]
|
|
65
|
+
else:
|
|
66
|
+
# Determine grid layout
|
|
67
|
+
grid_cols = min(3, n_plots) # Max 3 columns
|
|
68
|
+
grid_rows = int(np.ceil(n_plots / grid_cols))
|
|
69
|
+
# Create subplot grid
|
|
70
|
+
self.fig, axes = plt.subplots(
|
|
71
|
+
grid_rows, grid_cols, figsize=self.settings.figsize, squeeze=False
|
|
72
|
+
)
|
|
73
|
+
axes = axes.flatten()
|
|
74
|
+
# Hide unused axes
|
|
75
|
+
for i in range(n_plots, len(axes)):
|
|
76
|
+
axes[i].set_visible(False)
|
|
77
|
+
# Store first axis as default
|
|
78
|
+
self.ax = axes[0]
|
|
79
|
+
|
|
80
|
+
# Apply common settings to all axes
|
|
81
|
+
for i, ax in enumerate(axes):
|
|
82
|
+
if ax.get_visible():
|
|
83
|
+
ax.set_xlabel(self.settings.xlabel)
|
|
84
|
+
ax.set_ylabel(self.settings.ylabel)
|
|
85
|
+
# Apply grid if requested
|
|
86
|
+
if self.settings.grid:
|
|
87
|
+
ax.grid(True, alpha=self.settings.grid_alpha)
|
|
88
|
+
|
|
89
|
+
# Set subplot titles if available
|
|
90
|
+
if i < len(self.settings.subplot_titles):
|
|
91
|
+
ax.set_title(self.settings.subplot_titles[i])
|
|
92
|
+
|
|
93
|
+
# Set main title
|
|
94
|
+
if n_plots is None or n_plots <= 1:
|
|
95
|
+
if not self.settings.subplot_titles: # Only set if no subplot titles
|
|
96
|
+
self.ax.set_title(self.settings.title)
|
|
97
|
+
else:
|
|
98
|
+
plt.suptitle(self.settings.title, fontsize=16)
|
|
99
|
+
plt.tight_layout(rect=(0, 0, 1, 0.96)) # Make room for suptitle
|
|
100
|
+
|
|
101
|
+
if self.settings.legend_title is not None:
|
|
102
|
+
self.ax.legend(title=self.settings.legend_title)
|
|
103
|
+
|
|
104
|
+
return self.fig, axes
|
|
105
|
+
|
|
106
|
+
def plot_on(self, other_plot: "BasePlot", *args, **kwargs):
|
|
107
|
+
"""Combine BasePlot classes in a hierarchy.
|
|
108
|
+
|
|
109
|
+
Parameters:
|
|
110
|
+
-----------
|
|
111
|
+
other_plot : BasePlot
|
|
112
|
+
The plot class to use
|
|
113
|
+
*args, **kwargs : Arguments to pass to the plot method
|
|
114
|
+
|
|
115
|
+
Returns:
|
|
116
|
+
--------
|
|
117
|
+
The used plot instance (after plotting)
|
|
118
|
+
"""
|
|
119
|
+
# Create an instance of the plot class with our settings
|
|
120
|
+
# other_plot = plot_class(self.settings)
|
|
121
|
+
|
|
122
|
+
# Share our figure and axes
|
|
123
|
+
other_plot.fig = self.fig
|
|
124
|
+
other_plot.ax = self.ax
|
|
125
|
+
|
|
126
|
+
# Call the plot method with the provided arguments
|
|
127
|
+
other_plot.plot(*args, **kwargs)
|
|
128
|
+
|
|
129
|
+
# Return the plot in case further configuration is needed
|
|
130
|
+
return other_plot
|
|
131
|
+
|
|
132
|
+
def plot_on_axes(self, other_plot: "BasePlot", ax: Axes, *args, **kwargs):
|
|
133
|
+
"""Use another plot class to plot on a specific axis within this figure.
|
|
134
|
+
|
|
135
|
+
Parameters:
|
|
136
|
+
-----------
|
|
137
|
+
other_plot : BasePlot
|
|
138
|
+
The plot class to use
|
|
139
|
+
ax : matplotlib.axes.Axes
|
|
140
|
+
The specific axes to plot on
|
|
141
|
+
*args, **kwargs : Arguments to pass to the plot method
|
|
142
|
+
|
|
143
|
+
Returns:
|
|
144
|
+
--------
|
|
145
|
+
The used plot instance (after plotting)
|
|
146
|
+
"""
|
|
147
|
+
|
|
148
|
+
# Share our figure but use the provided axis
|
|
149
|
+
other_plot.fig = self.fig
|
|
150
|
+
other_plot.ax = ax
|
|
151
|
+
|
|
152
|
+
# Call the plot method with the provided arguments
|
|
153
|
+
other_plot.plot(*args, **kwargs)
|
|
154
|
+
|
|
155
|
+
# Return the plot in case further configuration is needed
|
|
156
|
+
return other_plot
|
|
157
|
+
|
|
158
|
+
@abstractmethod
|
|
159
|
+
def build(self, *args, **kwargs):
|
|
160
|
+
raise NotImplementedError("Plotting method must be implemented")
|
|
161
|
+
|
|
162
|
+
|
|
163
|
+
class LinePlot(BasePlot):
|
|
164
|
+
"""Plot a line plot using seaborn."""
|
|
165
|
+
|
|
166
|
+
def build(self, data: pd.DataFrame, **kwargs):
|
|
167
|
+
sns.lineplot(data=data, ax=self.ax, **kwargs)
|
|
168
|
+
return self.fig, self.ax
|
|
169
|
+
|
|
170
|
+
|
|
171
|
+
class ComparePlot(BasePlot):
|
|
172
|
+
def build(self, data: pd.DataFrame, x: str, y1: str, y2: str, **kwargs):
|
|
173
|
+
compare = LinePlot(self.settings)
|
|
174
|
+
self.plot_on(compare, data=data, x=x, y=y1, label=y1, **kwargs)
|
|
175
|
+
self.plot_on(compare, data=data, x=x, y=y2, label=y2, **kwargs)
|
|
176
|
+
plt.xticks(rotation=45)
|
|
177
|
+
|
|
178
|
+
return self.fig, self.ax
|
|
179
|
+
|
|
180
|
+
|
|
181
|
+
class BarWithDates(BasePlot):
|
|
182
|
+
def build(self, data, x: str, y: str, interval: int = 1, **kwargs):
|
|
183
|
+
sns.barplot(data=data, x=x, y=y, ax=self.ax, **kwargs)
|
|
184
|
+
if not self.ax:
|
|
185
|
+
raise ValueError("No axes available for plotting")
|
|
186
|
+
self.ax.xaxis.set_major_locator(mdates.MonthLocator(interval=interval))
|
|
187
|
+
|
|
188
|
+
|
|
189
|
+
class VerticalDate(BasePlot):
|
|
190
|
+
def build(self, date: str, label: str):
|
|
191
|
+
start_vaccination = pd.to_datetime(date).strftime("%Y-%m-%d")
|
|
192
|
+
plt.axvline(
|
|
193
|
+
x=start_vaccination, # type: ignore
|
|
194
|
+
color="red",
|
|
195
|
+
linestyle="--",
|
|
196
|
+
linewidth=2,
|
|
197
|
+
label=label,
|
|
198
|
+
)
|
|
199
|
+
|
|
200
|
+
|
|
201
|
+
class ResidualPlot(BasePlot):
|
|
202
|
+
def build(self, data, x: str, y: str, date: str, datelabel: str, interval: int = 1):
|
|
203
|
+
barplot = BarWithDates(self.settings)
|
|
204
|
+
self.plot_on(barplot, data=data, x=x, y=y, interval=interval)
|
|
205
|
+
vertical = VerticalDate(self.settings)
|
|
206
|
+
self.plot_on(vertical, date=date, label=datelabel)
|
|
207
|
+
plt.xticks(rotation=45)
|
|
208
|
+
return self.fig, self.ax
|
|
209
|
+
|
|
210
|
+
|
|
211
|
+
class HistogramPlot(BasePlot):
|
|
212
|
+
"""Plot a histogram using seaborn."""
|
|
213
|
+
|
|
214
|
+
def build(
|
|
215
|
+
self,
|
|
216
|
+
data: np.ndarray,
|
|
217
|
+
bins: Optional[int] = None,
|
|
218
|
+
kde: bool = False,
|
|
219
|
+
color: str = "skyblue",
|
|
220
|
+
alpha: float = 0.7,
|
|
221
|
+
**kwargs,
|
|
222
|
+
):
|
|
223
|
+
"""
|
|
224
|
+
Create a histogram plot of the provided data.
|
|
225
|
+
|
|
226
|
+
Parameters:
|
|
227
|
+
-----------
|
|
228
|
+
data : np.ndarray
|
|
229
|
+
Data to plot
|
|
230
|
+
bins : Optional[int]
|
|
231
|
+
Number of bins to use
|
|
232
|
+
kde : bool
|
|
233
|
+
Whether to overlay a KDE plot
|
|
234
|
+
color : str
|
|
235
|
+
Color of the histogram
|
|
236
|
+
alpha : float
|
|
237
|
+
Transparency of the histogram
|
|
238
|
+
**kwargs : Additional keyword arguments passed to sns.histplot
|
|
239
|
+
|
|
240
|
+
Returns:
|
|
241
|
+
--------
|
|
242
|
+
fig, ax : The created figure and axes
|
|
243
|
+
"""
|
|
244
|
+
# Calculate optimal bins if not specified
|
|
245
|
+
if bins is None:
|
|
246
|
+
bins = min(int(np.sqrt(len(data))), 50) # Reasonable default
|
|
247
|
+
|
|
248
|
+
# Plot histogram
|
|
249
|
+
sns.histplot(
|
|
250
|
+
data,
|
|
251
|
+
bins=bins,
|
|
252
|
+
kde=kde,
|
|
253
|
+
color=color,
|
|
254
|
+
alpha=alpha,
|
|
255
|
+
ax=self.ax,
|
|
256
|
+
stat="density", # Use density for overlay compatibility
|
|
257
|
+
**kwargs,
|
|
258
|
+
)
|
|
259
|
+
|
|
260
|
+
return self.fig, self.ax
|
|
261
|
+
|
|
262
|
+
|
|
263
|
+
class DistPlot(BasePlot):
|
|
264
|
+
"""Plot a parametric distribution."""
|
|
265
|
+
|
|
266
|
+
def build(
|
|
267
|
+
self,
|
|
268
|
+
distribution: Any,
|
|
269
|
+
x_range: Optional[Tuple[float, float]] = None,
|
|
270
|
+
samples: int = 1000,
|
|
271
|
+
color: str = "crimson",
|
|
272
|
+
linewidth: float = 2,
|
|
273
|
+
label: Optional[str] = None,
|
|
274
|
+
**kwargs,
|
|
275
|
+
) -> Tuple[Any, Any]:
|
|
276
|
+
"""
|
|
277
|
+
Create a KDE plot from a given distribution.
|
|
278
|
+
|
|
279
|
+
Parameters:
|
|
280
|
+
-----------
|
|
281
|
+
distribution : scipy.stats distribution
|
|
282
|
+
Distribution to plot (must have a pdf method)
|
|
283
|
+
x_range : Optional[Tuple[float, float]]
|
|
284
|
+
Range of x values to plot (if None, will estimate from distribution)
|
|
285
|
+
samples : int
|
|
286
|
+
Number of points to sample along x-axis
|
|
287
|
+
color : str
|
|
288
|
+
Color of the KDE line
|
|
289
|
+
linewidth : float
|
|
290
|
+
Width of the KDE line
|
|
291
|
+
label : Optional[str]
|
|
292
|
+
Label for the plot in legend
|
|
293
|
+
**kwargs : Additional keyword arguments passed to plt.plot
|
|
294
|
+
|
|
295
|
+
Returns:
|
|
296
|
+
--------
|
|
297
|
+
fig, ax : The created figure and axes
|
|
298
|
+
"""
|
|
299
|
+
x = self._estimate_x(distribution, samples, x_range)
|
|
300
|
+
|
|
301
|
+
# Calculate probability density
|
|
302
|
+
try:
|
|
303
|
+
y = distribution.pdf(x)
|
|
304
|
+
except AttributeError:
|
|
305
|
+
# If pdf not available, try pmf for discrete distributions
|
|
306
|
+
try:
|
|
307
|
+
y = distribution.pmf(x)
|
|
308
|
+
except (AttributeError, ValueError):
|
|
309
|
+
raise ValueError("Distribution must have either pdf or pmf method")
|
|
310
|
+
|
|
311
|
+
if self.ax is None:
|
|
312
|
+
raise ValueError("No axes available for plotting")
|
|
313
|
+
|
|
314
|
+
# Plot distribution
|
|
315
|
+
self.ax.plot(x, y, color=color, linewidth=linewidth, label=label, **kwargs)
|
|
316
|
+
|
|
317
|
+
if label is not None:
|
|
318
|
+
self.ax.legend()
|
|
319
|
+
|
|
320
|
+
return self.fig, self.ax
|
|
321
|
+
|
|
322
|
+
def _estimate_x(
|
|
323
|
+
self,
|
|
324
|
+
distribution: Any,
|
|
325
|
+
samples: int,
|
|
326
|
+
x_range: Optional[Tuple[float, float]] = None,
|
|
327
|
+
) -> np.ndarray:
|
|
328
|
+
"""
|
|
329
|
+
Estimate an appropriate x-range for plotting the distribution.
|
|
330
|
+
"""
|
|
331
|
+
if x_range is not None:
|
|
332
|
+
x = np.linspace(x_range[0], x_range[1], samples)
|
|
333
|
+
return x
|
|
334
|
+
|
|
335
|
+
# Try to use the percent point function (quantile function)
|
|
336
|
+
try:
|
|
337
|
+
lower = distribution.ppf(0.001)
|
|
338
|
+
upper = distribution.ppf(0.999)
|
|
339
|
+
# Add padding
|
|
340
|
+
padding = (upper - lower) * 0.1
|
|
341
|
+
# Create x values within range
|
|
342
|
+
x_range = (lower - padding, upper + padding)
|
|
343
|
+
x = np.linspace(x_range[0], x_range[1], samples)
|
|
344
|
+
return x
|
|
345
|
+
except (AttributeError, ValueError):
|
|
346
|
+
logger.warning(
|
|
347
|
+
"Falling back to generic range. Specify x_range for better results."
|
|
348
|
+
)
|
|
349
|
+
x_range = (-5, 5)
|
|
350
|
+
x = np.linspace(x_range[0], x_range[1], samples)
|
|
351
|
+
return x
|
|
352
|
+
|
|
353
|
+
|
|
354
|
+
@dataclass
|
|
355
|
+
class FitPlotSettings:
|
|
356
|
+
"""Settings for distribution fit plots."""
|
|
357
|
+
|
|
358
|
+
bins: Optional[int] = None
|
|
359
|
+
data_color: str = "lightgrey"
|
|
360
|
+
best_likelihood_color: str = "crimson"
|
|
361
|
+
best_ks_color: str = "darkblue"
|
|
362
|
+
other_color: str = "gray"
|
|
363
|
+
data_alpha: float = 0.6
|
|
364
|
+
max_fits: Optional[int] = None
|
|
365
|
+
|
|
366
|
+
def __repr__(self):
|
|
367
|
+
return (
|
|
368
|
+
f"FitPlotSettings(bins={self.bins},\n data_color={self.data_color},\n "
|
|
369
|
+
f"best_likelihood_color={self.best_likelihood_color},\n "
|
|
370
|
+
f"best_ks_color={self.best_ks_color},\n other_color={self.other_color},\n "
|
|
371
|
+
f"data_alpha={self.data_alpha},\n max_fits={self.max_fits})"
|
|
372
|
+
)
|
|
373
|
+
|
|
374
|
+
|
|
375
|
+
class PlotFits(BasePlot):
|
|
376
|
+
"""Plot histogram of data with fitted distributions overlaid."""
|
|
377
|
+
|
|
378
|
+
def plot(
|
|
379
|
+
self,
|
|
380
|
+
data: np.ndarray,
|
|
381
|
+
fit_results: List[Result],
|
|
382
|
+
fitplotsettings: "FitPlotSettings",
|
|
383
|
+
) -> Figure:
|
|
384
|
+
"""
|
|
385
|
+
Plot multiple fits on separate subplots.
|
|
386
|
+
|
|
387
|
+
Parameters:
|
|
388
|
+
-----------
|
|
389
|
+
data : np.ndarray
|
|
390
|
+
Data to plot
|
|
391
|
+
fit_results : List[FitResult]
|
|
392
|
+
List of fit results to plot (expect them to have best_likelihood and best_ks attributes)
|
|
393
|
+
fitplotsettings : FitPlotSettings
|
|
394
|
+
Settings for the fit plots
|
|
395
|
+
|
|
396
|
+
Returns:
|
|
397
|
+
--------
|
|
398
|
+
fig : The created figure
|
|
399
|
+
|
|
400
|
+
"""
|
|
401
|
+
# Filter and sort fits
|
|
402
|
+
sorted_fits = self._prepare_fits(fit_results, fitplotsettings.max_fits)
|
|
403
|
+
|
|
404
|
+
# Create subplot titles
|
|
405
|
+
subplot_titles = self._create_subplot_titles(sorted_fits)
|
|
406
|
+
self.settings.subplot_titles = subplot_titles
|
|
407
|
+
|
|
408
|
+
# Create figure with subplots
|
|
409
|
+
self.fig, axes = self.create_figure(n_plots=len(sorted_fits))
|
|
410
|
+
|
|
411
|
+
# Plot each fit
|
|
412
|
+
for i, fit in enumerate(sorted_fits):
|
|
413
|
+
self._plot_single_fit(
|
|
414
|
+
data=data, fit=fit, ax=axes[i], fitplotsettings=fitplotsettings
|
|
415
|
+
)
|
|
416
|
+
|
|
417
|
+
return self.fig
|
|
418
|
+
|
|
419
|
+
def build(self):
|
|
420
|
+
pass
|
|
421
|
+
|
|
422
|
+
def _prepare_fits(
|
|
423
|
+
self, fit_results: List[Result], max_fits: Optional[int] = None
|
|
424
|
+
) -> List[FitResult]:
|
|
425
|
+
# Filter only successful fits
|
|
426
|
+
successful_fits = [fit for fit in fit_results if isinstance(fit, FitResult)]
|
|
427
|
+
|
|
428
|
+
if not successful_fits:
|
|
429
|
+
raise ValueError("No successful fits to plot")
|
|
430
|
+
|
|
431
|
+
# Sort fits by log-likelihood (descending)
|
|
432
|
+
sorted_fits = sorted(
|
|
433
|
+
successful_fits,
|
|
434
|
+
key=lambda fit: fit.log_likelihood
|
|
435
|
+
if fit.log_likelihood is not None
|
|
436
|
+
else -np.inf,
|
|
437
|
+
reverse=True,
|
|
438
|
+
)
|
|
439
|
+
|
|
440
|
+
# Limit number of fits if specified
|
|
441
|
+
if max_fits is not None and max_fits > 0:
|
|
442
|
+
sorted_fits = sorted_fits[:max_fits]
|
|
443
|
+
|
|
444
|
+
return sorted_fits
|
|
445
|
+
|
|
446
|
+
def _create_subplot_titles(self, fits: List[FitResult]) -> List[str]:
|
|
447
|
+
subplot_titles = []
|
|
448
|
+
for fit in fits:
|
|
449
|
+
title = f"{fit.distribution}"
|
|
450
|
+
if fit.kstest:
|
|
451
|
+
title += f"\nKS p: {fit.kstest.p_value:.4f}"
|
|
452
|
+
if fit.log_likelihood is not None:
|
|
453
|
+
title += f", llh: {fit.log_likelihood:.2f}"
|
|
454
|
+
if hasattr(fit, "best_likelihood") and fit.best_likelihood:
|
|
455
|
+
title += "\n★ Best likelihood ★"
|
|
456
|
+
if hasattr(fit, "best_ks") and fit.best_ks:
|
|
457
|
+
title += "\n★ Best KS test ★"
|
|
458
|
+
subplot_titles.append(title)
|
|
459
|
+
return subplot_titles
|
|
460
|
+
|
|
461
|
+
def _plot_single_fit(
|
|
462
|
+
self, data: np.ndarray, fit: Any, ax: Axes, fitplotsettings: "FitPlotSettings"
|
|
463
|
+
) -> None:
|
|
464
|
+
# Choose color based on best fit status
|
|
465
|
+
dist_color = self._get_fit_color(fit, fitplotsettings)
|
|
466
|
+
|
|
467
|
+
# Create histogram
|
|
468
|
+
hist_plot = HistogramPlot(
|
|
469
|
+
PlotSettings(
|
|
470
|
+
xlabel=self.settings.xlabel,
|
|
471
|
+
ylabel=self.settings.ylabel,
|
|
472
|
+
grid=self.settings.grid,
|
|
473
|
+
grid_alpha=self.settings.grid_alpha,
|
|
474
|
+
)
|
|
475
|
+
)
|
|
476
|
+
# hist_plot.fig, hist_plot.ax = self.fig, ax
|
|
477
|
+
self.plot_on_axes(
|
|
478
|
+
hist_plot,
|
|
479
|
+
data=data,
|
|
480
|
+
ax=ax,
|
|
481
|
+
bins=fitplotsettings.bins,
|
|
482
|
+
kde=False,
|
|
483
|
+
color=fitplotsettings.data_color,
|
|
484
|
+
alpha=fitplotsettings.data_alpha,
|
|
485
|
+
)
|
|
486
|
+
|
|
487
|
+
# Overlay distribution
|
|
488
|
+
kde_plot = DistPlot(PlotSettings())
|
|
489
|
+
# kde_plot.fig, kde_plot.ax = self.fig, ax
|
|
490
|
+
# kde_plot.plot(fit.frozen_dist, color=dist_color, label=fit.distribution)
|
|
491
|
+
self.plot_on_axes(
|
|
492
|
+
kde_plot,
|
|
493
|
+
ax=ax,
|
|
494
|
+
distribution=fit.frozen_dist,
|
|
495
|
+
color=dist_color,
|
|
496
|
+
label=fit.distribution,
|
|
497
|
+
)
|
|
498
|
+
|
|
499
|
+
def _get_fit_color(self, fit: Any, fitplotsettings: "FitPlotSettings") -> str:
|
|
500
|
+
if hasattr(fit, "best_likelihood") and fit.best_likelihood:
|
|
501
|
+
return fitplotsettings.best_likelihood_color
|
|
502
|
+
elif hasattr(fit, "best_ks") and fit.best_ks:
|
|
503
|
+
return fitplotsettings.best_ks_color
|
|
504
|
+
else:
|
|
505
|
+
return fitplotsettings.other_color
|