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.
@@ -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