Files
kernbench2/tests/analytical_visualization/topology_map.py
T
mukesh 9fdde44922 analytical-viz: interactive Streamlit tool for SIP transformer analysis
New tests/analytical_visualization/ module - an interactive dashboard
for exploring memory / latency tradeoffs of transformer inference on
the SIP architecture.

Highlights:
- 30+ model presets (Qwen, Llama 2/3/3.1, Mistral, Gemma 2, Phi 3,
  DeepSeek MLA, Mixtral/Qwen 3 MoE, Grok-1, ...)
- Placement toggles: TP and CP each on PE-level vs cube-level
- CP ring variant: K/V ring vs Q+O/m/l ring (prefill); in decode the
  O/m/l all-reduce is folded into S8 (no separate C1 row)
- SIP interconnect: ring / mesh2d / torus2d with matching link drawing
- Per-stage latency table with compute + memory + comm formulas,
  auto-scaled ns/us/ms, colored by dominant bound
- Ring attention loop indicator on the pipeline diagram (purple arc
  over S5-S8 with 'xN hops' badge)
- Tensor sharding view with optional physical PE/cube annotations
- Replication-waste + optimization-hints panel
- Save & compare configurations (config1, config2, ...): summary table
  plus side-by-side per-stage attention and FFN latency, best-in-row
  highlighting
- Symbol glossary with current values for every symbol used in formulas

Not tied to production sim_engine or runtime API; purely analytical
tooling for design-space exploration.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-07-26 22:12:08 -07:00

