Source code for knowledgespaces.viz.hasse

"""
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