plot-misc 2.2.1__tar.gz → 2.3.0__tar.gz

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 (46) hide show
  1. {plot_misc-2.2.1/plot_misc.egg-info → plot_misc-2.3.0}/PKG-INFO +26 -10
  2. {plot_misc-2.2.1 → plot_misc-2.3.0}/README.md +21 -7
  3. plot_misc-2.3.0/plot_misc/__init__.py +12 -0
  4. plot_misc-2.3.0/plot_misc/_version.py +1 -0
  5. {plot_misc-2.2.1 → plot_misc-2.3.0}/plot_misc/constants.py +6 -0
  6. {plot_misc-2.2.1 → plot_misc-2.3.0}/plot_misc/example_data/examples.py +48 -0
  7. {plot_misc-2.2.1 → plot_misc-2.3.0}/plot_misc/forest.py +0 -14
  8. plot_misc-2.3.0/plot_misc/heatmap.py +466 -0
  9. {plot_misc-2.2.1 → plot_misc-2.3.0}/plot_misc/machine_learning.py +10 -14
  10. {plot_misc-2.2.1 → plot_misc-2.3.0}/plot_misc/survival.py +1 -0
  11. {plot_misc-2.2.1 → plot_misc-2.3.0}/plot_misc/utils/formatting.py +1 -1
  12. {plot_misc-2.2.1 → plot_misc-2.3.0}/plot_misc/utils/utils.py +275 -135
  13. {plot_misc-2.2.1 → plot_misc-2.3.0}/plot_misc/volcano.py +31 -12
  14. {plot_misc-2.2.1 → plot_misc-2.3.0/plot_misc.egg-info}/PKG-INFO +26 -10
  15. {plot_misc-2.2.1 → plot_misc-2.3.0}/plot_misc.egg-info/requires.txt +3 -0
  16. {plot_misc-2.2.1 → plot_misc-2.3.0}/pyproject.toml +4 -3
  17. {plot_misc-2.2.1 → plot_misc-2.3.0}/requirements.txt +1 -0
  18. plot_misc-2.2.1/plot_misc/__init__.py +0 -1
  19. plot_misc-2.2.1/plot_misc/_version.py +0 -1
  20. plot_misc-2.2.1/plot_misc/heatmap.py +0 -254
  21. {plot_misc-2.2.1 → plot_misc-2.3.0}/LICENSE +0 -0
  22. {plot_misc-2.2.1 → plot_misc-2.3.0}/MANIFEST.in +0 -0
  23. {plot_misc-2.2.1 → plot_misc-2.3.0}/plot_misc/barchart.py +0 -0
  24. {plot_misc-2.2.1 → plot_misc-2.3.0}/plot_misc/errors.py +0 -0
  25. {plot_misc-2.2.1 → plot_misc-2.3.0}/plot_misc/example_data/__init__.py +0 -0
  26. {plot_misc-2.2.1 → plot_misc-2.3.0}/plot_misc/example_data/example_datasets/bar_points.tsv.gz +0 -0
  27. {plot_misc-2.2.1 → plot_misc-2.3.0}/plot_misc/example_data/example_datasets/barchart.tsv.gz +0 -0
  28. {plot_misc-2.2.1 → plot_misc-2.3.0}/plot_misc/example_data/example_datasets/calibration_bins.tsv.gz +0 -0
  29. {plot_misc-2.2.1 → plot_misc-2.3.0}/plot_misc/example_data/example_datasets/calibration_data.tsv.gz +0 -0
  30. {plot_misc-2.2.1 → plot_misc-2.3.0}/plot_misc/example_data/example_datasets/forest_data.tsv.gz +0 -0
  31. {plot_misc-2.2.1 → plot_misc-2.3.0}/plot_misc/example_data/example_datasets/group_bar.tsv.gz +0 -0
  32. {plot_misc-2.2.1 → plot_misc-2.3.0}/plot_misc/example_data/example_datasets/heatmap_data.tsv.gz +0 -0
  33. {plot_misc-2.2.1 → plot_misc-2.3.0}/plot_misc/example_data/example_datasets/incidence_matrix_data.tsv.gz +0 -0
  34. {plot_misc-2.2.1 → plot_misc-2.3.0}/plot_misc/example_data/example_datasets/lollipop_data.tsv.gz +0 -0
  35. {plot_misc-2.2.1 → plot_misc-2.3.0}/plot_misc/example_data/example_datasets/mace_associations.tsv.gz +0 -0
  36. {plot_misc-2.2.1 → plot_misc-2.3.0}/plot_misc/example_data/example_datasets/net_benefit.tsv.gz +0 -0
  37. {plot_misc-2.2.1 → plot_misc-2.3.0}/plot_misc/example_data/example_datasets/volcano.tsv.gz +0 -0
  38. {plot_misc-2.2.1 → plot_misc-2.3.0}/plot_misc/incidencematrix.py +0 -0
  39. {plot_misc-2.2.1 → plot_misc-2.3.0}/plot_misc/piechart.py +0 -0
  40. {plot_misc-2.2.1 → plot_misc-2.3.0}/plot_misc/utils/__init__.py +0 -0
  41. {plot_misc-2.2.1 → plot_misc-2.3.0}/plot_misc/utils/colour.py +0 -0
  42. {plot_misc-2.2.1 → plot_misc-2.3.0}/plot_misc.egg-info/SOURCES.txt +0 -0
  43. {plot_misc-2.2.1 → plot_misc-2.3.0}/plot_misc.egg-info/dependency_links.txt +0 -0
  44. {plot_misc-2.2.1 → plot_misc-2.3.0}/plot_misc.egg-info/top_level.txt +0 -0
  45. {plot_misc-2.2.1 → plot_misc-2.3.0}/requirements-dev.txt +0 -0
  46. {plot_misc-2.2.1 → plot_misc-2.3.0}/setup.cfg +0 -0