612 lines
24 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Draw the SIP topology diagram, color-coded by parallelism group.
Color scheme:
- CP ring: orange arrows (unchanged)
- Inter-SIP: red arrows across SIP boundaries
- PP stage: cube border color (each PP stage gets a distinct hue)
- TP group: PE fill color (each TP group within a cube gets one hue,
matches its cube's PP stage but slightly desaturated)
- EP group: a small badge on the cube (for MoE)
- DP replica: cubes of each DP replica get a hatched pattern
"""
from __future__ import annotations
import matplotlib.pyplot as plt
import matplotlib.patches as patches
from matplotlib.lines import Line2D
from .model_config import FullConfig
MESH_W = 4
MESH_H = 4
PE_COLS = 4
PE_ROWS = 2
def _group_label(cfg, group_idx: int) -> str:
"""Label for a cube-level group under the current placement.
- cp_placement=cube -> CP rank (or CP/TP if both on cube)
- cp_placement=pe, tp_placement=cube -> TP rank
- both on pe -> single group ("PE")
"""
parts = []
if cfg.topo.cp_placement == "cube":
cp_r = group_idx % max(1, cfg.topo.cp)
parts.append(f"CP{cp_r}")
if cfg.topo.tp_placement == "cube":
div = cfg.topo.cp if cfg.topo.cp_placement == "cube" else 1
tp_r = (group_idx // max(1, div)) % max(1, cfg.topo.tp)
parts.append(f"TP{tp_r}")
return "/".join(parts) if parts else "PE"
# Fully-distinct palette for (PP stage x CP rank) — 24 unique colors.
# Each cube's (pp, cp) pair maps to one entry; wraps if you exceed 24 groups.
_GROUP_PALETTE = [
"#e6194b", "#3cb44b", "#ffe119", "#4363d8",
"#f58231", "#911eb4", "#42d4f4", "#f032e6",
"#bfef45", "#fabed4", "#469990", "#dcbeff",
"#9a6324", "#fffac8", "#800000", "#aaffc3",
"#808000", "#ffd8b1", "#000075", "#a9a9a9",
"#004d40", "#c62828", "#4527a0", "#00695c",
]
# Kept for backward-compat / PP-stage arrows
_PP_HUES = _GROUP_PALETTE[:8]
def _shade(hue_hex: str, t: float) -> str:
"""Mix hue_hex toward white by t in [0, 1] (0 = original, 1 = white)."""
h = hue_hex.lstrip("#")
r, g, b = int(h[0:2], 16), int(h[2:4], 16), int(h[4:6], 16)
rn = int(r + (255 - r) * t)
gn = int(g + (255 - g) * t)
bn = int(b + (255 - b) * t)
return f"#{rn:02x}{gn:02x}{bn:02x}"
def _cp_color(pp_stage: int, cp_rank: int, cp_size: int) -> tuple[str, str]:
"""Return (border_color, pe_fill_color) for a cube at (pp_stage, cp_rank).
Each unique (pp, cp) pair gets a **fully distinct color** from the
24-entry palette. Border is the pure color; PE fill is a lightened
version so the cube border stays visible around the PE grid.
"""
idx = pp_stage * max(1, cp_size) + cp_rank
color = _GROUP_PALETTE[idx % len(_GROUP_PALETTE)]
pe_fill = _shade(color, 0.30) # slightly lightened for PE fill
return color, pe_fill
# Inactive (unused) colors.
_INACTIVE_CUBE_FACE = "#f7f7f7"
_INACTIVE_CUBE_EDGE = "#c8c8c8"
_INACTIVE_PE = "#eaeaea"
def _snake_path(n: int, mesh_w: int, mesh_h: int) -> list[int]:
path: list[int] = []
for r in range(mesh_h):
cols = range(mesh_w) if r % 2 == 0 else range(mesh_w - 1, -1, -1)
for c in cols:
path.append(r * mesh_w + c)
if len(path) == n:
return path
return path[:n]
def _rect_shape(n: int, max_w: int, max_h: int) -> tuple[int, int]:
"""Return (rows, cols) for n cubes packed as square as possible in max_wxmax_h."""
if n <= 0:
return (0, 0)
# Prefer near-square shapes; cap at mesh dims.
best = None
for cols in range(1, min(n, max_w) + 1):
rows = (n + cols - 1) // cols
if rows > max_h:
continue
# Score: penalise non-squareness + wasted cells
waste = rows * cols - n
aspect = abs(rows - cols)
score = aspect * 10 + waste
if best is None or score < best[0]:
best = (score, rows, cols)
if best is None:
# Fallback: single row
return (1, min(n, max_w))
return best[1], best[2]
def _pack_groups_2d(items_per_group: int, n_groups: int,
mesh_w: int, mesh_h: int,
layout_mode: str = "compact") -> list[tuple[int, int]]:
"""Return list of (row, col) positions for n_groups × items_per_group
cubes in a mesh_w x mesh_h mesh.
layout_mode:
- "compact": each group is a near-square rectangle; groups are also
tiled 2D as near-square. Ideal for 2x2, 2x4 etc.
- "linear": each group is a single row of cubes; groups tile down
as rows (original sequential layout).
"""
if n_groups == 0 or items_per_group == 0:
return []
if layout_mode == "linear":
# Each group is a single horizontal row of cubes.
tp_rows, tp_cols = 1, min(items_per_group, mesh_w)
else: # compact
tp_rows, tp_cols = _rect_shape(items_per_group, mesh_w, mesh_h)
if tp_rows == 0:
return []
# Group grid: how many groups per row/col of GROUPS
max_group_cols = max(1, mesh_w // tp_cols)
max_group_rows = max(1, mesh_h // tp_rows)
if layout_mode == "linear":
# Row-major: fill each mesh row with groups, then wrap down.
g_cols = max_group_cols
g_rows = max_group_rows
else:
g_rows, g_cols = _rect_shape(n_groups, max_group_cols, max_group_rows)
positions: list[tuple[int, int]] = []
for g in range(n_groups):
gr = g // g_cols
gc = g % g_cols
base_row = gr * tp_rows
base_col = gc * tp_cols
for i in range(items_per_group):
r = i // tp_cols
c = i % tp_cols
row = base_row + r
col = base_col + c
if row >= mesh_h or col >= mesh_w:
fallback = (g * items_per_group + i)
row = fallback // mesh_w
col = fallback % mesh_w
positions.append((row, col))
return positions
def _draw_one_sip(ax, cfg: FullConfig, sip_idx: int,
sip_x0: float, sip_y0: float,
sip_width: float, sip_height: float,
cube_pp_cp: dict[int, tuple[int, int]],
ring_paths: dict[int, list[int]],
tp_group_cubes: dict[int, list[int]]):
"""Draw one SIP with PP hue + CP shade + TP group boundaries."""
cube_size = min(sip_width / MESH_W, sip_height / MESH_H) * 0.85
cube_gap = cube_size * 0.15
total_w = MESH_W * (cube_size + cube_gap) - cube_gap
total_h = MESH_H * (cube_size + cube_gap) - cube_gap
ox = sip_x0 + (sip_width - total_w) / 2
oy = sip_y0 + (sip_height - total_h) / 2
sip_rect = patches.FancyBboxPatch(
(sip_x0, sip_y0), sip_width, sip_height,
boxstyle="round,pad=0.02", facecolor="#f8f9fa",
edgecolor="#212529", linewidth=1.5,
)
ax.add_patch(sip_rect)
ax.text(sip_x0 + sip_width / 2, sip_y0 + sip_height + 0.05,
f"SIP {sip_idx}", ha="center", va="bottom",
fontsize=10, fontweight="bold")
def _cube_center(cube_local: int) -> tuple[float, float]:
r = cube_local // MESH_W
c = cube_local % MESH_W
x = ox + c * (cube_size + cube_gap)
y = oy + (MESH_H - 1 - r) * (cube_size + cube_gap)
return x, y
for cube_local in range(MESH_W * MESH_H):
x, y = _cube_center(cube_local)
info = cube_pp_cp.get(cube_local, None)
is_used = info is not None
if is_used:
pp_stage, cp_rank = info
cube_edge, pe_fill = _cp_color(
pp_stage, cp_rank,
max(1, cfg.topo.inter_cube_dims),
)
cube_face = _shade(cube_edge, 0.85)
cube_lw = 2.0
else:
cube_edge = _INACTIVE_CUBE_EDGE
cube_face = _INACTIVE_CUBE_FACE
cube_lw = 0.6
pe_fill = _INACTIVE_PE
rect = patches.FancyBboxPatch(
(x, y), cube_size, cube_size,
boxstyle="round,pad=0.01",
facecolor=cube_face, edgecolor=cube_edge, linewidth=cube_lw,
)
ax.add_patch(rect)
if is_used:
pp_stage, cp_rank = info
# Small cube ID above
ax.text(x + cube_size / 2, y + cube_size + 0.01,
f"cube {cube_local}", ha="center", va="bottom",
fontsize=5.5, color=cube_edge, fontweight="bold")
# Group label (CP/TP depending on placement)
_lbl = _group_label(cfg, cp_rank)
ax.text(x + 0.03, y + cube_size - 0.03,
_lbl, ha="left", va="top",
fontsize=11, color=cube_edge, fontweight="bold",
bbox=dict(boxstyle="round,pad=0.15",
facecolor="white", edgecolor=cube_edge,
linewidth=1.0, alpha=0.9))
if cfg.topo.pp > 1:
ax.text(x + cube_size - 0.03, y + cube_size - 0.03,
f"PP{pp_stage}", ha="right", va="top",
fontsize=9, color=cube_edge, fontweight="bold",
bbox=dict(boxstyle="round,pad=0.1",
facecolor="white", edgecolor=cube_edge,
linewidth=0.8, alpha=0.9))
else:
ax.text(x + cube_size / 2, y + cube_size + 0.008,
f"{cube_local}", ha="center", va="bottom",
fontsize=6, color=cube_edge)
# PEs
pe_pad = cube_size * 0.08
pe_gap = cube_size * 0.02
inner_w = cube_size - 2 * pe_pad
inner_h = cube_size - 2 * pe_pad
pe_w = (inner_w - (PE_COLS - 1) * pe_gap) / PE_COLS
pe_h = (inner_h - (PE_ROWS - 1) * pe_gap) / PE_ROWS
pes_used = cfg.topo.pes_per_cube_used if is_used else 0
for pr in range(PE_ROWS):
for pc in range(PE_COLS):
pe_id = pr * PE_COLS + pc
px = x + pe_pad + pc * (pe_w + pe_gap)
py = y + pe_pad + (PE_ROWS - 1 - pr) * (pe_h + pe_gap)
fill = pe_fill if pe_id < pes_used else _INACTIVE_PE
pe_rect = patches.Rectangle(
(px, py), pe_w, pe_h,
facecolor=fill, edgecolor="#666", linewidth=0.3,
)
ax.add_patch(pe_rect)
if is_used and cfg.topo.ep > 1:
ax.text(x + cube_size - 0.02, y + cube_size - 0.02,
f"EP{cfg.topo.ep}", ha="right", va="top",
fontsize=6, color="#5f0f40", fontweight="bold")
# (CP ring arrows removed per user request)
# Draw TP group bounding boxes (RED dashed) when TP > 8
for tp_g, cubes in tp_group_cubes.items():
if len(cubes) <= 1:
continue # TP fits in one cube, redundant with cube border
# Bounding box around all cubes in this TP group
xs, ys = [], []
for cube_local in cubes:
x, y = _cube_center(cube_local)
xs.extend([x, x + cube_size])
ys.extend([y, y + cube_size])
pad = 0.03
bbox_x = min(xs) - pad
bbox_y = min(ys) - pad
bbox_w = (max(xs) - min(xs)) + 2 * pad
bbox_h = (max(ys) - min(ys)) + 2 * pad
tp_rect = patches.Rectangle(
(bbox_x, bbox_y), bbox_w, bbox_h,
fill=False, edgecolor="#d90429", linewidth=1.8,
linestyle="--",
)
ax.add_patch(tp_rect)
ax.text(bbox_x + bbox_w / 2, bbox_y - 0.05,
f"TP group {tp_g}", ha="center", va="top",
color="#d90429", fontsize=7, fontweight="bold")
def _sip_grid_layout(n_sips: int, topology: str) -> tuple[int, int]:
"""(rows, cols) for arranging SIPs on the page.
- ring: horizontal chain (1 x N)
- mesh2d / torus2d: near-square grid
"""
if n_sips <= 0:
return 0, 0
if topology == "ring" or n_sips <= 2:
return 1, n_sips
import math
cols = int(math.ceil(math.sqrt(n_sips)))
rows = (n_sips + cols - 1) // cols
return rows, cols
def _draw_sip_links(ax, sip_xy: list[tuple[float, float]],
topology: str, sip_w: float, sip_h: float,
grid_rows: int, grid_cols: int):
"""Draw inter-SIP interconnect lines.
Solid lines: adjacent (grid-neighbor) links.
Dashed lines: wrap-around links (ring, torus).
"""
n = len(sip_xy)
if n <= 1:
return
color = "#c62828"
lw = 2.0
def center(i):
x, y = sip_xy[i]
return x + sip_w / 2, y + sip_h / 2
if topology == "ring":
# 1D chain (adjacent solid)
for i in range(n - 1):
x1, y1 = center(i)
x2, y2 = center(i + 1)
ax.plot([x1 + sip_w / 2, x2 - sip_w / 2], [y1, y2],
color=color, linewidth=lw)
# wrap-around from last back to first (dashed arc below)
if n > 2:
x_first, y_first = center(0)
x_last, y_last = center(n - 1)
arc_y = min(y_first, y_last) - sip_h / 2 - 0.35
ax.plot([x_first, x_first], [y_first - sip_h / 2, arc_y],
color=color, linewidth=lw, linestyle="--")
ax.plot([x_first, x_last], [arc_y, arc_y],
color=color, linewidth=lw, linestyle="--")
ax.plot([x_last, x_last], [arc_y, y_last - sip_h / 2],
color=color, linewidth=lw, linestyle="--")
elif n == 2:
# 2 SIPs: single link
x1, y1 = center(0)
x2, y2 = center(1)
# wrap link: dashed loop below
arc_y = y1 - sip_h / 2 - 0.35
ax.plot([x1, x1], [y1 - sip_h / 2, arc_y],
color=color, linewidth=lw, linestyle="--")
ax.plot([x1, x2], [arc_y, arc_y],
color=color, linewidth=lw, linestyle="--")
ax.plot([x2, x2], [arc_y, y2 - sip_h / 2],
color=color, linewidth=lw, linestyle="--")
return
# mesh2d / torus2d: grid neighbours
def idx(r, c):
return r * grid_cols + c
for r in range(grid_rows):
for c in range(grid_cols):
i = idx(r, c)
if i >= n:
continue
# right
if c + 1 < grid_cols and idx(r, c + 1) < n:
x1, y1 = center(i)
x2, y2 = center(idx(r, c + 1))
ax.plot([x1 + sip_w / 2, x2 - sip_w / 2], [y1, y2],
color=color, linewidth=lw)
# down (visually — grid row+1 is drawn below)
if r + 1 < grid_rows and idx(r + 1, c) < n:
x1, y1 = center(i)
x2, y2 = center(idx(r + 1, c))
ax.plot([x1, x2], [y1 - sip_h / 2, y2 + sip_h / 2],
color=color, linewidth=lw)
if topology == "torus2d":
# row wrap: last col -> first col in each row (dashed arc BELOW that row)
for r in range(grid_rows):
left_i = idx(r, 0)
right_i = idx(r, grid_cols - 1)
if left_i >= n or right_i >= n or left_i == right_i:
continue
x1, y1 = center(left_i)
x2, y2 = center(right_i)
arc_y = min(y1, y2) - sip_h / 2 - 0.35
ax.plot([x1, x1], [y1 - sip_h / 2, arc_y],
color=color, linewidth=lw, linestyle="--")
ax.plot([x1, x2], [arc_y, arc_y],
color=color, linewidth=lw, linestyle="--")
ax.plot([x2, x2], [arc_y, y2 - sip_h / 2],
color=color, linewidth=lw, linestyle="--")
# col wrap: last row -> first row in each col (dashed arc to the RIGHT)
for c in range(grid_cols):
top_i = idx(0, c)
bot_i = idx(grid_rows - 1, c)
if top_i >= n or bot_i >= n or top_i == bot_i:
continue
x1, y1 = center(top_i) # top row (higher y)
x2, y2 = center(bot_i) # bottom row (lower y)
arc_x = max(x1, x2) + sip_w / 2 + 0.35
ax.plot([x1 + sip_w / 2, arc_x], [y1, y1],
color=color, linewidth=lw, linestyle="--")
ax.plot([arc_x, arc_x], [y1, y2],
color=color, linewidth=lw, linestyle="--")
ax.plot([arc_x, x2 + sip_w / 2], [y2, y2],
color=color, linewidth=lw, linestyle="--")
def draw_topology(cfg: FullConfig, ax=None, layout_mode: str = "linear"):
"""Draw one or more SIPs. Each (PP,CP) group gets a unique color.
SIP grid arrangement + inter-SIP links depend on cfg.topo.sip_topology:
"ring" -> 1D horizontal chain, wrap arc below
"mesh2d" -> near-square grid, no wrap
"torus2d" -> grid + row/col wrap arcs
"""
sips_used = max(1, cfg.topo.sips_used)
total_pes = cfg.topo.total_pes
pes_per_stage = cfg.topo.pes_per_stage # CP*TP
n_replicas = cfg.topo.dp
n_pp_stages = cfg.topo.pp
pes_per_cube = cfg.topo.pes_per_cube_hw
# Placement-aware: how many cubes per "group" (was: cubes-per-TP).
# A group is one cube-level unit (one CP rank if cp on cube, or one TP
# rank if only tp on cube, etc.). Cubes per group == intra-cube spill.
cubes_per_group = max(
1, (cfg.topo.intra_cube_dims + pes_per_cube - 1) // pes_per_cube
)
n_groups_per_stage = max(1, cfg.topo.inter_cube_dims)
cubes_per_stage = cfg.topo.cubes_per_stage
cubes_used = cfg.topo.cubes_used
cubes_per_sip = MESH_W * MESH_H
sip_topo = cfg.topo.sip_topology
grid_rows, grid_cols = _sip_grid_layout(sips_used, sip_topo)
sip_w = 4.0
sip_h = 4.0
sip_gap = 0.6
fig_w = grid_cols * sip_w + (grid_cols - 1) * sip_gap + 1.5
fig_h = grid_rows * sip_h + (grid_rows - 1) * sip_gap + 2.0
if ax is None:
fig, ax = plt.subplots(figsize=(max(6, fig_w), max(6, fig_h)))
else:
fig = ax.figure
# 2D packing: each group's cubes form a rectangle inside the SIP mesh.
# If more groups than fit in one SIP, spill to next SIP.
tp_rows, tp_cols = _rect_shape(cubes_per_group, MESH_W, MESH_H)
if tp_rows == 0:
tp_rows, tp_cols = 1, 1
groups_per_sip = max(1, (MESH_W // tp_cols) * (MESH_H // tp_rows))
cube_pp_cp_global: dict[int, tuple[int, int]] = {}
ring_globals_by_pp: dict[int, list[int]] = {}
tp_group_cubes_global: dict[int, list[int]] = {}
current_sip = 0
for rep in range(n_replicas):
for pp_s in range(n_pp_stages):
# Chunk cube-level groups into per-SIP batches.
for chunk_start in range(0, n_groups_per_stage, groups_per_sip):
chunk = min(groups_per_sip, n_groups_per_stage - chunk_start)
positions = _pack_groups_2d(cubes_per_group, chunk,
MESH_W, MESH_H,
layout_mode=layout_mode)
for local_gi in range(chunk):
group_idx = chunk_start + local_gi
tp_gid = ((rep * n_pp_stages) + pp_s) * n_groups_per_stage + group_idx
cubes_here: list[int] = []
for i in range(cubes_per_group):
idx_in_group = local_gi * cubes_per_group + i
r, c = positions[idx_in_group]
local_cube = r * MESH_W + c
gc = current_sip * cubes_per_sip + local_cube
cube_pp_cp_global[gc] = (pp_s, group_idx)
cubes_here.append(gc)
tp_group_cubes_global[tp_gid] = cubes_here
ring_globals_by_pp.setdefault(pp_s, []).append(cubes_here[0])
current_sip += 1 # next SIP for the next chunk (or (rep, pp))
# Partition (pp, cp) map by SIP using cube_positions_global
cubes_by_sip: list[dict[int, tuple[int, int]]] = [dict() for _ in range(sips_used)]
for gc, ppcp in cube_pp_cp_global.items():
sip_idx = gc // cubes_per_sip
local_cube = gc % cubes_per_sip
if sip_idx < sips_used:
cubes_by_sip[sip_idx][local_cube] = ppcp
# Per-SIP snake-path CP rings (organized by pp_stage)
ring_by_sip: list[dict[int, list[int]]] = [dict() for _ in range(sips_used)]
for pp_s, gc_list in ring_globals_by_pp.items():
by_sip_local: dict[int, list[int]] = {}
for gc in gc_list:
si = gc // cubes_per_sip
lc = gc % cubes_per_sip
by_sip_local.setdefault(si, []).append(lc)
for si, locals_list in by_sip_local.items():
if si < sips_used:
# Snake through the packed positions in order
ring_by_sip[si][pp_s] = locals_list
# TP group cubes by SIP
tp_by_sip: list[dict[int, list[int]]] = [dict() for _ in range(sips_used)]
for tp_g, gc_list in tp_group_cubes_global.items():
by_sip_local: dict[int, list[int]] = {}
for gc in gc_list:
si = gc // cubes_per_sip
lc = gc % cubes_per_sip
by_sip_local.setdefault(si, []).append(lc)
for si, locals_list in by_sip_local.items():
if si < sips_used:
tp_by_sip[si][tp_g] = locals_list
# Compute (x, y) top-left for each SIP based on grid layout.
# Row 0 is drawn at the TOP visually, so y decreases with row index.
sip_xy: list[tuple[float, float]] = []
for sip_idx in range(sips_used):
r = sip_idx // grid_cols
c = sip_idx % grid_cols
x = c * (sip_w + sip_gap)
y = (grid_rows - 1 - r) * (sip_h + sip_gap)
sip_xy.append((x, y))
for sip_idx in range(sips_used):
x0, y0 = sip_xy[sip_idx]
_draw_one_sip(
ax, cfg, sip_idx,
sip_x0=x0, sip_y0=y0,
sip_width=sip_w, sip_height=sip_h,
cube_pp_cp=cubes_by_sip[sip_idx],
ring_paths=ring_by_sip[sip_idx],
tp_group_cubes=tp_by_sip[sip_idx],
)
# Inter-SIP interconnect links per user-selected topology.
_draw_sip_links(ax, sip_xy, sip_topo, sip_w, sip_h, grid_rows, grid_cols)
total_w = grid_cols * sip_w + (grid_cols - 1) * sip_gap
total_h = grid_rows * sip_h + (grid_rows - 1) * sip_gap
# Extra room: bottom for wrap arcs, right for torus col-wrap arcs, top for SIP labels
ax.set_xlim(-0.3, total_w + 0.9)
ax.set_ylim(-0.9, total_h + 0.7)
ax.set_aspect("equal")
ax.axis("off")
# Legend — one entry per unique (PP, cube-group) + TP boundary + SIP links
legend_handles = []
for pp_s in range(n_pp_stages):
for gi in range(n_groups_per_stage):
color, _ = _cp_color(pp_s, gi, n_groups_per_stage)
legend_handles.append(Line2D(
[], [], color=color, marker="s", linestyle="",
markersize=10,
label=(f"PP{pp_s} " if n_pp_stages > 1 else "")
+ _group_label(cfg, gi),
))
if cfg.topo.tp > cfg.topo.pes_per_cube_hw:
legend_handles.append(Line2D(
[], [], color="#d90429", marker="s", linestyle="--",
markersize=10, markerfacecolor="none", markeredgewidth=1.5,
label="TP group (dashed red)",
))
if cfg.topo.ep > 1:
legend_handles.append(Line2D(
[], [], color="#5f0f40", marker="s", linestyle="",
markersize=8, label=f"EP={cfg.topo.ep}",
))
if sips_used > 1:
_topo_label = {
"ring": "Inter-SIP: ring (chain + wrap)",
"mesh2d": "Inter-SIP: 2D mesh (no wrap)",
"torus2d": "Inter-SIP: 2D torus (grid + wrap)",
}.get(sip_topo, f"Inter-SIP: {sip_topo}")
legend_handles.append(Line2D(
[], [], color="#c62828", linewidth=2.0,
label=_topo_label,
))
# Legend below the figure; wrap ncols so it's readable
n_items = len(legend_handles)
ncol = min(6, max(3, n_items))
ax.legend(handles=legend_handles, loc="upper center",
bbox_to_anchor=(0.5, -0.02), ncol=ncol,
fontsize=8, frameon=False)
title = (
f"Layout: {total_pes} PEs = {n_replicas} DP replica(s) x "
f"{n_pp_stages} PP stage(s) x {cubes_per_stage} cubes (CP={cfg.topo.cp}) "
f"x TP={cfg.topo.tp} | cubes: {cubes_used}, SIPs: {sips_used}"
)
ax.set_title(title, fontsize=10)
return fig