plot-misc 2.2.1__tar.gz → 2.2.2__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 (45) hide show
  1. {plot_misc-2.2.1/plot_misc.egg-info → plot_misc-2.2.2}/PKG-INFO +22 -8
  2. {plot_misc-2.2.1 → plot_misc-2.2.2}/README.md +21 -7
  3. plot_misc-2.2.2/plot_misc/_version.py +1 -0
  4. {plot_misc-2.2.1 → plot_misc-2.2.2}/plot_misc/constants.py +6 -0
  5. {plot_misc-2.2.1 → plot_misc-2.2.2}/plot_misc/example_data/examples.py +48 -0
  6. plot_misc-2.2.2/plot_misc/heatmap.py +467 -0
  7. {plot_misc-2.2.1 → plot_misc-2.2.2}/plot_misc/utils/utils.py +206 -95
  8. {plot_misc-2.2.1 → plot_misc-2.2.2}/plot_misc/volcano.py +5 -5
  9. {plot_misc-2.2.1 → plot_misc-2.2.2/plot_misc.egg-info}/PKG-INFO +22 -8
  10. {plot_misc-2.2.1 → plot_misc-2.2.2}/pyproject.toml +1 -1
  11. plot_misc-2.2.1/plot_misc/_version.py +0 -1
  12. plot_misc-2.2.1/plot_misc/heatmap.py +0 -254
  13. {plot_misc-2.2.1 → plot_misc-2.2.2}/LICENSE +0 -0
  14. {plot_misc-2.2.1 → plot_misc-2.2.2}/MANIFEST.in +0 -0
  15. {plot_misc-2.2.1 → plot_misc-2.2.2}/plot_misc/__init__.py +0 -0
  16. {plot_misc-2.2.1 → plot_misc-2.2.2}/plot_misc/barchart.py +0 -0
  17. {plot_misc-2.2.1 → plot_misc-2.2.2}/plot_misc/errors.py +0 -0
  18. {plot_misc-2.2.1 → plot_misc-2.2.2}/plot_misc/example_data/__init__.py +0 -0
  19. {plot_misc-2.2.1 → plot_misc-2.2.2}/plot_misc/example_data/example_datasets/bar_points.tsv.gz +0 -0
  20. {plot_misc-2.2.1 → plot_misc-2.2.2}/plot_misc/example_data/example_datasets/barchart.tsv.gz +0 -0
  21. {plot_misc-2.2.1 → plot_misc-2.2.2}/plot_misc/example_data/example_datasets/calibration_bins.tsv.gz +0 -0
  22. {plot_misc-2.2.1 → plot_misc-2.2.2}/plot_misc/example_data/example_datasets/calibration_data.tsv.gz +0 -0
  23. {plot_misc-2.2.1 → plot_misc-2.2.2}/plot_misc/example_data/example_datasets/forest_data.tsv.gz +0 -0
  24. {plot_misc-2.2.1 → plot_misc-2.2.2}/plot_misc/example_data/example_datasets/group_bar.tsv.gz +0 -0
  25. {plot_misc-2.2.1 → plot_misc-2.2.2}/plot_misc/example_data/example_datasets/heatmap_data.tsv.gz +0 -0
  26. {plot_misc-2.2.1 → plot_misc-2.2.2}/plot_misc/example_data/example_datasets/incidence_matrix_data.tsv.gz +0 -0
  27. {plot_misc-2.2.1 → plot_misc-2.2.2}/plot_misc/example_data/example_datasets/lollipop_data.tsv.gz +0 -0
  28. {plot_misc-2.2.1 → plot_misc-2.2.2}/plot_misc/example_data/example_datasets/mace_associations.tsv.gz +0 -0
  29. {plot_misc-2.2.1 → plot_misc-2.2.2}/plot_misc/example_data/example_datasets/net_benefit.tsv.gz +0 -0
  30. {plot_misc-2.2.1 → plot_misc-2.2.2}/plot_misc/example_data/example_datasets/volcano.tsv.gz +0 -0
  31. {plot_misc-2.2.1 → plot_misc-2.2.2}/plot_misc/forest.py +0 -0
  32. {plot_misc-2.2.1 → plot_misc-2.2.2}/plot_misc/incidencematrix.py +0 -0
  33. {plot_misc-2.2.1 → plot_misc-2.2.2}/plot_misc/machine_learning.py +0 -0
  34. {plot_misc-2.2.1 → plot_misc-2.2.2}/plot_misc/piechart.py +0 -0
  35. {plot_misc-2.2.1 → plot_misc-2.2.2}/plot_misc/survival.py +0 -0
  36. {plot_misc-2.2.1 → plot_misc-2.2.2}/plot_misc/utils/__init__.py +0 -0
  37. {plot_misc-2.2.1 → plot_misc-2.2.2}/plot_misc/utils/colour.py +0 -0
  38. {plot_misc-2.2.1 → plot_misc-2.2.2}/plot_misc/utils/formatting.py +0 -0
  39. {plot_misc-2.2.1 → plot_misc-2.2.2}/plot_misc.egg-info/SOURCES.txt +0 -0
  40. {plot_misc-2.2.1 → plot_misc-2.2.2}/plot_misc.egg-info/dependency_links.txt +0 -0
  41. {plot_misc-2.2.1 → plot_misc-2.2.2}/plot_misc.egg-info/requires.txt +0 -0
  42. {plot_misc-2.2.1 → plot_misc-2.2.2}/plot_misc.egg-info/top_level.txt +0 -0
  43. {plot_misc-2.2.1 → plot_misc-2.2.2}/requirements-dev.txt +0 -0
  44. {plot_misc-2.2.1 → plot_misc-2.2.2}/requirements.txt +0 -0
  45. {plot_misc-2.2.1 → plot_misc-2.2.2}/setup.cfg +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: plot-misc