@@ -1,7 +1,7 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: plot-misc
3
- Version: 2.2.1
4
- Summary: Various plotting templates built on top of matplotlib
3
+ Version: 2.3.0
4
+ Summary: Various plotting archetypes built on top of matplotlib
5
5
  Author-email: A Floriaan Schmidt <floriaanschmidt@gmail.com>
6
6
  License-Expression: GPL-3.0-or-later
7
7
  Project-URL: Homepage, https://gitlab.com/SchmidtAF/plot-misc
@@ -11,8 +11,9 @@ Classifier: Programming Language :: Python :: 3
11
11
  Classifier: Programming Language :: Python :: 3.10
12
12
  Classifier: Programming Language :: Python :: 3.11
13
13
  Classifier: Programming Language :: Python :: 3.12
14
+ Classifier: Programming Language :: Python :: 3.13
14
15
  Classifier: Programming Language :: Python :: Implementation :: PyPy
15
- Requires-Python: <3.13,>=3.10
16
+ Requires-Python: <3.14,>=3.10
16
17
  Description-Content-Type: text/markdown
17
18
  License-File: LICENSE
18
19
  Requires-Dist: pandas>=1.3
@@ -22,6 +23,7 @@ Requires-Dist: scipy>=1.5
22
23
  Requires-Dist: statsmodels>=0.1
23
24
  Requires-Dist: scikit-learn>=1.4
24
25
  Requires-Dist: adjustText>=1.3
26
+ Requires-Dist: typing_extensions>=4; python_version < "3.11"
25
27
  Provides-Extra: dev
26
28
  Requires-Dist: python-build; extra == "dev"
27
29
  Requires-Dist: twine; extra == "dev"
@@ -48,15 +50,29 @@ Dynamic: license-file
48
50
  <img src="https://schmidtaf.gitlab.io/plot-misc/_images/icon.png" alt="plot-misc icon" width="250"/>
49
51
 
50
52
  # A collection of plotting functions
51
- __version__: `2.2.1`
53
+ __version__: `2.3.0`
52
54
 
53
55
  This repository collects plotting modules written on top of `matplotlib`.
54
- The functions are intended to set up light-touch, basic illustrations that
55
- can be customised using the standard matplotlib interface via axes and figures.
56
- Functionality is included to create illustrations commonly used in medical research,
57
- covering forest plots, volcano plots, incidence matrices/bubble charts,
58
- illustrations to evaluate prediction models (e.g. feature importance, net benefit, calibration plots),
59
- and more.
56
+ The functions describe plotting archetypes intended to set up light-touch,
57
+ illustrations that can be customised using the standard matplotlib interface
58
+ via axes and figures.
59
+ Because the implementation is matplotlib-first, the API is consistent with
60
+ matplotlib conventions, and users already familiar with the library will find
61
+ the learning curve minimal.
62
+
63
+ The functionality is geared towards illustrations commonly used in biomedical
64
+ research:
65
+
66
+ * Bar charts
67
+ * Bubble charts
68
+ * Forest plots (with optional side-tables)
69
+ * Heatmaps (with optional annotations)
70
+ * Incidence matrix plots
71
+ * Machine learning plots (calibration, feature importance, net benefit)
72
+ * Pie charts
73
+ * Survival plots (with optional survival table)
74
+ * Tree/compatibility plots
75
+ * Volcano plots
60
76
 
