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>
This commit is contained in:
2026-07-26 22:11:33 -07:00
parent ae131aa2af
commit 9fdde44922
12 changed files with 4515 additions and 0 deletions
@@ -0,0 +1,611 @@
"""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