3
- Version: 2.2.1
3
+ Version: 2.2.2
4
4
  Summary: Various plotting templates built on top of matplotlib
5
5
  Author-email: A Floriaan Schmidt <floriaanschmidt@gmail.com>
6
6
  License-Expression: GPL-3.0-or-later
@@ -48,15 +48,29 @@ Dynamic: license-file
48
48
  <img src="https://schmidtaf.gitlab.io/plot-misc/_images/icon.png" alt="plot-misc icon" width="250"/>
49
49
 
50
50
  # A collection of plotting functions
51
- __version__: `2.2.1`
51
+ __version__: `2.2.2`
52
52
 
53
53
  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.
54
+ The functions describe plotting archetypes intended to set up light-touch,
55
+ illustrations that can be customised using the standard matplotlib interface
56
+ via axes and figures.
57
+ Because the implementation is matplotlib-first, the API is consistent with
58
+ matplotlib conventions, and users already familiar with the library will find
59
+ the learning curve minimal.
60
+
61
+ The functionality is geared towards illustrations commonly used in biomedical
62
+ research:
63
+
64
+ * Bar charts
65
+ * Bubble charts
66
+ * Forest plots (with optional side-tables)
67
+ * Heatmaps (with optional annotations)
68
+ * Incidence matrix plots
69
+ * Machine learning plots (calibration, feature importance, net benefit)
70
+ * Pie charts
71
+ * Survival plots (with optional survival table)
72
+ * Tree/compatibility plots
73
+ * Volcano plots
60
74
 