61
77
  Please consult the **[documentation](https://SchmidtAF.gitlab.io/plot-misc/)**
62
78
  for plot-misc.
@@ -1,15 +1,29 @@
1
1
  <img src="https://schmidtaf.gitlab.io/plot-misc/_images/icon.png" alt="plot-misc icon" width="250"/>
2
2
 
3
3
  # A collection of plotting functions
4
- __version__: `2.2.1`
4
+ __version__: `2.3.0`
5
5
 
6
6
  This repository collects plotting modules written on top of `matplotlib`.
7
- The functions are intended to set up light-touch, basic illustrations that
8
- can be customised using the standard matplotlib interface via axes and figures.
9
- Functionality is included to create illustrations commonly used in medical research,
10
- covering forest plots, volcano plots, incidence matrices/bubble charts,
11
- illustrations to evaluate prediction models (e.g. feature importance, net benefit, calibration plots),
12
- and more.
7
+ The functions describe plotting archetypes intended to set up light-touch,
8
+ illustrations that can be customised using the standard matplotlib interface
9
+ via axes and figures.
10
+ Because the implementation is matplotlib-first, the API is consistent with
11
+ matplotlib conventions, and users already familiar with the library will find
12
+ the learning curve minimal.
13
+
14
+ The functionality is geared towards illustrations commonly used in biomedical
15
+ research:
16
+
17
+ * Bar charts
18
+ * Bubble charts
19
+ * Forest plots (with optional side-tables)
20
+ * Heatmaps (with optional annotations)
21
+ * Incidence matrix plots
22
+ * Machine learning plots (calibration, feature importance, net benefit)
23
+ * Pie charts
24
+ * Survival plots (with optional survival table)
25
+ * Tree/compatibility plots
26
+ * Volcano plots
13
27
 
14
28
  Please consult the **[documentation](https://SchmidtAF.gitlab.io/plot-misc/)**
15
29
  for plot-misc.
@@ -0,0 +1,12 @@
1
+ """
2
+ plot-misc: matplotlib-based plotting archetypes for scientific figures.
3
+
4
+ A curated collection of user-oriented plotting functions and
5
+ classes built on top of `matplotlib`. Each function returns standard
6
+ matplotlib `Figure`/`Axes` objects, so results can be further customised
7
+ with familiar matplotlib methods. Per the package design callables are limited
8
+ to creating illustrations, and should data be internally calculated this is
9
+ done with options for user overwrites, while making the derived data available
10
+ for inspection and re-use.
11
+ """
12
+ from ._version import __version__
@@ -0,0 +1 @@
1
+ __version__ = '2.3.0'
@@ -58,6 +58,8 @@ class UtilsNames(object):
58
58
  annot_pval = 'matrix_pvalue'
59
59
  annot_effect = 'matrix_point_estimate'
60
60
  value_point = 'curated_matrix_point_estimate_value'
61
+ value_unsigned_log = 'curated_matrix_value_unsigned_log'
62
+ value_raw = 'curated_matrix_value_raw'
61
63
  value_original = 'crude_point_estimate'
62
64
  source_data = 'source_data'
63
65
  mat_point = 'point'
@@ -67,8 +69,12 @@ class UtilsNames(object):
67
69
  mat_outcome = 'outcome'
68
70
  mat_exposure_list = ['IL2ra', 'IP10', 'SCF', 'TRAIL']
69
71
  mat_outcome_list = ['HDL-C', 'LDL-C']
72
+ mat_annot_symbol = 'symbol'
70
73
  mat_annot_star = 'star'
71
74
  mat_annot_pval = 'pvalues'
75
+ mat_annot_pval_signed = 'pvalues_signed'
76
+ mat_annot_pval_unsigned = 'pvalues_unsigned'
77
+ mat_annot_pval_raw = 'pvalues_raw'
72
78
  mat_annot_point = 'point_estimates'
73
79
  mat_annot_none = '`NoneType`'
74
80
  roc_false_positive = 'false_positive'
@@ -508,6 +508,54 @@ def heatmap_pvalue_matrix(**kwargs):
508
508
  data.index.name = UtilsNames.mat_outcome
509
509
  return data
510
510
 
511
+ # ~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
512
+ @dataset
513
+ def qc_matrix(seed=2026):
514
+ """
515
+ Creates a dummy quality-control (QC) data set to showcase
516
+ `plot_misc.heatmap.masked_heatmap`.
517
+
518
+ The values are signed, standardised QC deviations for a set of samples
519
+ (rows) across several QC metrics (columns). The accompanying indicator
520
+ flags the cells that failed QC (an absolute deviation above 2), i.e. the
521
+ cells `masked_heatmap` should highlight; the passing cells are left to the
522
+ background layer.
523
+
524
+ Parameters
525
+ ----------
526
+ seed : `int`, default 2026
527
+ Seed for the random number generator, ensuring a reproducible matrix.
528
+
529
+ Returns
530
+ -------
531
+ values : `pd.DataFrame`
532
+ Signed standardised QC deviations of shape (12, 6).
533
+ indicator : `pd.DataFrame`
534
+ A binary (0/1) table of the same shape as `values`, equal to 1 where
535
+ the metric failed QC and 0 otherwise.
536
+ """
537
+ rng = np.random.default_rng(seed)
538
+ samples = ['Sample_{:02d}'.format(i) for i in range(1, 13)]
539
+ metrics = [
540
+ 'CallRate', 'Heterozygosity', 'Contamination', 'MeanDepth',
541
+ 'DuplicationRate', 'InsertSize',
542
+ ]
543
+ values = pd.DataFrame(
544
+ rng.normal(loc=0.0, scale=1.3, size=(len(samples), len(metrics))),
545
+ index=samples, columns=metrics,
546
+ )
547
+ # inject a handful of unambiguous QC failures so the showcase always has
548
+ # highlighted cells regardless of the random draw
549
+ values.iloc[0, 2] = 3.4
550
+ values.iloc[3, 0] = -3.1
551
+ values.iloc[5, 4] = 2.8
552
+ values.iloc[7, 1] = -2.6
553
+ values.iloc[9, 5] = 3.0
554
+ values.iloc[11, 3] = -2.9
555
+ values = values.round(3)
556
+ indicator = (values.abs() > 2).astype(int)
557
+ return values, indicator
558
+
511
559
  # ~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
512
560
  @dataset
513
561
  def load_calibration_data(**kwargs):
@@ -982,20 +982,6 @@ class EmpiricalSupport(object):
982
982
  results_ : `EmpiricalSupportResults`
983
983
  An EmpiricalSupportResults instance.
984
984
 
985
- Methods
986
- -------
987
- calc_empirical_support(estimate, standard_error, alpha)
988
- Computes the range of confidence intervals and compatibility metrics
989
- over the supplied alpha values.
990
-
991
- plot_tree(...)
992
- Creates a 'tree plot' summarising the parameter space supported by
993
- the data, with options for CI and estimate annotations.
994
-
995
- _plot_empirical_support(...)
996
- Generates a visualisation of confidence intervals and their overlap
997
- across varying alpha values.
998
-
999
985
  Notes
1000
986
  -----
1001
987
  This implementation is based on the concept of compatibility (or
@@ -0,0 +1,466 @@
1
+ """
2
+ Heatmap drawing and annotation tools built on top of matplotlib and seaborn.
3
+
4
+ This module provides flexible functions to create and annotate heatmaps using
5
+ either `matplotlib` or `seaborn`, with extensive support for customisation and
6
+ publication-quality output.
7
+
8
+ Functions
9
+ ---------
10
+ heatmap(data, row_labels, col_labels, ...)
11
+ Draws a standard heatmap using matplotlib's `imshow`, with options for
12
+ gridlines, tick formatting, and embedded colourbars.
13
+
14
+ masked_heatmap(data, indicator, row_labels, col_labels, ...)
15
+ Draws a two-layer heatmap: a single-colour background and, on top of it,
16
+ the heatmap restricted to the cells flagged by a binary indicator table.
17
+
18
+ annotate_heatmap(im, data=None, valfmt=None, ...)
19
+ Adds text annotations to an existing heatmap image (AxesImage object),
20
+ with configurable formatting and colour thresholding.
21
+
22
+ Notes
23
+ -----
24
+ The base structure of the `heatmap` and `annotate_heatmap` functions is derived
25
+ from the example published in the official matplotlib gallery [1]_.
26
+
27
+ References
28
+ ----------
29
+ .. [1] Matplotlib contributors. "Creating annotated heatmaps." Matplotlib
30
+ Gallery. https://matplotlib.org/stable/gallery/images_contours_and_fields/image_annotated_heatmap.html
31
+ """
32
+
33
+ # modules
34
+ import numpy as np
35
+ import pandas as pd
36
+ import matplotlib
37
+ import matplotlib.pyplot as plt
38
+ from matplotlib.colors import ListedColormap
39
+ from matplotlib.patches import Rectangle
40
+ from plot_misc.utils.utils import _update_kwargs
41
+ from plot_misc.errors import (
42
+ is_type,
43
+ is_df,
44
+ InputValidationError,
45
+ )
46
+ from plot_misc.constants import Real
47
+ from typing import Any
48
+
49
+ # ~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
50
+ def heatmap(data:pd.DataFrame | np.ndarray, row_labels:list[str] | np.ndarray,
51
+ col_labels:list[str] | np.ndarray, grid_col:str='white',
52
+ grid_linestyle:str='-', grid_linewidth:float=3,
53
+ cbar_bool:bool=False, cbar_label:str="",
54
+ ax:plt.Axes | None = None,
55
+ figsize:tuple[float,float] | None = None,
56
+ grid_kw:dict[Any,Any] | None = None,
57
+ cbar_kw:dict[Any,Any] | None = None,
58
+ **kwargs:Any,
59
+ ) -> tuple[matplotlib.image.AxesImage,
60
+ matplotlib.colorbar.Colorbar]:
61
+ """
62
+ Plot a heatmap with row and column labels using matplotlib.
63
+
64
+ This function draws a heatmap using `imshow`, with options to configure
65
+ grid lines, colourbars, and axis labels. It accepts both NumPy arrays
66
+ and pandas DataFrames as input.
67
+
68
+ Parameters
69
+ ----------
70
+ data : `pd.DataFrame` or `np.array`
71
+ A 2D array of shape (M, N) containing the values to plot.
72
+ row_labels : `list` [`str`] or `np.ndarray`
73
+ A list or array of length M with the labels for the rows.
74
+ col_labels : `list` [`str`] or `np.ndarray`
75
+ A list or array of length N with the labels for the rows.
76
+ grid_col : `str`, default 'white'
77
+ The colour of the grid lines
78
+ grid_linestyle : `str`, default '-'
79
+ The linestyle of the grid lines
80
+ grid_linewidth : `float`, default 3
81
+ The width of the grid lines.
82
+ cbar_bool : `bool`, default `False`
83
+ If `True`, add a colourbar to the figure.
84
+ cbar_label : `str`, default " "
85
+ The label for the colorbar.
86
+ ax : `plt.Axes` or `None`, default None
87
+ A `matplotlib.axes.Axes` instance to which the heatmap is plotted. If
88
+ not provided, use current axes or create a new one.
89
+ figsize : `tuple` [`float`, `float`] or `None`, default `None`
90
+ Figure size in inches (width, height). Ignored if `ax` is provided.
91
+ grid_kw : `dict` [`str`,`any`] or `None`, default None
92
+ A dictionary with arguments to `matplotlib.Axes.grid`.
93
+ cbar_kw : `dict` [`str`, `any`] or `None`, default `None`
94
+ A dictionary with arguments to `matplotlib.Figure.colorbar`.
95
+ **kwargs : `any`
96
+ All other arguments are forwarded to `imshow`.
97
+
98
+ Returns
99
+ -------
100
+ im : `matplotlib.image.AxesImage`
101
+ The heatmap image object.
102
+ cbar : `matplotlib.colorbar.Colorbar` or `None`
103
+ The colourbar object if `cbar_bool` is `True`, otherwise `None`.
104
+
105
+ Notes
106
+ -----
107
+ The returned objects can be used to annotate the cells using for example
108
+ `annotate_heatmap`.
109
+
110
+ This function is adapted from the matplotlib gallery example [HM1]_.
111
+
112
+ References
113
+ ----------
114
+ .. [HM1] Matplotlib contributors. "Creating annotated heatmaps."
115
+ Matplotlib Gallery.
116
+ https://matplotlib.org/stable/gallery/images_contours_and_fields/image_annotated_heatmap.html
117
+ """
118
+ # check in put
119
+ is_type(data, (pd.DataFrame, np.ndarray))
120
+ is_type(row_labels, (list, np.ndarray))
121
+ is_type(col_labels, (list, np.ndarray))
122
+ is_type(grid_col, str)
123
+ is_type(grid_linestyle, str)
124
+ is_type(grid_linewidth, Real)
125
+ is_type(cbar_bool, bool)
126
+ is_type(cbar_label, str)
127
+ # create a axes if needed
128
+ if ax is None:
129
+ _, ax = plt.subplots(figsize=figsize)
130
+ else:
131
+ f = ax.figure
132
+ # check input
133
+ if isinstance(data, pd.DataFrame):
134
+ matrix = data.copy().to_numpy()
135
+ else:
136
+ matrix = data
137
+ # copy
138
+ row_lab = row_labels
139
+ col_lab = col_labels
140
+ # map None to dict
141
+ grid_kw = grid_kw or {}
142
+ cbar_kw = cbar_kw or {}
143
+ # ### Plot the heatmap
144
+ im = ax.imshow(matrix, **kwargs)
145
+ # Create colorbar
146
+ if cbar_bool:
147
+ # NOTE if the kwargs for colobar is extended use `_update_kwargs
148
+ cbar = ax.figure.colorbar(im, ax=ax, **cbar_kw)
149
+ cbar.ax.set_ylabel(cbar_label, rotation=-90, va="bottom")
150
+ else:
151
+ cbar = None
152
+ # Show all ticks and label them with the respective list entries.
153
+ ax.set_xticks(np.arange(matrix.shape[1]))
154
+ ax.set_xticklabels(col_lab)
155
+ ax.set_yticks(np.arange(matrix.shape[0]))
156
+ ax.set_yticklabels(row_lab)
157
+ # Let the horizontal axes labeling appear on top.
158
+ # ax.tick_params(top=True, bottom=False,
159
+ # labeltop=True, labelbottom=False)
160
+ # Rotate the tick labels and set their alignment.
161
+ plt.setp(ax.get_xticklabels(), rotation=30, ha="right",
162
+ rotation_mode="anchor")
163
+ # Turn spines off and create white grid.
164
+ ax.spines['top'].set_visible(False)
165
+ ax.spines['bottom'].set_visible(False)
166
+ ax.spines['right'].set_visible(False)
167
+ ax.spines['left'].set_visible(False)
168
+ # set tick marks
169
+ ax.set_xticks(np.arange(matrix.shape[1]+1)-.5, minor=True)
170
+ ax.set_yticks(np.arange(matrix.shape[0]+1)-.5, minor=True)
171
+ # grid
172
+ new_grid_kwargs = _update_kwargs(
173
+ update_dict=grid_kw, which="minor", color=grid_col,
174
+ linestyle=grid_linestyle, linewidth=grid_linewidth,
175
+ clip_on=False,)
176
+ ax.grid(**new_grid_kwargs)
177
+ ax.tick_params(which="minor", bottom=False, left=False)
178
+ # return stuff
179
+ return im, cbar
180
+
181
+ # ~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
182
+ def masked_heatmap(data:pd.DataFrame | np.ndarray,
183
+ indicator:pd.DataFrame | np.ndarray,
184
+ row_labels:list[str] | np.ndarray,
185
+ col_labels:list[str] | np.ndarray,
186
+ background_col:str='white', background_gridcol:str='white',
187
+ background_linestyle:str='-', background_linewidth:float=0.5,
188
+ background_zorder:Real = 1,
189
+ outline_col:str='black', outline_linestyle:str='-',
190
+ outline_linewidth:float=1.5, outline_zorder:Real = 2,
191
+ frame: bool=False,
192
+ cbar_bool:bool=False, cbar_label:str="",
193
+ ax:plt.Axes | None = None,
194
+ figsize:tuple[float,float] | None = None,
195
+ grid_kw:dict[Any,Any] | None = None,
196
+ cbar_kw:dict[Any,Any] | None = None,
197
+ background_kw:dict[Any,Any] | None = None,
198
+ outline_kw:dict[Any,Any] | None = None,
199
+ **kwargs: Any,
200
+ ) -> tuple[matplotlib.image.AxesImage,
201
+ matplotlib.colorbar.Colorbar]:
202
+ """
203
+ Plot a two-layer heatmap masked by a binary indicator table.
204
+
205
+ The function draws two layers. First a single-colour background covering
206
+ every cell (carrying an optional grid lattice). Second, the heatmap of
207
+ `data`, restricted to the cells where `indicator` equals 1.
208
+
209
+ Parameters
210
+ ----------
211
+ data : `pd.DataFrame` or `np.ndarray`
212
+ A 2D array of shape (M, N) containing the values to plot.
213
+ indicator : `pd.DataFrame` or `np.ndarray`
214
+ A binary (0/1, booleans accepted) array of the same shape as `data`.
215
+ Only cells equal to 1 are drawn and outlined.
216
+ row_labels : `list` [`str`] or `np.ndarray`
217
+ A list or array of length M with the labels for the rows.
218
+ col_labels : `list` [`str`] or `np.ndarray`
219
+ A list or array of length N with the labels for the columns.
220
+ background_col : `str`, default 'white'
221
+ The fill colour of the background layer.
222
+ background_gridcol : `str`, default 'white'
223
+ The colour of the background grid lattice lines.
224
+ background_linestyle : `str`, default '-'
225
+ The linestyle of the background grid lattice.
226
+ background_linewidth : `float`, default 0.5
227
+ The width of the background grid lattice. Set to 0 to suppress it.
228
+ background_zorder : `int`, `float` default `1`
229
+ The draw order of the background grid lattice.
230
+ outline_col : `str`, default 'black'
231
+ The edge colour of the per-cell outlines drawn on `indicator == 1`
232
+ cells.
233
+ outline_linestyle : `str`, default '-'
234
+ The linestyle of the per-cell outlines.
235
+ outline_linewidth : `float`, default 1.5
236
+ The width of the per-cell outlines. Set to 0 to suppress them.
237
+ outline_zorder : `int`, `float`, default `2`
238
+ The draw order of the per-cell outlines.
239
+ frame : `bool`, default `False`
240
+ Whether to plot the spines.
241
+ cbar_bool : `bool`, default `False`
242
+ If `True`, add a colourbar (built from the masked heatmap layer).
243
+ cbar_label : `str`, default ""
244
+ The label for the colourbar.
245
+ ax : `plt.Axes` or `None`, default `None`
246
+ A `matplotlib.axes.Axes` instance to draw on. If `None`, a new figure
247
+ and axes are created.
248
+ figsize : `tuple` [`float`, `float`] or `None`, default `None`
249
+ Figure size in inches (width, height). Ignored if `ax` is provided.
250
+ grid_kw : `dict` [`str`, `any`] or `None`, default `None`
251
+ Additional arguments forwarded to `matplotlib.Axes.grid` for the
252
+ background lattice.
253
+ outline_kw : `dict` [`str`, `any`] or `None`, default `None`
254
+ Additional arguments forwarded to each `matplotlib.patches.Rectangle`
255
+ outline. Outlines default to `clip_on=False` so the borders of cells on
256
+ the matrix boundary are not clipped by the axes edge; pass
257
+ `{'clip_on': True}` to restore clipping.
258
+ cbar_kw : `dict` [`str`, `any`] or `None`, default `None`
259
+ A dictionary with arguments to `matplotlib.Figure.colorbar`.
260
+ background_kw : `dict` [`str`, `any`] or `None`, default `None`,
261
+ A dictionary with arguments to `heatmap.heatmap`.
262
+ **kwargs : `any`,
263
+ All other arguments passed to `masking ax.imshow`.
264
+
265
+ Returns
266
+ -------
267
+ im : `matplotlib.image.AxesImage`
268
+ The masked (foreground) heatmap image object.
269
+ cbar : `matplotlib.colorbar.Colorbar` or `None`
270
+ The colourbar object if `cbar_bool` is `True`, otherwise `None`.
271
+
272
+ Notes
273
+ -----
274
+ The masking is achieved by a separate imshow call setting the cells to
275
+ transparent, revealing the background.
276
+
277
+ The returned `im` mirrors the contract of `heatmap` and can therefore be
278
+ annotated using `annotate_heatmap`.
279
+ """
280
+ # create an axes if needed
281
+ if ax is None:
282
+ _, ax = plt.subplots(figsize=figsize)
283
+ else:
284
+ f = ax.figure
285
+ # check input types
286
+ is_type(data, (pd.DataFrame, np.ndarray))
287
+ is_type(indicator, (pd.DataFrame, np.ndarray))
288
+ _ = [is_type(k, (dict, type(None))) for k in\
289
+ (grid_kw, cbar_kw, outline_kw, background_kw)]
290
+ # the indicator must match the data shape exactly (full 2D shape, not just
291
+ # the row count)
292
+ if np.shape(data) != np.shape(indicator):
293
+ raise InputValidationError(
294
+ f"`indicator` shape {np.shape(indicator)} does not match `data` "
295
+ f"shape {np.shape(data)}."
296
+ )
297
+ # coerce the data and indicator to numpy arrays
298
+ if isinstance(data, pd.DataFrame):
299
+ matrix = data.copy().to_numpy()
300
+ else:
301
+ matrix = data
302
+ if isinstance(indicator, pd.DataFrame):
303
+ flag = indicator.copy().to_numpy()
304
+ else:
305
+ flag = indicator
306
+ # flag should only contain 0 and 1
307
+ unique_flags = set(np.unique(flag).tolist())
308
+ if not unique_flags.issubset({0, 1}):
309
+ raise InputValidationError(
310
+ f"`indicator` must only contain binary (0/1) values, got "
311
+ f"{sorted(unique_flags)}."
312
+ )
313
+ # setup the kwargs None to dict
314
+ background_kw = background_kw or {}
315
+ # masking_kw = masking_kw or {}
316
+ grid_kw = grid_kw or {}
317
+ cbar_kw = cbar_kw or {}
318
+ outline_kw = outline_kw or {}
319
+ outline_kw = _update_kwargs(update_dict=outline_kw,
320
+ zorder=outline_zorder)
321
+ # the background grid lattice carries the background draw order
322
+ grid_kw = _update_kwargs(update_dict=grid_kw, zorder=background_zorder)
323
+ # ### Layer 1: a single-colour background covering every cell. Reusing
324
+ # `heatmap` keeps a single source of truth for the tick, label, spine and
325
+ # grid (lattice) cosmetics.
326
+ background = np.zeros_like(matrix, dtype=float)
327
+ layer1_kwargs = _update_kwargs(
328
+ update_dict=background_kw,
329
+ data=background, row_labels=row_labels, col_labels=col_labels,
330
+ grid_col=background_gridcol, grid_linestyle=background_linestyle,
331
+ grid_linewidth=background_linewidth, cbar_bool=False, ax=ax,
332
+ grid_kw=grid_kw, cmap=ListedColormap([background_col]),
333
+ )
334
+ heatmap(**layer1_kwargs, )
335
+ # ### Layer 2: the heatmap, masked so only `indicator == 1` cells are drawn.
336
+ masked = np.ma.masked_where(flag == 0, matrix)
337
+ # creating an alpha matrix.
338
+ # user_alpha = masking_kw.pop('alpha', 1.0)
339
+ user_alpha = kwargs.pop('alpha', 1.0)
340
+ alpha = (flag == 1).astype(float) * np.asarray(user_alpha, dtype=float)
341
+ layer2_kwargs = _update_kwargs(
342
+ update_dict=kwargs, alpha=alpha,
343
+ )
344
+ im = ax.imshow(masked, **layer2_kwargs)
345
+ # Create colorbar from the foreground (masked) layer
346
+ if cbar_bool:
347
+ cbar = ax.figure.colorbar(im, ax=ax, **cbar_kw)
348
+ cbar.ax.set_ylabel(cbar_label, rotation=-90, va="bottom")
349
+ else:
350
+ cbar = None
351
+ # ### Outline each `indicator == 1` cell. Zero cells get no patch, so they
352
+ # carry no outline; a zero `outline_linewidth` hides the borders.
353
+ rect_kw = _update_kwargs(update_dict=outline_kw, facecolor='none',
354
+ edgecolor=outline_col,
355
+ linestyle=outline_linestyle,
356
+ linewidth=outline_linewidth,
357
+ clip_on=False,
358
+ )
359
+ rows, cols = np.where(flag == 1)
360
+ # NOTE the 0.5 and 1.0 are imshow fixed convention and should be hardcoded
361
+ for i, j in zip(rows, cols):
362
+ ax.add_patch(Rectangle((j-.5, i-.5), 1, 1, **rect_kw))
363
+ # Show the spines
364
+ if frame:
365
+ for spine in ax.spines.values():
366
+ spine.set_visible(True)
367
+ # return stuff
368
+ return im, cbar
369
+
370
+ # ~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
371
+ def annotate_heatmap(
372
+ im:plt.Axes.imshow,
373
+ data:pd.DataFrame | np.ndarray | None = None,
374
+ valfmt:str | matplotlib.ticker.Formatter | None = None,
375
+ textcolors:tuple[str,str] | list[str,str]=("black","white"),
376
+ threshold: float | None = None,
377
+ **kwargs:Any,
378
+ ) -> list[plt.Text]:
379
+ """
380
+ Annotate each cell in a heatmap image with its value.
381
+
382
+ This function adds text annotations to an existing `AxesImage` object,
383
+ such as those created by the `heatmap` function. The text colour may
384
+ be adjusted dynamically based on a threshold value and the image’s colour
385
+ map.
386
+
387
+ Parameters
388
+ ----------
389
+ im : `plt.Axes.imshow`
390
+ The AxesImage to be labeled.
391
+ data : `pd.DataFrame`, `np.array`, or `None`, default `Nonetype`
392
+ A 2D numpy array of shape (M, N). If `None`, the function uses the
393
+ array embedded in `im`.
394
+ valfmt : `str`, `matplotlib.ticker.Formatter` or `None`, default `None`
395
+ The format of the annotations inside the heatmap. This should either
396
+ use the string format method, e.g. "$ {x:.2f}" - (note the `x` is needs
397
+ to be included to represent the numerical), or be a
398
+ `matplotlib.ticker.Formatter`.
399
+ textcolors : `list` or `tuple` [`str`, `str`], default `('black', 'white')`
400
+ A pair of colors. The first is used for values below a threshold,
401
+ the second for those above.
402
+ threshold : `float` or `None`, default `None`
403
+ The absolute value in data units according to which the colors from
404
+ textcolors are applied. If None (the default) uses the middle of the
405
+ colormap as separation.
406
+ **kwargs : `any`
407
+ All other arguments are forwarded to each call to `text` used to create
408
+ the text labels.
409
+
410
+ Returns
411
+ -------
412
+ texts : `list` of `matplotlib.text.Text`
413
+ A list of text annotation objects added to the heatmap.
414
+ """
415
+
416
+ # mapping data to matrix
417
+ values = im.get_array()
418
+ # masked cells (e.g. from `masked_heatmap`) are not drawn and must not be
419
+ # annotated; `getmaskarray` yields a full boolean mask for masked arrays and
420
+ # an all-False mask for plain arrays, leaving the unmasked path unchanged
421
+ mask = np.ma.getmaskarray(values)
422
+ if data is None:
423
+ matrix = im.get_array()
424
+ elif isinstance(data, pd.DataFrame):
425
+ matrix = data.copy().to_numpy()
426
+ else:
427
+ matrix = data
428
+ # Compare raw data values against the raw threshold
429
+ # This bypases any kind of value normalisation - which we should skip
430
+ # because the string are not normalised only the values
431
+ if threshold is None:
432
+ try:
433
+ threshold = values.max() / 2.
434
+ except np.core._exceptions.UFuncTypeError:
435
+ threshold = None
436
+ # Set default alignment to center, but allow it to be
437
+ # overwritten by text_kw.
438
+ kw = _update_kwargs(update_dict=kwargs,
439
+ horizontalalignment="center",
440
+ verticalalignment="center",
441
+ )
442
+ # Get the formatter in case a string is supplied
443
+ if valfmt is not None:
444
+ if isinstance(valfmt, str):
445
+ valfmt = matplotlib.ticker.StrMethodFormatter(valfmt)
446
+ # Loop over the data and create a `Text` for each "pixel".
447
+ # Change the text's color depending on the data.
448
+ texts = []
449
+ for i in range(matrix.shape[0]):
450
+ for j in range(matrix.shape[1]):
451
+ # skip masked cells, which carry no drawn value to annotate
452
+ if mask[i, j]:
453
+ continue
454
+ # only run if threshold exists
455
+ if threshold is not None:
456
+ kw.update(color=textcolors[int(abs(values[i, j]) >= threshold)])
457
+ # format text or not
458
+ if valfmt is not None:
459
+ # NOTE text takes x, y - whereas values takes rows (y), col (x)
460
+ text = im.axes.text(j, i, valfmt(matrix[i, j], None), **kw)
461
+ else:
462
+ text = im.axes.text(j, i, matrix[i,j], **kw)
463
+ texts.append(text)
464
+ # returning stuff
465
+ return texts
466
+