diff --git a/tests/analytical_visualization/tensor_sharding.py b/tests/analytical_visualization/tensor_sharding.py index 675ebcc..528e519 100644 --- a/tests/analytical_visualization/tensor_sharding.py +++ b/tests/analytical_visualization/tensor_sharding.py @@ -279,12 +279,38 @@ def draw_tensor_sharding(cfg: FullConfig, ax=None, my_pe: int = 0, row_label=f"rows split by {ffn_txt}", note="row-parallel") - # ── KV cache: (S_kv, H_kv*d_head), 2D shard ─ + # ── KV cache: (B, S_kv, H_kv*d_head), 2D shard on S_kv x heads ─ + # Batch (B) is a stacking dimension — each concurrent request holds + # its own KV slice per PE, so total KV per PE = B * (per-slice size). kv_row_splits = cp kv_col_splits = min(tp, m.h_kv) if m.h_kv >= tp else tp + _B = max(1, cfg.topo.b) + _batch_suffix = f", B={_B}" if _B > 1 else "" + _batch_note = (f"; batch B={_B} => {_B}x KV bytes per PE" + if _B > 1 else "") + # Draw B-1 offset "shadow" rectangles behind the KV cache to give + # a visual hint of the batch stacking dimension. Cap at 5 shadows + # so a large B doesn't overwhelm the diagram. + if _B > 1: + _kv_shadow_x = px(3) + _kv_shadow_y = py(0) + _n_shadows = min(_B - 1, 5) + _offset_step = 0.15 + for _si in range(_n_shadows, 0, -1): + _dx = _si * _offset_step + _dy = _si * _offset_step + _shadow = patches.Rectangle( + (_kv_shadow_x + _dx, _kv_shadow_y + _dy), + cell_w, cell_h, + facecolor="#f5f5f5", + edgecolor="#adb5bd", + linewidth=0.8, linestyle="--", alpha=0.55, + ) + ax.add_patch(_shadow) _draw_tensor(ax, px(3), py(0), cell_w, cell_h, "KV cache (K and V)", - f"({m.h_kv * m.d_head}, S_kv={cfg.topo.s_kv:,})", + f"({m.h_kv * m.d_head}, S_kv={cfg.topo.s_kv:,}" + f"{_batch_suffix})", row_splits=kv_row_splits, col_splits=kv_col_splits if show_physical else 1, my_row_shard=kv_row_shard, @@ -292,7 +318,8 @@ def draw_tensor_sharding(cfg: FullConfig, ax=None, my_pe: int = 0, row_label=f"rows: S split by {cp_txt}", col_label=f"cols: heads split by {tp_txt}", col_line_color=_TP_LINE_COLOR, - note="2D shard: seq axis (CP) x head axis (TP)", + note=f"2D shard: seq axis (CP) x head axis (TP)" + f"{_batch_note}", cell_annot=_kv_ann) # In non-physical mode, KV was drawn with col_splits=1; overlay TP col # splits as dotted lines to hint at the 2D nature. In physical mode @@ -311,8 +338,11 @@ def draw_tensor_sharding(cfg: FullConfig, ax=None, my_pe: int = 0, ax.set_aspect("auto") ax.axis("off") + _B_title = max(1, cfg.topo.b) + _b_title_suffix = f", B={_B_title}" if _B_title > 1 else "" ax.set_title( - f"Tensor sharding view (PE {my_pe}) | TP={tp}, CP={cp}, " + f"Tensor sharding view (PE {my_pe}) | TP={tp}, CP={cp}" + f"{_b_title_suffix}, " f"FFN scope={cfg.topo.ffn_shard_scope} | " f"blue = this PE's shard", fontsize=11, fontweight="bold",