"""
Hasse diagram visualization for knowledge structures and surmise relations.
Requires the ``viz`` extra: ``pip install knowledgespaces[viz]``
The Hasse diagram of a knowledge structure plots states as nodes and
draws the covering relation of state inclusion: no intermediate state
lies strictly between the endpoints. For surmise relations, it shows the
transitive reduction.
References:
Falmagne, J.-C., & Doignon, J.-P. (2011).
Learning Spaces, Chapter 1. Springer-Verlag.
"""
from __future__ import annotations
from collections.abc import Collection, Mapping
from dataclasses import dataclass
from itertools import pairwise
from math import isfinite
from typing import TYPE_CHECKING, Literal, cast
from knowledgespaces.structures.knowledge_base import KnowledgeBase
from knowledgespaces.structures.knowledge_structure import KnowledgeStructure
from knowledgespaces.structures.relations import SurmiseRelation
from knowledgespaces.structures.set_family import SetFamily
if TYPE_CHECKING:
from matplotlib.axes import Axes
from matplotlib.figure import Figure
_MATPLOTLIB_IMPORT_ERROR = (
"matplotlib is required for visualization. Install it with: pip install 'knowledgespaces[viz]'"
)
def _state_label(state: frozenset[str]) -> str:
"""Human-readable label for a knowledge state."""
if not state:
return "∅"
return "{" + ", ".join(sorted(state)) + "}"
[docs]
@dataclass(frozen=True)
class HasseData:
"""Exact plotted members, cover edges and coordinates, without matplotlib.
Edges use integer indices into ``states``. ``as_dict`` exports JSON-ready
lists; item identities are preserved independently of display labels.
"""
domain: frozenset[str]
states: tuple[frozenset[str], ...]
edges: tuple[tuple[int, int], ...]
positions: tuple[tuple[float, float], ...]
def as_dict(self) -> dict[str, object]:
return {
"domain": sorted(self.domain),
"states": [sorted(s) for s in self.states],
"edges": [list(edge) for edge in self.edges],
"positions": [list(position) for position in self.positions],
}
[docs]
def hasse_data(
structure: KnowledgeStructure | SetFamily | KnowledgeBase,
*,
orientation: Literal["vertical", "horizontal"] = "vertical",
) -> HasseData:
"""Compute inclusion covers of the supplied family, preserving its members.
For ``KnowledgeBase`` this plots the base rows, never its union span.
Empty families are allowed. Coordinates use cardinality layers, swapped
for horizontal orientation. This layout does not promise to avoid every
edge crossing. No optional plotting dependency is imported.
"""
if orientation not in ("vertical", "horizontal"):
raise ValueError("orientation must be 'vertical' or 'horizontal'.")
family = structure.as_family() if isinstance(structure, KnowledgeBase) else structure
states = tuple(sorted(family, key=lambda s: (len(s), sorted(s))))
layers: dict[int, list[int]] = {}
for i, state in enumerate(states):
layers.setdefault(len(state), []).append(i)
width = max((len(layer) for layer in layers.values()), default=1)
positions = [(0.0, 0.0)] * len(states)
for size, layer in sorted(layers.items()):
for index, node in enumerate(layer):
across = (index - (len(layer) - 1) / 2) * width / len(layer)
positions[node] = (
(across, float(size)) if orientation == "vertical" else (float(size), across)
)
edges = []
for i, left in enumerate(states):
covers: list[frozenset[str]] = []
for j, right in enumerate(states):
if left < right and not any(middle < right for middle in covers):
covers.append(right)
edges.append((i, j))
return HasseData(structure.domain, states, tuple(edges), tuple(positions))
[docs]
def plot_hasse(
structure: KnowledgeStructure | SetFamily | KnowledgeBase,
*,
ax: Axes | None = None,
figsize: tuple[float, float] = (10, 8),
node_color: str = "#4A90D9",
edge_color: str = "#888888",
highlight_states: Collection[frozenset[str]] | None = None,
highlight_color: str = "#E74C3C",
title: str | None = None,
show_labels: bool = True,
font_size: int = 8,
orientation: Literal["vertical", "horizontal"] = "vertical",
node_labels: Mapping[frozenset[str], str] | None = None,
node_colors: Mapping[frozenset[str], str] | None = None,
node_values: Mapping[frozenset[str], float] | None = None,
cmap: str = "viridis",
value_label: str = "Node value",
highlight_paths: Collection[Collection[Collection[str]]] | None = None,
centre: Collection[str] | None = None,
) -> Figure:
"""Plot the Hasse diagram of a knowledge structure.
Nodes are knowledge states, arranged in layers by cardinality.
Edges connect states with no intermediate state in the inclusion order.
They differ by one item on a well-graded family, but may differ by
multiple items on a general structure or with equivalent items.
Parameters
----------
structure : KnowledgeStructure, SetFamily or KnowledgeBase
Exact members to visualize. A base is not expanded into its span.
ax : matplotlib Axes or None
Axes to draw on. If None, a new figure is created.
figsize : tuple
Figure size (width, height) in inches.
node_color : str
Color for state nodes.
edge_color : str
Color for edges.
highlight_states : collection of frozensets or None
States to highlight (e.g., atoms, base, current state).
highlight_color : str
Color for highlighted states.
title : str or None
Plot title. ``None`` (the default) uses a standard title; pass
an empty string to suppress the title entirely.
show_labels : bool
Whether to show state labels on nodes.
font_size : int
Font size for labels.
orientation : {'vertical', 'horizontal'}
Direction of increasing cardinality.
node_labels, node_colors : mapping or None
Optional display labels or colors, keyed by exact member states.
node_values : mapping or None
Finite values for every node, mapped to ``cmap`` with a colorbar.
Useful for state masses; not normalized or treated as probabilities.
Mutually exclusive with node_colors. Highlights then use outlines.
cmap, value_label : str
Matplotlib colormap name and colorbar label.
highlight_paths : collection of paths or None
Highlight states and edges of supplied chains of inclusion covers.
centre : collection or None
A member to mark with a diamond, for example a neighbourhood centre.
Returns
-------
matplotlib.figure.Figure
The figure containing the Hasse diagram.
"""
try:
import matplotlib.patches as mpatches
import matplotlib.pyplot as plt
except ImportError as e:
raise ImportError(_MATPLOTLIB_IMPORT_ERROR) from e
data = hasse_data(structure, orientation=orientation)
states = data.states
state_set = set(states)
highlight_set = {frozenset(s) for s in (highlight_states or ())}
centre_state = None if centre is None else frozenset(centre)
if not highlight_set <= state_set or (
centre_state is not None and centre_state not in state_set
):
raise ValueError("Highlighted states and centre must be members of the plotted family.")
for mapping in (node_labels, node_colors, node_values):
if mapping is not None and not set(mapping) <= state_set:
raise ValueError("Node mapping contains a state outside the plotted family.")
if node_colors is not None and node_values is not None:
raise ValueError("Use either node_colors or node_values, not both.")
if node_values is not None and (
set(node_values) != state_set or not all(isfinite(v) for v in node_values.values())
):
raise ValueError("node_values must supply one finite value for every state.")
edges = tuple((states[i], states[j]) for i, j in data.edges)
marked_edges = set()
for raw_path in highlight_paths or ():
path = tuple(frozenset(s) for s in raw_path)
if not path or not set(path) <= state_set:
raise ValueError("Each highlighted path must contain members of the plotted family.")
for edge in pairwise(path):
if edge not in edges:
raise ValueError("Highlighted path steps must be inclusion covers of the family.")
marked_edges.add(edge)
highlight_set.update(path)
if ax is None:
fig, ax = plt.subplots(1, 1, figsize=figsize)
else:
fig = cast("Figure", ax.get_figure())
positions = dict(zip(states, data.positions, strict=True))
colors: Mapping[frozenset[str], object] = node_colors or {}
if node_values:
import numpy as np
from matplotlib.cm import ScalarMappable
from matplotlib.colors import Normalize
values = list(node_values.values())
scale = ScalarMappable(norm=Normalize(min(values), max(values)), cmap=cmap)
rgba = scale.to_rgba(np.array([node_values[s] for s in states]))
colors = dict(zip(states, rgba, strict=True))
fig.colorbar(scale, ax=ax, label=value_label)
# Draw edges
for s1, s2 in edges:
x1, y1 = positions[s1]
x2, y2 = positions[s2]
marked = (s1, s2) in marked_edges
ax.plot(
[x1, x2],
[y1, y2],
color=highlight_color if marked else edge_color,
linewidth=2.5 if marked else 0.8,
zorder=1,
)
# Draw nodes
for s in states:
x, y = positions[s]
color = colors.get(s, highlight_color if s in highlight_set else node_color)
ax.scatter(
x,
y,
s=300,
color=color,
zorder=2,
marker="D" if s == centre_state else "o",
edgecolors=highlight_color if s in highlight_set else "white",
linewidth=1.5,
)
if show_labels:
ax.annotate(
(node_labels or {}).get(s, _state_label(s)),
(x, y),
textcoords="offset points",
xytext=(0, 12),
ha="center",
fontsize=font_size,
fontweight="bold" if s in highlight_set else "normal",
)
# Legend for highlights
if highlight_set:
handles = [
mpatches.Patch(
facecolor="none", edgecolor=highlight_color, label="Highlighted outline / path"
)
]
if not colors:
handles.insert(0, mpatches.Patch(color=node_color, label="States"))
ax.legend(handles=handles, loc="upper left", fontsize=font_size)
if title is None:
title = "Hasse Diagram"
if title:
ax.set_title(title, fontsize=12, fontweight="bold")
if orientation == "vertical":
ax.set_ylabel("State size")
ax.set_xticks([])
else:
ax.set_xlabel("State size")
ax.set_yticks([])
ax.margins(0.15)
fig.tight_layout()
return fig
[docs]
def plot_relation(
relation: SurmiseRelation,
*,
ax: Axes | None = None,
figsize: tuple[float, float] = (8, 6),
node_color: str = "#2ECC71",
edge_color: str = "#555555",
title: str | None = None,
font_size: int = 10,
collapse_equivalent: bool = False,
orientation: Literal["vertical", "horizontal"] = "vertical",
node_labels: Mapping[str, str] | None = None,
node_colors: Mapping[str, str] | None = None,
) -> Figure:
"""Plot the Hasse diagram of a surmise relation.
Nodes are items, arranged by topological level. Edges show
the transitive reduction (direct prerequisites only).
Parameters
----------
relation : SurmiseRelation
The surmise relation (will be transitively reduced).
ax : matplotlib Axes or None
Axes to draw on. If None, a new figure is created.
figsize : tuple
Figure size in inches.
node_color : str
Node color.
edge_color : str
Edge color.
title : str or None
Plot title. ``None`` (the default) uses a standard title; pass
an empty string to suppress the title entirely.
font_size : int
Font size for item labels.
collapse_equivalent : bool
If True, draw the partial-order quotient by mutual prerequisites.
Each node displays every item in its class. Class representatives
and members are separately available via relation.quotient().
Default False preserves rejection of cycles. Closure and reduction
operate directly on the relation, without enumerating knowledge states.
orientation : {'vertical', 'horizontal'}
Direction of increasing prerequisite level; arrows keep their meaning.
node_labels, node_colors : mapping or None
Optional labels or colors keyed by items (quotient representatives
when collapse_equivalent=True). Identities and order are unchanged.
Returns
-------
matplotlib.figure.Figure
"""
try:
import matplotlib.pyplot as plt
except ImportError as e:
raise ImportError(_MATPLOTLIB_IMPORT_ERROR) from e
classes: dict[str, frozenset[str]] | None = None
if collapse_equivalent:
relation, classes = relation.quotient()
elif not relation.is_antisymmetric():
raise ValueError(
"Cannot plot Hasse diagram: the relation contains cycles "
"(is not antisymmetric). A Hasse diagram is only defined "
"for partial orders. Use collapse_equivalent=True to draw its quotient."
)
if orientation not in ("vertical", "horizontal"):
raise ValueError("orientation must be 'vertical' or 'horizontal'.")
for mapping in (node_labels, node_colors):
if mapping is not None and not set(mapping) <= relation.items:
raise ValueError("Node mapping contains unknown relation items.")
if ax is None:
fig, ax = plt.subplots(1, 1, figsize=figsize)
else:
fig = cast("Figure", ax.get_figure())
# Use transitive reduction for clean Hasse diagram
hasse = relation.transitive_reduction()
levels = relation.transitive_closure().levels()
# Group items by level
level_groups: dict[int, list[str]] = {}
for item, lvl in levels.items():
level_groups.setdefault(lvl, []).append(item)
# Compute positions
max_width = max(len(v) for v in level_groups.values()) if level_groups else 1
positions: dict[str, tuple[float, float]] = {}
for lvl, items in sorted(level_groups.items()):
n = len(items)
for i, item in enumerate(sorted(items)):
x = (i - (n - 1) / 2) * (max_width / max(n, 1))
positions[item] = (x, lvl) if orientation == "vertical" else (lvl, x)
# Fit all class labels inside their node; shorten arrows to the boundary
# so their heads are not hidden by the node drawn above them.
node_sizes = {
item: 500.0
if classes is None
else max(500.0, (1.2 * font_size * len(classes[item]) + 12) ** 2)
for item in positions
}
# Draw edges (arrows pointing upward: prerequisite → successor)
for a, b in hasse:
x1, y1 = positions[a]
x2, y2 = positions[b]
ax.annotate(
"",
xy=(x2, y2),
xytext=(x1, y1),
arrowprops={
"arrowstyle": "->",
"color": edge_color,
"lw": 1.2,
"shrinkA": node_sizes[a] ** 0.5 / 2 + 2,
"shrinkB": node_sizes[b] ** 0.5 / 2 + 2,
},
zorder=1,
)
# Draw nodes
for item, (x, y) in positions.items():
ax.scatter(
x,
y,
s=node_sizes[item],
c=(node_colors or {}).get(item, node_color),
zorder=2,
edgecolors="white",
linewidth=2,
)
label = (
item if classes is None or len(classes[item]) == 1 else "\n".join(sorted(classes[item]))
)
label = (node_labels or {}).get(item, label)
ax.text(x, y, label, ha="center", va="center", fontsize=font_size, fontweight="bold")
if title is None:
title = (
"Surmise Relation (Quotient Hasse Diagram)"
if collapse_equivalent
else "Surmise Relation (Hasse Diagram)"
)
if title:
ax.set_title(title, fontsize=12, fontweight="bold")
if orientation == "vertical":
ax.set_ylabel("Prerequisite level")
ax.set_xticks([])
else:
ax.set_xlabel("Prerequisite level")
ax.set_yticks([])
ax.margins(0.2)
fig.tight_layout()
return fig