61
75
  Please consult the **[documentation](https://SchmidtAF.gitlab.io/plot-misc/)**
62
76
  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.2.2`
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 @@
1
+ __version__ = '2.2.2'
@@ -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):
@@ -0,0 +1,467 @@
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."
30
+ Matplotlib Gallery.
31
+ https://matplotlib.org/stable/gallery/images_contours_and_fields/image_annotated_heatmap.html
32
+ """
33
+
34
+ # modules
35
+ import numpy as np
36
+ import pandas as pd
37
+ import matplotlib
38
+ import matplotlib.pyplot as plt
39
+ from matplotlib.colors import ListedColormap
40
+ from matplotlib.patches import Rectangle
41
+ from plot_misc.utils.utils import _update_kwargs
42
+ from plot_misc.errors import (
43
+ is_type,
44
+ is_df,
45
+ InputValidationError,
46
+ )
47
+ from plot_misc.constants import Real
48
+ from typing import Any
49
+
50
+ # ~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
51
+ def heatmap(data:pd.DataFrame | np.ndarray, row_labels:list[str] | np.ndarray,
52
+ col_labels:list[str] | np.ndarray, grid_col:str='white',
53
+ grid_linestyle:str='-', grid_linewidth:float=3,
54
+ cbar_bool:bool=False, cbar_label:str="",
55
+ ax:plt.Axes | None = None,
56
+ figsize:tuple[float,float] | None = None,
57
+ grid_kw:dict[Any,Any] | None = None,
58
+ cbar_kw:dict[Any,Any] | None = None,
59
+ **kwargs:Any,
60
+ ) -> tuple[matplotlib.image.AxesImage,
61
+ matplotlib.colorbar.Colorbar]:
62
+ """
63
+ Plot a heatmap with row and column labels using matplotlib.
64
+
65
+ This function draws a heatmap using `imshow`, with options to configure
66
+ grid lines, colourbars, and axis labels. It accepts both NumPy arrays
67
+ and pandas DataFrames as input.
68
+
69
+ Parameters
70
+ ----------
71
+ data : `pd.DataFrame` or `np.array`
72
+ A 2D array of shape (M, N) containing the values to plot.
73
+ row_labels : `list` [`str`] or `np.ndarray`
74
+ A list or array of length M with the labels for the rows.
75
+ col_labels : `list` [`str`] or `np.ndarray`
76
+ A list or array of length N with the labels for the rows.
77
+ grid_col : `str`, default 'white'
78
+ The colour of the grid lines
79
+ grid_linestyle : `str`, default '-'
80
+ The linestyle of the grid lines
81
+ grid_linewidth : `float`, default 3
82
+ The width of the grid lines.
83
+ cbar_bool : `bool`, default `False`
84
+ If `True`, add a colourbar to the figure.
85
+ cbar_label : `str`, default " "
86
+ The label for the colorbar.
87
+ ax : `plt.Axes` or `None`, default None
88
+ A `matplotlib.axes.Axes` instance to which the heatmap is plotted. If
89
+ not provided, use current axes or create a new one.
90
+ figsize : `tuple` [`float`, `float`] or `None`, default `None`
91
+ Figure size in inches (width, height). Ignored if `ax` is provided.
92
+ grid_kw : `dict` [`str`,`any`] or `None`, default None
93
+ A dictionary with arguments to `matplotlib.Axes.grid`.
94
+ cbar_kw : `dict` [`str`, `any`] or `None`, default `None`
95
+ A dictionary with arguments to `matplotlib.Figure.colorbar`.
96
+ **kwargs : `any`
97
+ All other arguments are forwarded to `imshow`.
98
+
99
+ Returns
100
+ -------
101
+ im : `matplotlib.image.AxesImage`
102
+ The heatmap image object.
103
+ cbar : `matplotlib.colorbar.Colorbar` or `None`
104
+ The colourbar object if `cbar_bool` is `True`, otherwise `None`.
105
+
106
+ Notes
107
+ -----
108
+ The returned objects can be used to annotate the cells using for example
109
+ `annotate_heatmap`.
110
+
111
+ This function is adapted from the matplotlib gallery example [HM1]_.
112
+
113
+ References
114
+ ----------
115
+ .. [HM1] Matplotlib contributors. "Creating annotated heatmaps."
116
+ Matplotlib Gallery.
117
+ https://matplotlib.org/stable/gallery/images_contours_and_fields/image_annotated_heatmap.html
118
+ """
119
+ # check in put
120
+ is_type(data, (pd.DataFrame, np.ndarray))
121
+ is_type(row_labels, (list, np.ndarray))
122
+ is_type(col_labels, (list, np.ndarray))
123
+ is_type(grid_col, str)
124
+ is_type(grid_linestyle, str)
125
+ is_type(grid_linewidth, Real)
126
+ is_type(cbar_bool, bool)
127
+ is_type(cbar_label, str)
128
+ # create a axes if needed
129
+ if ax is None:
130
+ _, ax = plt.subplots(figsize=figsize)
131
+ else:
132
+ f = ax.figure
133
+ # check input
134
+ if isinstance(data, pd.DataFrame):
135
+ matrix = data.copy().to_numpy()
136
+ else:
137
+ matrix = data
138
+ # copy
139
+ row_lab = row_labels
140
+ col_lab = col_labels
141
+ # map None to dict
142
+ grid_kw = grid_kw or {}
143
+ cbar_kw = cbar_kw or {}
144
+ # ### Plot the heatmap
145
+ im = ax.imshow(matrix, **kwargs)
146
+ # Create colorbar
147
+ if cbar_bool:
148
+ # NOTE if the kwargs for colobar is extended use `_update_kwargs
149
+ cbar = ax.figure.colorbar(im, ax=ax, **cbar_kw)
150
+ cbar.ax.set_ylabel(cbar_label, rotation=-90, va="bottom")
151
+ else:
152
+ cbar = None
153
+ # Show all ticks and label them with the respective list entries.
154
+ ax.set_xticks(np.arange(matrix.shape[1]))
155
+ ax.set_xticklabels(col_lab)
156
+ ax.set_yticks(np.arange(matrix.shape[0]))
157
+ ax.set_yticklabels(row_lab)
158
+ # Let the horizontal axes labeling appear on top.
159
+ # ax.tick_params(top=True, bottom=False,
160
+ # labeltop=True, labelbottom=False)
161
+ # Rotate the tick labels and set their alignment.
162
+ plt.setp(ax.get_xticklabels(), rotation=30, ha="right",
163
+ rotation_mode="anchor")
164
+ # Turn spines off and create white grid.
165
+ ax.spines['top'].set_visible(False)
166
+ ax.spines['bottom'].set_visible(False)
167
+ ax.spines['right'].set_visible(False)
168
+ ax.spines['left'].set_visible(False)
169
+ # set tick marks
170
+ ax.set_xticks(np.arange(matrix.shape[1]+1)-.5, minor=True)
171
+ ax.set_yticks(np.arange(matrix.shape[0]+1)-.5, minor=True)
172
+ # grid
173
+ new_grid_kwargs = _update_kwargs(
174
+ update_dict=grid_kw, which="minor", color=grid_col,
175
+ linestyle=grid_linestyle, linewidth=grid_linewidth,
176
+ clip_on=False,)
177
+ ax.grid(**new_grid_kwargs)
178
+ ax.tick_params(which="minor", bottom=False, left=False)
179
+ # return stuff
180
+ return im, cbar
181
+
182
+ # ~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
183
+ def masked_heatmap(data:pd.DataFrame | np.ndarray,
184
+ indicator:pd.DataFrame | np.ndarray,
185
+ row_labels:list[str] | np.ndarray,
186
+ col_labels:list[str] | np.ndarray,
187
+ background_col:str='white', background_gridcol:str='white',
188
+ background_linestyle:str='-', background_linewidth:float=0.5,
189
+ background_zorder:Real = 1,
190
+ outline_col:str='black', outline_linestyle:str='-',
191
+ outline_linewidth:float=1.5, outline_zorder:Real = 2,
192
+ frame: bool=False,
193
+ cbar_bool:bool=False, cbar_label:str="",
194
+ ax:plt.Axes | None = None,
195
+ figsize:tuple[float,float] | None = None,
196
+ grid_kw:dict[Any,Any] | None = None,
197
+ cbar_kw:dict[Any,Any] | None = None,
198
+ background_kw:dict[Any,Any] | None = None,
199
+ outline_kw:dict[Any,Any] | None = None,
200
+ **kwargs: Any,
201
+ ) -> tuple[matplotlib.image.AxesImage,
202
+ matplotlib.colorbar.Colorbar]:
203
+ """
204
+ Plot a two-layer heatmap masked by a binary indicator table.
205
+
206
+ The function draws two layers. First a single-colour background covering
207
+ every cell (carrying an optional grid lattice). Second, the heatmap of
208
+ `data`, restricted to the cells where `indicator` equals 1.
209
+
210
+ Parameters
211
+ ----------
212
+ data : `pd.DataFrame` or `np.ndarray`
213
+ A 2D array of shape (M, N) containing the values to plot.
214
+ indicator : `pd.DataFrame` or `np.ndarray`
215
+ A binary (0/1, booleans accepted) array of the same shape as `data`.
216
+ Only cells equal to 1 are drawn and outlined.
217
+ row_labels : `list` [`str`] or `np.ndarray`
218
+ A list or array of length M with the labels for the rows.
219
+ col_labels : `list` [`str`] or `np.ndarray`
220
+ A list or array of length N with the labels for the columns.
221
+ background_col : `str`, default 'white'
222
+ The fill colour of the background layer.
223
+ background_gridcol : `str`, default 'white'
224
+ The colour of the background grid lattice lines.
225
+ background_linestyle : `str`, default '-'
226
+ The linestyle of the background grid lattice.
227
+ background_linewidth : `float`, default 0.5
228
+ The width of the background grid lattice. Set to 0 to suppress it.
229
+ background_zorder : `int`, `float` default `1`
230
+ The draw order of the background grid lattice.
231
+ outline_col : `str`, default 'black'
232
+ The edge colour of the per-cell outlines drawn on `indicator == 1`
233
+ cells.
234
+ outline_linestyle : `str`, default '-'
235
+ The linestyle of the per-cell outlines.
236
+ outline_linewidth : `float`, default 1.5
237
+ The width of the per-cell outlines. Set to 0 to suppress them.
238
+ outline_zorder : `int`, `float`, default `2`
239
+ The draw order of the per-cell outlines.
240
+ frame : `bool`, default `False`
241
+ Whether to plot the spines.
242
+ cbar_bool : `bool`, default `False`
243
+ If `True`, add a colourbar (built from the masked heatmap layer).
244
+ cbar_label : `str`, default ""
245
+ The label for the colourbar.
246
+ ax : `plt.Axes` or `None`, default `None`
247
+ A `matplotlib.axes.Axes` instance to draw on. If `None`, a new figure
248
+ and axes are created.
249
+ figsize : `tuple` [`float`, `float`] or `None`, default `None`
250
+ Figure size in inches (width, height). Ignored if `ax` is provided.
251
+ grid_kw : `dict` [`str`, `any`] or `None`, default `None`
252
+ Additional arguments forwarded to `matplotlib.Axes.grid` for the
253
+ background lattice.
254
+ outline_kw : `dict` [`str`, `any`] or `None`, default `None`
255
+ Additional arguments forwarded to each `matplotlib.patches.Rectangle`
256
+ outline. Outlines default to `clip_on=False` so the borders of cells on
257
+ the matrix boundary are not clipped by the axes edge; pass
258
+ `{'clip_on': True}` to restore clipping.
259
+ cbar_kw : `dict` [`str`, `any`] or `None`, default `None`
260
+ A dictionary with arguments to `matplotlib.Figure.colorbar`.
261
+ background_kw : `dict` [`str`, `any`] or `None`, default `None`,
262
+ A dictionary with arguments to `heatmap.heatmap`.
263
+ **kwargs : `any`,
264
+ All other arguments passed to `masking ax.imshow`.
265
+
266
+ Returns
267
+ -------
268
+ im : `matplotlib.image.AxesImage`
269
+ The masked (foreground) heatmap image object.
270
+ cbar : `matplotlib.colorbar.Colorbar` or `None`
271
+ The colourbar object if `cbar_bool` is `True`, otherwise `None`.
272
+
273
+ Notes
274
+ -----
275
+ The masking is achieved by a separate imshow call setting the cells to
276
+ transparent, revealing the background.
277
+
278
+ The returned `im` mirrors the contract of `heatmap` and can therefore be
279
+ annotated using `annotate_heatmap`.
280
+ """
281
+ # create an axes if needed
282
+ if ax is None:
283
+ _, ax = plt.subplots(figsize=figsize)
284
+ else:
285
+ f = ax.figure
286
+ # check input types
287
+ is_type(data, (pd.DataFrame, np.ndarray))
288
+ is_type(indicator, (pd.DataFrame, np.ndarray))
289
+ _ = [is_type(k, (dict, type(None))) for k in\
290
+ (grid_kw, cbar_kw, outline_kw, background_kw)]
291
+ # the indicator must match the data shape exactly (full 2D shape, not just
292
+ # the row count)
293
+ if np.shape(data) != np.shape(indicator):
294
+ raise InputValidationError(
295
+ f"`indicator` shape {np.shape(indicator)} does not match `data` "
296
+ f"shape {np.shape(data)}."
297
+ )
298
+ # coerce the data and indicator to numpy arrays
299
+ if isinstance(data, pd.DataFrame):
300
+ matrix = data.copy().to_numpy()
301
+ else:
302
+ matrix = data
303
+ if isinstance(indicator, pd.DataFrame):
304
+ flag = indicator.copy().to_numpy()
305
+ else:
306
+ flag = indicator
307
+ # flag should only contain 0 and 1
308
+ unique_flags = set(np.unique(flag).tolist())
309
+ if not unique_flags.issubset({0, 1}):
310
+ raise InputValidationError(
311
+ f"`indicator` must only contain binary (0/1) values, got "
312
+ f"{sorted(unique_flags)}."
313
+ )
314
+ # setup the kwargs None to dict
315
+ background_kw = background_kw or {}
316
+ # masking_kw = masking_kw or {}
317
+ grid_kw = grid_kw or {}
318
+ cbar_kw = cbar_kw or {}
319
+ outline_kw = outline_kw or {}
320
+ outline_kw = _update_kwargs(update_dict=outline_kw,
321
+ zorder=outline_zorder)
322
+ # the background grid lattice carries the background draw order
323
+ grid_kw = _update_kwargs(update_dict=grid_kw, zorder=background_zorder)
324
+ # ### Layer 1: a single-colour background covering every cell. Reusing
325
+ # `heatmap` keeps a single source of truth for the tick, label, spine and
326
+ # grid (lattice) cosmetics.
327
+ background = np.zeros_like(matrix, dtype=float)
328
+ layer1_kwargs = _update_kwargs(
329
+ update_dict=background_kw,
330
+ data=background, row_labels=row_labels, col_labels=col_labels,
331
+ grid_col=background_gridcol, grid_linestyle=background_linestyle,
332
+ grid_linewidth=background_linewidth, cbar_bool=False, ax=ax,
333
+ grid_kw=grid_kw, cmap=ListedColormap([background_col]),
334
+ )
335
+ heatmap(**layer1_kwargs, )
336
+ # ### Layer 2: the heatmap, masked so only `indicator == 1` cells are drawn.
337
+ masked = np.ma.masked_where(flag == 0, matrix)
338
+ # creating an alpha matrix.
339
+ # user_alpha = masking_kw.pop('alpha', 1.0)
340
+ user_alpha = kwargs.pop('alpha', 1.0)
341
+ alpha = (flag == 1).astype(float) * np.asarray(user_alpha, dtype=float)
342
+ layer2_kwargs = _update_kwargs(
343
+ update_dict=kwargs, alpha=alpha,
344
+ )
345
+ im = ax.imshow(masked, **layer2_kwargs)
346
+ # Create colorbar from the foreground (masked) layer
347
+ if cbar_bool:
348
+ cbar = ax.figure.colorbar(im, ax=ax, **cbar_kw)
349
+ cbar.ax.set_ylabel(cbar_label, rotation=-90, va="bottom")
350
+ else:
351
+ cbar = None
352
+ # ### Outline each `indicator == 1` cell. Zero cells get no patch, so they
353
+ # carry no outline; a zero `outline_linewidth` hides the borders.
354
+ rect_kw = _update_kwargs(update_dict=outline_kw, facecolor='none',
355
+ edgecolor=outline_col,
356
+ linestyle=outline_linestyle,
357
+ linewidth=outline_linewidth,
358
+ clip_on=False,
359
+ )
360
+ rows, cols = np.where(flag == 1)
361
+ # NOTE the 0.5 and 1.0 are imshow fixed convention and should be hardcoded
362
+ for i, j in zip(rows, cols):
363
+ ax.add_patch(Rectangle((j-.5, i-.5), 1, 1, **rect_kw))
364
+ # Show the spines
365
+ if frame:
366
+ for spine in ax.spines.values():
367
+ spine.set_visible(True)
368
+ # return stuff
369
+ return im, cbar
370
+
371
+ # ~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
372
+ def annotate_heatmap(
373
+ im:plt.Axes.imshow,
374
+ data:pd.DataFrame | np.ndarray | None = None,
375
+ valfmt:str | matplotlib.ticker.Formatter | None = None,
376
+ textcolors:tuple[str,str] | list[str,str]=("black","white"),
377
+ threshold: float | None = None,
378
+ **kwargs:Any,
379
+ ) -> list[plt.Text]:
380
+ """
381
+ Annotate each cell in a heatmap image with its value.
382
+
383
+ This function adds text annotations to an existing `AxesImage` object,
384
+ such as those created by the `heatmap` function. The text colour may
385
+ be adjusted dynamically based on a threshold value and the image’s colour
386
+ map.
387
+
388
+ Parameters
389
+ ----------
390
+ im : `plt.Axes.imshow`
391
+ The AxesImage to be labeled.
392
+ data : `pd.DataFrame`, `np.array`, or `None`, default `Nonetype`
393
+ A 2D numpy array of shape (M, N). If `None`, the function uses the
394
+ array embedded in `im`.
395
+ valfmt : `str`, `matplotlib.ticker.Formatter` or `None`, default `None`
396
+ The format of the annotations inside the heatmap. This should either
397
+ use the string format method, e.g. "$ {x:.2f}" - (note the `x` is needs
398
+ to be included to represent the numerical), or be a
399
+ `matplotlib.ticker.Formatter`.
400
+ textcolors : `list` or `tuple` [`str`, `str`], default `('black', 'white')`
401
+ A pair of colors. The first is used for values below a threshold,
402
+ the second for those above.
403
+ threshold : `float` or `None`, default `None`
404
+ The absolute value in data units according to which the colors from
405
+ textcolors are applied. If None (the default) uses the middle of the
406
+ colormap as separation.
407
+ **kwargs : `any`
408
+ All other arguments are forwarded to each call to `text` used to create
409
+ the text labels.
410
+
411
+ Returns
412
+ -------
413
+ texts : `list` of `matplotlib.text.Text`
414
+ A list of text annotation objects added to the heatmap.
415
+ """
416
+
417
+ # mapping data to matrix
418
+ values = im.get_array()
419
+ # masked cells (e.g. from `masked_heatmap`) are not drawn and must not be
420
+ # annotated; `getmaskarray` yields a full boolean mask for masked arrays and
421
+ # an all-False mask for plain arrays, leaving the unmasked path unchanged
422
+ mask = np.ma.getmaskarray(values)
423
+ if data is None:
424
+ matrix = im.get_array()
425
+ elif isinstance(data, pd.DataFrame):
426
+ matrix = data.copy().to_numpy()
427
+ else:
428
+ matrix = data
429
+ # Compare raw data values against the raw threshold
430
+ # This bypases any kind of value normalisation - which we should skip
431
+ # because the string are not normalised only the values
432
+ if threshold is None:
433
+ try:
434
+ threshold = values.max() / 2.
435
+ except np.core._exceptions.UFuncTypeError:
436
+ threshold = None
437
+ # Set default alignment to center, but allow it to be
438
+ # overwritten by text_kw.
439
+ kw = _update_kwargs(update_dict=kwargs,
440
+ horizontalalignment="center",
441
+ verticalalignment="center",
442
+ )
443
+ # Get the formatter in case a string is supplied
444
+ if valfmt is not None:
445
+ if isinstance(valfmt, str):
446
+ valfmt = matplotlib.ticker.StrMethodFormatter(valfmt)
447
+ # Loop over the data and create a `Text` for each "pixel".
448
+ # Change the text's color depending on the data.
449
+ texts = []
450
+ for i in range(matrix.shape[0]):
451
+ for j in range(matrix.shape[1]):
452
+ # skip masked cells, which carry no drawn value to annotate
453
+ if mask[i, j]:
454
+ continue
455
+ # only run if threshold exists
456
+ if threshold is not None:
457
+ kw.update(color=textcolors[int(abs(values[i, j]) >= threshold)])
458
+ # format text or not
459
+ if valfmt is not None:
460
+ # NOTE text takes x, y - whereas values takes rows (y), col (x)
461
+ text = im.axes.text(j, i, valfmt(matrix[i, j], None), **kw)
462
+ else:
463
+ text = im.axes.text(j, i, matrix[i,j], **kw)
464
+ texts.append(text)
465
+ # returning stuff
466
+ return texts
467
+