"""Toxoplasma-specific visualisation helpers.
Every figure here is built inside :func:`_house`, which is
:mod:`spacr.figures.style` applied as a context manager. Read that module
before adding a panel; the rule it exists to enforce is that **everything is
grey except what the sentence is about**, and this file is where breaking it
was measured -- see :func:`custom_volcano_plot`.
"""
import contextlib
import os
import matplotlib.pyplot as plt
import seaborn as sns
import numpy as np
from adjustText import adjust_text
import pandas as pd
from scipy.stats import fisher_exact
from .figures.style import (ROLES, TYPE_SCALE, Palette, reference_line,
resolve_ink, rotate_ticks, text_legend,
theme_target)
from .figures.style import rc as style_rc
from . import tabular
from .plot import save_figure # noqa: F401
#: The page width the published type scale was measured at: 180 mm, the double
#: column of Cell and Nature Microbiology, which is what
#: :data:`spacr.figures.style.TYPE_SCALE` pins its absolute points to.
REFERENCE_WIDTH_IN = 7.09
@contextlib.contextmanager
def _house(width, *, frame='L'):
"""The house style, with the type scaled to the canvas this module asks for.
Apply styling through a context so a figure does not modify process-wide
``rcParams``. Scale every typography tier by the same factor relative to
the 180 mm reference width; this preserves the house-style ratios on the
larger canvases used for volcano plots and related figures.
:param width: the figure's width in inches.
:param frame: ``'L'`` (left and bottom spines) or ``'box'``.
:yields: ``(ink, scale)`` -- the resolved text/axis colour, and the factor
every explicit ``fontsize`` in the block has to be multiplied by.
"""
target = theme_target()
params = style_rc(target, frame=frame)
scale = max(1.0, float(width) / REFERENCE_WIDTH_IN)
for key in ('font.size', 'axes.labelsize', 'axes.titlesize',
'xtick.labelsize', 'ytick.labelsize', 'legend.fontsize'):
params[key] = params[key] * scale
with plt.rc_context(params):
yield resolve_ink(target), scale
#: Leader lines from a label to its point. Grey and hairline, because a
#: leader is furniture: it says which point a name belongs to and nothing
#: else. They were solid black at the default width.
_LEADER = dict(arrowstyle='-', color=ROLES['reference'], linewidth=0.6)
#: Marker shapes for a genuinely categorical second variable, in the order
#: the published figures reach for them. Shape, not hue -- the colour budget
#: belongs to the claim.
_MARKERS = ('o', 's', '^', 'D', 'v', 'P')
def _scaled_sizes(values, low=50.0, high=200.0):
"""Marker areas spanning ``low`` to ``high`` across ``values``.
What seaborn's ``sizes=(50, 200)`` did for the enrichment panels, kept so
the points are the same sizes after the hue came out. A constant column
maps to the small end rather than dividing by a zero range.
"""
values = np.asarray(values, dtype=float)
if values.size == 0:
return values
span = float(np.nanmax(values) - np.nanmin(values))
if not np.isfinite(span) or span <= 0:
return np.full(values.shape, low)
return low + (values - np.nanmin(values)) / span * (high - low)
def _sized_text_legend(ax, entries, scale, **kwargs):
""":func:`spacr.figures.style.text_legend`, at this canvas's type size.
``text_legend`` writes at the pinned 6 pt annotation size, which on the
20-inch volcano is invisible. The entries are re-sized after the fact
rather than re-implemented, so the placement rule stays in one place.
"""
first = len(ax.texts)
text_legend(ax, entries, **kwargs)
for text in ax.texts[first:]:
text.set_fontsize(TYPE_SCALE['annotation'] * scale)
def _compartment_mask(values, wanted):
"""Boolean mask of the rows whose localisation is one of ``wanted``.
:param values: the merged metadata column, one localisation per row.
:param wanted: ``None``, one compartment name, or a sequence of them.
:returns: ``(mask, label)``, or ``(None, '')`` when nothing was asked for.
"""
if wanted is None:
return None, ''
names = [wanted] if isinstance(wanted, str) else list(wanted)
names = [str(name) for name in names if str(name)]
if not names:
return None, ''
mask = values.astype(str).isin(names)
return mask, ', '.join(names)
[docs]
def custom_volcano_plot(
data_path,
metadata_path,
metadata_column='tagm_location',
point_size=50,
figsize=20,
threshold=0,
save_path=None,
x_lim=None,
y_lims=None,
draw=True,
highlight_location=None,
):
"""Render a volcano plot and return the significant feature names.
Plot each feature at ``(coefficient, -log10(p_value))``. Features that do
not pass the call rule are grey; called positive effects are green and
called negative or zero effects are rust. This direction-based palette
makes significance and effect direction the primary visual encoding.
Localization is optional: ``highlight_location`` overlays selected
compartments in blue instead of assigning simultaneous colours to every
category.
Parameters
----------
data_path : pandas.DataFrame or path-like
Regression table containing ``feature``, ``coefficient``, and
``p_value`` columns. DataFrame input is copied.
metadata_path : pandas.DataFrame or path-like
Gene metadata containing one row per ``gene_nr`` and the selected
metadata column. DataFrame input is copied.
metadata_column : str, optional
Localization or annotation column used by ``highlight_location``.
point_size : float, optional
Marker area passed to ``Axes.scatter``.
figsize : float, optional
Width and height of the square figure in inches. Typography scales
with this value.
threshold : float, optional
Absolute coefficient threshold for calls. A row is returned when
``p_value <= 0.05`` and
``abs(coefficient) >= abs(threshold)``.
save_path : path-like, optional
Destination passed to :func:`spacr.figures.scene.write_figure`, which
draws the scene the screen would show and falls back to
:func:`spacr.plot.save_figure`. The written extension follows the
configured figure format either way.
x_lim : sequence of float, optional
Two x-axis limits. The default is ``[-0.5, 0.5]``.
y_lims : sequence, optional
Use ``[low, high]`` for one axis or
``[[lower_low, lower_high], [upper_low, upper_high]]`` for a broken
y-axis. By default, fit one axis to the finite values.
draw : bool, optional
Build, optionally save, and show the figure. If false, return the hit
list before constructing a figure.
highlight_location : str or sequence of str, optional
Values from ``metadata_column`` to overlay in the highlight colour and
name in the in-panel legend.
Returns
-------
list of str
Feature-derived ``variable`` names that satisfy the call rule, in
table order.
Raises
------
pandas.errors.MergeError
If the metadata contains duplicate ``gene_nr`` values and therefore
cannot be joined many-to-one.
ValueError
If ``y_lims`` does not match a supported form.
"""
if x_lim is None:
x_lim = [-0.5, 0.5]
from matplotlib.gridspec import GridSpec
if isinstance(data_path, pd.DataFrame):
data = data_path.copy()
else:
data = tabular.read_table(data_path, report=None)
data['variable'] = data['feature'].str.extract(r'\[(.*?)\]')
data['variable'] = data['variable'].fillna(data['feature'])
data['gene_nr'] = data['variable'].str.split('_').str[0]
data = data[data['variable'] != 'Intercept']
if isinstance(metadata_path, pd.DataFrame):
metadata = metadata_path.copy()
else:
metadata = tabular.read_table(metadata_path, report=None)
metadata['gene_nr'] = metadata['gene_nr'].astype(str)
data['gene_nr'] = data['gene_nr'].astype(str)
try:
merged_data = pd.merge(
data,
metadata[['gene_nr', metadata_column]],
on='gene_nr',
how='left',
validate='many_to_one',
)
except pd.errors.MergeError as exc:
duplicated = metadata.loc[
metadata['gene_nr'].duplicated(keep=False), 'gene_nr']
if duplicated.empty:
raise
examples = duplicated.unique()[:5].tolist()
raise pd.errors.MergeError(
f"The gene metadata lists {duplicated.nunique()} gene_nr value(s) "
f"more than once (e.g. {examples}), so it cannot say which "
f"{metadata_column!r} belongs to a gene. Joining it anyway would "
f"plot those genes once per duplicate row and return each of them "
f"more than once in the hit list. De-duplicate the metadata on "
f"gene_nr before plotting. (pandas: {exc})"
) from exc
merged_data[metadata_column] = merged_data[metadata_column].fillna('unknown')
merged_data['neg_log_p'] = -np.log10(merged_data['p_value'])
called = ((merged_data['p_value'] <= 0.05)
& (merged_data['coefficient'].abs() >= abs(threshold)))
hit_list = list(merged_data.loc[called, 'variable'])
if not draw:
return hit_list
is_broken, lower_lim, upper_lim = _normalize_y_lims(
y_lims, merged_data['neg_log_p'])
with _house(figsize) as (ink, scale):
if is_broken:
fig = plt.figure(figsize=(figsize, figsize))
from .figures.bundle import _register_figure_data
_register_figure_data(fig, merged_data, x="coefficient", y="neg_log_p", kind="scatter")
gs = GridSpec(2, 1, height_ratios=[1, 3], hspace=0.05)
ax_upper = fig.add_subplot(gs[0])
ax_lower = fig.add_subplot(gs[1], sharex=ax_upper)
ax_upper.tick_params(axis='x', which='both', bottom=False, labelbottom=False)
all_axes = [ax_lower, ax_upper]
else:
fig, ax_lower = plt.subplots(figsize=(figsize, figsize))
from .figures.bundle import _register_figure_data
_register_figure_data(fig, merged_data, x="coefficient", y="neg_log_p", kind="scatter")
ax_upper = None
all_axes = [ax_lower]
coefficient = merged_data['coefficient'].to_numpy(dtype=float)
neg_log_p = merged_data['neg_log_p'].to_numpy(dtype=float)
called_mask = called.to_numpy(dtype=bool)
on_upper = (neg_log_p > upper_lim[0]) if is_broken \
else np.zeros(len(neg_log_p), dtype=bool)
up = called_mask & (coefficient > 0)
down = called_mask & (coefficient <= 0)
wanted, wanted_label = _compartment_mask(
merged_data[metadata_column], highlight_location)
layers = [(~called_mask, ROLES['data'], 1, None),
(up, ROLES['up'], 2, f"called, positive ({int(up.sum())})"),
(down, ROLES['down'], 2, f"called, negative ({int(down.sum())})")]
if wanted is not None:
layers.append((wanted.to_numpy(dtype=bool), ROLES['highlight'], 3,
wanted_label))
entries = []
for selected, colour, zorder, label in layers:
for axis, side in ((ax_lower, ~on_upper), (ax_upper, on_upper)):
if axis is None:
continue
take = selected & side
if not take.any():
continue
axis.scatter(coefficient[take], neg_log_p[take], color=colour,
marker='o', s=point_size, linewidths=0,
zorder=zorder)
if label and selected.any():
entries.append((label, colour))
ax_lower.set_ylim(lower_lim)
ax_lower.set_xlim(x_lim)
ax_lower.set_xlabel('Coefficient')
ax_lower.set_ylabel('-log10(p-value)')
if is_broken:
ax_upper.set_ylim(upper_lim)
ax_upper.set_ylabel('-log10(p-value)')
ax_upper.spines['bottom'].set_visible(False)
for axis in all_axes:
if threshold:
reference_line(axis, x=-abs(threshold))
reference_line(axis, x=abs(threshold))
else:
reference_line(axis, x=0.0)
reference_line(ax_lower, y=-np.log10(0.05))
texts_upper, texts_lower = [], []
label_size = TYPE_SCALE['annotation'] * scale
for index in np.flatnonzero(called_mask):
axis = ax_upper if on_upper[index] else ax_lower
text = axis.text(
coefficient[index],
neg_log_p[index],
merged_data['variable'].iat[index],
fontsize=label_size,
color=ink,
ha='center',
va='bottom',
)
if axis is ax_upper:
texts_upper.append(text)
else:
texts_lower.append(text)
leader = dict(arrowstyle='-', color=ROLES['reference'],
linewidth=0.6)
if texts_lower:
adjust_text(texts_lower, ax=ax_lower, arrowprops=leader)
if is_broken and texts_upper:
adjust_text(texts_upper, ax=ax_upper, arrowprops=leader)
if entries:
_sized_text_legend(ax_lower, entries, scale)
if save_path:
save_path = _write_the_figure(fig, save_path,
bbox_inches='tight')
plt.show()
return hit_list
def _write_the_figure(fig, save_path, title=None, **kwargs):
"""Write ``fig`` through the SCENE renderer, not straight to matplotlib.
Every generated figure on this path is asked for in pyqtgraph, so the
file a run writes is the same scene the screen would show. The
translation is a whitelist: an artist nobody thought about leaves the
figure incomplete and the matplotlib page is written instead, with the
reason recorded rather than a piece of the picture quietly missing.
``save_figure`` is still what runs in that fallback, so the format
preference and the DPI reach the file either way.
:param fig: the finished matplotlib figure.
:param save_path: destination; the extension follows the preference.
:param title: the name the gallery tile carries.
:returns: the path actually written.
"""
from .figures.scene import write_figure
written, _drew, _why = write_figure(fig, os.fspath(save_path),
title=title, **kwargs)
return written or save_path
def _fit_outside_legend(fig, legend, pad=0.02, min_axes_width=0.45):
"""Make room inside ``fig`` for a legend anchored outside the axes.
A legend at ``bbox_to_anchor=(1.02, 1)`` sits beyond the axes and therefore
beyond the figure. ``bbox_inches='tight'`` grows the saved file to cover it,
so the PDF on disk looks right -- but nothing rescues the figure shown in
the application, which is drawn at the figure's own extent. The legend is
simply clipped, reported as "the volcano plot is always cut off on the
right side".
Measuring the legend and shrinking the axes to fit it means the figure is
correct as drawn, on screen and on disk alike, instead of being correct
only after a save-time rescue.
NO FIGURE IN THIS MODULE ANCHORS A LEGEND OUTSIDE ITS AXES ANY MORE: the
house style puts a two-line text legend inside the panel, which is what
made the 27-swatch column that this helper was written to survive
unnecessary in the first place. It is kept for a panel that has to.
:param fig: Figure holding the legend.
:param legend: The legend to make room for.
:param pad: Extra figure-width fraction left beside the legend.
:param min_axes_width: Never shrink the axes below this fraction, so a
runaway legend cannot squeeze the data down to nothing.
"""
if legend is None:
return
try:
fig.canvas.draw()
extent = legend.get_window_extent()
fig_width = fig.get_figwidth() * fig.dpi
if fig_width <= 0:
return
right = 1.0 - (extent.width / fig_width) - pad
fig.subplots_adjust(right=max(min(right, 0.98), min_axes_width))
except Exception:
pass
def _normalize_y_lims(y_lims, neg_log_p):
"""Coerce y_lims into ``(is_broken, lower_lim, upper_lim)`` for volcano plotting.
- ``None``: auto-fit a single panel from the data.
- ``[low, high]``: single panel with explicit limits.
- ``[[low1, high1], [low2, high2]]``: broken axis (lower, upper).
:raises ValueError: When ``y_lims`` does not match one of the supported forms.
"""
if y_lims is None:
finite = neg_log_p[np.isfinite(neg_log_p)]
if len(finite) == 0:
return False, [0.0, 1.0], None
ymax = float(finite.max()) * 1.05
return False, [0.0, max(ymax, 1.0)], None
if not (isinstance(y_lims, (list, tuple)) and len(y_lims) == 2):
raise ValueError(
"y_lims must be None, [low, high], or [[low1, high1], [low2, high2]]; "
f"got {y_lims!r}"
)
a, b = y_lims
if all(isinstance(v, (int, float)) or v is None for v in (a, b)):
return False, [a, b], None
if all(isinstance(v, (list, tuple)) and len(v) == 2 for v in (a, b)):
return True, list(a), list(b)
raise ValueError(
"y_lims must be None, [low, high], or [[low1, high1], [low2, high2]]; "
f"got {y_lims!r}"
)
[docs]
def go_term_enrichment_by_column(significant_df, metadata_path, go_term_columns=None):
"""Compute and plot GO-term enrichment for each requested metadata column.
For every ``go_term_column`` counts occurrences among hit vs background
genes, runs Fisher's exact test per term, and produces scatter plots of
enrichment vs ``-log10(p)`` both per column and combined.
:param significant_df: DataFrame of screen hits with a ``n_gene`` column.
:param metadata_path: CSV path holding ``Gene ID`` plus GO-term columns.
:param go_term_columns: Columns to test. Defaults to the four standard
Computed/Curated GO categories.
:returns: None. Results are displayed as Matplotlib figures.
"""
if go_term_columns is None:
go_term_columns = ['Computed GO Processes', 'Curated GO Components', 'Curated GO Functions', 'Curated GO Processes']
significant_df = significant_df.loc[
significant_df['n_gene'].notna()].copy()
gene_list = significant_df['n_gene'].to_list()
metadata = tabular.read_table(metadata_path, report=None)
split_columns = metadata['Gene ID'].str.split('_', expand=True)
metadata['gene_nr'] = split_columns[1]
hits_metadata = metadata.loc[
metadata['gene_nr'].isin(gene_list)].copy()
combined_results = []
for go_term_column in go_term_columns:
go_terms = []
enrichment_scores = []
p_values = []
metadata[go_term_column] = metadata[go_term_column].fillna('')
hits_metadata[go_term_column] = hits_metadata[go_term_column].fillna('')
all_go_terms = metadata[go_term_column].str.split(';').explode()
hit_go_terms = hits_metadata[go_term_column].str.split(';').explode()
all_go_term_counts = all_go_terms.value_counts()
hit_go_term_counts = hit_go_terms.value_counts()
for go_term in all_go_term_counts.index:
total_with_go_term = all_go_term_counts.get(go_term, 0)
hits_with_go_term = hit_go_term_counts.get(go_term, 0)
total_genes = len(metadata)
total_hits = len(hits_metadata)
contingency_table = [[hits_with_go_term, total_hits - hits_with_go_term],
[total_with_go_term - hits_with_go_term, total_genes - total_hits - (total_with_go_term - hits_with_go_term)]]
_, p_value = fisher_exact(contingency_table)
if total_with_go_term > 0 and total_hits > 0:
enrichment_score = (hits_with_go_term / total_hits) / (total_with_go_term / total_genes)
else:
enrichment_score = 0.0
if enrichment_score > 0.0:
go_terms.append(go_term)
enrichment_scores.append(enrichment_score)
p_values.append(p_value)
results_df = pd.DataFrame({
'GO Term': go_terms,
'Enrichment Score': enrichment_scores,
'P-value': p_values,
'GO Column': go_term_column
})
results_df = results_df.sort_values(by='Enrichment Score', ascending=False)
combined_results.append(results_df)
with _house(10) as (ink, scale):
fig = plt.figure(figsize=(10, 6))
from .figures.bundle import _register_figure_data
_register_figure_data(fig, results_df, x="Enrichment Score", y="P-value", kind="scatter")
ax = fig.gca()
enrichment = results_df['Enrichment Score'].to_numpy(dtype=float)
significance = -np.log10(results_df['P-value'].to_numpy(dtype=float))
called = results_df['P-value'].to_numpy(dtype=float) <= 0.05
ax.scatter(enrichment, significance, s=_scaled_sizes(enrichment),
color=[ROLES['highlight'] if hit else ROLES['data']
for hit in called],
linewidths=0, zorder=2)
reference_line(ax, y=-np.log10(0.05))
reference_line(ax, x=1.0)
ax.set_title(f'GO Term Enrichment Analysis for {go_term_column}')
ax.set_xlabel('Enrichment Score')
ax.set_ylabel('-log10(P-value)')
texts = [ax.text(enrichment[i], significance[i],
results_df['GO Term'].iat[i],
fontsize=TYPE_SCALE['annotation'] * scale,
color=ink)
for i in np.flatnonzero(called)]
if texts:
adjust_text(texts, ax=ax, arrowprops=_LEADER)
_sized_text_legend(
ax,
[(f'p <= 0.05 ({int(called.sum())})', ROLES['highlight']),
(f'not called ({int((~called).sum())})', ROLES['data'])],
scale)
fig.tight_layout()
plt.show()
print(f'Results for {go_term_column}')
combined_df = pd.concat(combined_results)
with _house(12) as (ink, scale):
fig = plt.figure(figsize=(12, 8))
from .figures.bundle import _register_figure_data
_register_figure_data(fig, combined_df, x="Enrichment Score", y="P-value", kind="scatter")
ax = fig.gca()
enrichment = combined_df['Enrichment Score'].to_numpy(dtype=float)
significance = -np.log10(combined_df['P-value'].to_numpy(dtype=float))
called = combined_df['P-value'].to_numpy(dtype=float) <= 0.05
sizes = _scaled_sizes(enrichment)
colours = [ROLES['highlight'] if hit else ROLES['data']
for hit in called]
for index, column in enumerate(dict.fromkeys(combined_df['GO Column'])):
rows = np.flatnonzero(
(combined_df['GO Column'] == column).to_numpy(dtype=bool))
ax.scatter(enrichment[rows], significance[rows], s=sizes[rows],
color=[colours[i] for i in rows],
marker=_MARKERS[index % len(_MARKERS)],
linewidths=0, zorder=2)
reference_line(ax, y=-np.log10(0.05))
reference_line(ax, x=1.0)
ax.set_title('Combined GO Term Enrichment Analysis')
ax.set_xlabel('Enrichment Score')
ax.set_ylabel('-log10(P-value)')
texts = [ax.text(enrichment[i], significance[i],
combined_df['GO Term'].iat[i],
fontsize=TYPE_SCALE['annotation'] * scale, color=ink)
for i in range(len(combined_df))]
adjust_text(texts, ax=ax, arrowprops=_LEADER)
fig.tight_layout()
plt.show()
[docs]
def plot_gene_phenotypes(data, gene_list, x_column='Gene ID', data_column='T.gondii GT1 CRISPR Phenotype - Mean Phenotype',error_column='T.gondii GT1 CRISPR Phenotype - Standard Error', save_path=None):
"""Plot ranked mean phenotype with SE shading and highlight selected genes.
Parameters
----------
data : pandas.DataFrame
Gene identifiers and phenotype/error columns. The frame is copied
before numeric conversion.
gene_list : iterable of str
Gene names or ``TGGT1_<id>`` identifiers to highlight.
x_column : str, default='Gene ID'
Column used to match gene identifiers.
data_column : str
Mean phenotype column plotted on the y-axis.
error_column : str
Standard-error column used for the uncertainty band.
save_path : path-like, optional
Figure destination. The configured figure format controls the final
extension.
Notes
-----
The complete ranked phenotype curve is drawn in grey and selected genes
use the spaCR accent colour. The figure is displayed after optional save.
"""
def extract_gene_id(gene):
"""Return the numeric portion of a ``TGGT1_<id>`` tag, or ``gene`` itself."""
if isinstance(gene, str) and '_' in gene:
return gene.split('_')[1]
return str(gene)
data = data.copy()
data.loc[:, data_column] = pd.to_numeric(data[data_column], errors='coerce')
data = data.dropna(subset=[data_column])
data.loc[:, error_column] = pd.to_numeric(data[error_column], errors='coerce')
data = data.dropna(subset=[error_column])
data['x'] = data[x_column].apply(extract_gene_id)
data = data.sort_values(by=data_column).reset_index(drop=True)
data['rank'] = range(1, len(data) + 1)
x = data['rank']
y = data[data_column]
yerr = data[error_column]
with _house(10) as (ink, scale):
fig = plt.figure(figsize=(10, 10))
from .figures.bundle import _register_figure_data
_register_figure_data(fig, data, x="rank", y=data_column, kind="line")
plt.plot(x, y, label='Mean Phenotype', color=Palette.GREY_DARK,
linewidth=1.2)
plt.fill_between(
x, y - yerr, y + yerr,
color=Palette.GREY_DARK, alpha=0.25, label='Standard Error',
linewidth=0,
)
texts = []
for gene in gene_list:
gene_id = extract_gene_id(gene)
gene_data = data[data['x'] == gene_id]
if not gene_data.empty:
plt.scatter(
gene_data['rank'],
gene_data[data_column],
color=ROLES['highlight'],
s=200,
linewidths=0,
label=f'Highlighted Gene: {gene}',
zorder=3
)
texts.append(
plt.text(
gene_data['rank'].values[0],
gene_data[data_column].values[0],
gene,
fontsize=TYPE_SCALE['annotation'] * scale,
color=ink,
ha='right',
)
)
adjust_text(texts, arrowprops=_LEADER)
plt.xlabel('Rank')
plt.ylabel('Mean Phenotype')
plt.legend().remove()
plt.tight_layout()
if save_path:
save_path = _write_the_figure(fig, save_path,
bbox_inches='tight')
print(f"Figure saved to {save_path}")
plt.show()
[docs]
def plot_gene_heatmaps(data, gene_list, columns, x_column='Gene ID', normalize=False, save_path=None):
"""Render a heatmap for selected genes across selected metadata columns.
THE RAMP IS SINGLE-HUE. It was viridis, a rainbow that reads as five
categories where the quantity is one ordered score; the house style's
``Blues`` runs light to dark, so "more" is one direction rather than a
tour of the spectrum. A diverging map would be right only if the values
were signed, and after ``normalize`` they run 0 to 1.
:param data: DataFrame containing per-gene rows. Copied before use -- the
row-key column this adds used to appear on the caller's frame.
:param gene_list: Genes to include as heatmap rows.
:param columns: Column names to include as heatmap columns.
:param x_column: Column holding gene identifiers for row matching.
:param normalize: When True, min-max scale each gene's row to [0, 1].
:param save_path: Optional destination for the figure. Saving goes through
:func:`spacr.figures.scene.write_figure` and, on its fallback,
:func:`spacr.plot.save_figure`, so the format and the file extension
follow the figure preference rather than always being PDF.
:returns: None. Displays the Matplotlib figure.
"""
def extract_gene_id(gene):
"""Return the numeric portion of a ``TGGT1_<id>`` tag, or ``gene`` itself."""
if isinstance(gene, str) and '_' in gene:
return gene.split('_')[1]
return str(gene)
data = data.copy()
data['x'] = data[x_column].apply(extract_gene_id)
filtered_data = data[data['x'].isin(gene_list)].set_index('x')[columns]
if normalize:
filtered_data = filtered_data.apply(lambda x: (x - x.min()) / (x.max() - x.min()), axis=1)
width = len(columns) * 4
height = len(gene_list) * 1
with _house(width, frame='box') as (ink, scale):
fig = plt.figure(figsize=(width, height))
from .figures.bundle import _register_figure_data
_register_figure_data(fig, filtered_data, kind="heatmap", matrix=True)
cmap = sns.color_palette(Palette.SEQUENTIAL, as_cmap=True)
ax = sns.heatmap(
filtered_data,
cmap=cmap,
cbar=True,
annot=False,
linewidths=0,
square=True
)
rotate_ticks(ax, 45)
plt.yticks(rotation=0)
plt.xlabel('')
plt.ylabel('')
for bar in fig.axes[len(fig.axes) - 1:]:
bar.tick_params(colors=ink, labelsize=TYPE_SCALE['tick'] * scale)
for spine in bar.spines.values():
spine.set_visible(False)
plt.tight_layout()
if save_path:
save_path = _write_the_figure(fig, save_path,
bbox_inches='tight')
print(f"Figure saved to {save_path}")
plt.show()
[docs]
def generate_score_heatmap(settings):
"""Build combined classification-score and control-fraction heatmaps for a plate.
Thin wrapper around :func:`spacr.submodules.generate_score_heatmap`, kept
only so the historic ``spacr.toxo`` import path keeps working. This module
used to carry a second copy of that function, identical to it line for line
apart from the key names and the colormap, and left behind by the
``column_name`` -> ``columnID`` rename: it filtered, grouped and merged on
``column_name``, a key no spaCR CSV carries any more, and one helper even
created ``columnID`` and then immediately indexed ``column_name``. Every
call raised ``KeyError('column_name')`` on a canonical input. Rather than
repair a second copy, delegate to the one that was migrated.
The only behavioural difference between the two copies was the colormap:
this one hard-coded ``'viridis'`` and ignored ``settings['cmap']``, while
the ``submodules`` version requires it. The default below preserves what
``toxo`` callers used to get while now honouring ``cmap`` when they pass it.
Imported inside the function on purpose: ``spacr.submodules`` pulls in
cellpose, torch and shap at import time, and ``spacr.ml`` imports this
module.
:param settings: Config dict with keys ``folders``, ``csv_name``,
``data_column``, ``csv``, ``cv_csv``, ``data_column_cv``, ``plateID``,
``columnID``, ``control_sgrnas``, ``fraction_grna``, ``dst`` and,
optionally, ``cmap``.
:returns: merged DataFrame joining reads, classifier scores and CV scores per well.
"""
from .submodules import generate_score_heatmap as _generate_score_heatmap
settings = dict(settings)
settings.setdefault('cmap', 'viridis')
return _generate_score_heatmap(settings)