Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
65 changes: 44 additions & 21 deletions MHCXGraph/scripts/create_heatmaps.py
Original file line number Diff line number Diff line change
Expand Up @@ -471,6 +471,19 @@ def create_heatmap(args):
comp_df_full = build_component_matrix(comp_df)
ratio_df_full = build_ratio_matrix(comp_df)

# A fixed figure size makes tick labels and cell annotations illegible
# once there are many structures, so scale it with the matrix size.
# 0.25 in/structure keeps 11pt labels readable; small matrices keep
# the original 10x8 minimum.
structure_count = len(labels)
figure_width = max(10, structure_count * 0.25)
figure_height = max(8, structure_count * 0.25)

# Cell-value annotations turn into unreadable overlapping text once
# there are too many cells to fit them individually, so skip them
# past this size and let the color alone convey the value.
max_annotated_structures = 50

# Seaborn receives a precomputed SciPy linkage matrix, so it does not
# recalculate pairwise distances. It still uses SciPy to interpret and
# render the dendrogram layout.
Expand All @@ -479,7 +492,7 @@ def create_heatmap(args):
cmap="viridis",
row_linkage=linkage_matrix,
col_linkage=linkage_matrix,
figsize=(10, 8),
figsize=(figure_width, figure_height),
dendrogram_ratio=0.15,
cbar_kws={"label": "Similarity index"},
)
Expand All @@ -492,36 +505,46 @@ def create_heatmap(args):
columns=ordered_labels,
)

cluster_grid.ax_heatmap.set_xticklabels(
heatmap_axis = cluster_grid.ax_heatmap

# Fix the tick positions before assigning labels. sns.clustermap's
# FixedLocator only reliably matches the cell count once the matrix
# is large (~40+ categories); below that it works by chance. Ticks
# are offset by 0.5 to land on cell centers, matching the convention
# clustermap already uses for its own small-matrix ticks.
n = len(ordered_labels)
heatmap_axis.set_xticks([i + 0.5 for i in range(n)])
heatmap_axis.set_yticks([i + 0.5 for i in range(n)])

heatmap_axis.set_xticklabels(
ordered_labels,
fontsize=11,
rotation=90,
)

cluster_grid.ax_heatmap.set_yticklabels(
heatmap_axis.set_yticklabels(
ordered_labels,
fontsize=11,
rotation=0,
)

heatmap_axis = cluster_grid.ax_heatmap

for row_index in range(len(ordered_labels)):
for column_index in range(len(ordered_labels)):
value = component_ordered.iloc[
row_index,
column_index,
]

heatmap_axis.text(
column_index + 0.5,
row_index + 0.5,
f"{int(value)}",
horizontalalignment="center",
verticalalignment="center",
fontsize=11,
color="black",
)
if structure_count <= max_annotated_structures:
for row_index in range(len(ordered_labels)):
for column_index in range(len(ordered_labels)):
value = component_ordered.iloc[
row_index,
column_index,
]

heatmap_axis.text(
column_index + 0.5,
row_index + 0.5,
f"{int(value)}",
horizontalalignment="center",
verticalalignment="center",
fontsize=11,
color="black",
)

output_file = output_directory / args.name

Expand Down
Loading