Keiro / architecture /sparse_moe /visualization.py
iamrahulreddy's picture
add: sparse_moe architecture source
d8f717b verified
Raw
History Blame Contribute Delete
11.4 kB
from __future__ import annotations
import warnings
from typing import List, Optional, Tuple, Union, Dict, Any
from pathlib import Path
import numpy as np
try:
import matplotlib.pyplot as plt
import matplotlib.colors as mcolors
from matplotlib.gridspec import GridSpec
HAS_MATPLOTLIB = True
except ImportError:
HAS_MATPLOTLIB = False
warnings.warn(
"Matplotlib not installed. Visualization features unavailable. "
"Install with: pip install matplotlib"
)
def _check_matplotlib():
if not HAS_MATPLOTLIB:
raise RuntimeError(
"Matplotlib is required for visualization. "
"Install with: pip install matplotlib"
)
class KeiroPalette:
# Primary line colors
PRIMARY = {
"dense": "#E74C3C", # Red - Dense baseline ("Before")
"moe": "#3498DB", # Blue - Sparse MoE ("After")
"expert_avg": "#9B59B6",# Purple - Expert average
"memory": "#2ECC71", # Green - Memory bounds
}
# Generative expert spectrum for heatmap/routing
EXPERTS = [
"#3498DB", "#2980B9", "#1ABC9C", "#27AE60",
"#F39C12", "#D35400", "#E74C3C", "#8E44AD"
]
BG_LIGHT = "#FAFAFA"
BG_DARK = "#1A1A2E"
GRID_LIGHT = "#E0E0E0"
GRID_DARK = "#2D2D44"
class BasePlot:
def __init__(self, figsize=(12, 8), dpi=150, theme="dark", title=None):
_check_matplotlib()
self.figsize = figsize
self.dpi = dpi
self.theme = theme
self.fig, self.ax = plt.subplots(figsize=figsize)
self._apply_theme()
if title:
self.ax.set_title(title, fontsize=16, fontweight='bold', pad=20)
def _apply_theme(self):
bg = KeiroPalette.BG_DARK if self.theme == "dark" else KeiroPalette.BG_LIGHT
grid = KeiroPalette.GRID_DARK if self.theme == "dark" else KeiroPalette.GRID_LIGHT
fg = 'white' if self.theme == "dark" else 'black'
self.fig.patch.set_facecolor(bg)
self.ax.set_facecolor(bg)
self.ax.tick_params(colors=fg)
self.ax.xaxis.label.set_color(fg)
self.ax.yaxis.label.set_color(fg)
self.ax.title.set_color(fg)
self.ax.grid(True, alpha=0.2, color=grid)
for spine in self.ax.spines.values():
spine.set_color(grid)
def save(self, filepath: Union[str, Path]):
filepath = Path(filepath)
filepath.parent.mkdir(parents=True, exist_ok=True)
self.fig.savefig(filepath, dpi=self.dpi, bbox_inches="tight", facecolor=self.fig.get_facecolor())
def close(self):
plt.close(self.fig)
class ResourceUtilizationPlot(BasePlot):
def __init__(self, title="Resource Utilization (Before vs After)", **kwargs):
super().__init__(title=title, **kwargs)
self.ax.set_xlabel("Time (seconds)", fontsize=12)
self.ax2 = self.ax.twinx()
self.ax.set_ylabel("Memory Allocated (MB)", fontsize=12)
self.ax2.set_ylabel("GPU Utilization (%)", fontsize=12)
if self.theme == "dark":
self.ax2.tick_params(colors='white')
self.ax2.yaxis.label.set_color('white')
for spine in self.ax2.spines.values():
spine.set_color(KeiroPalette.GRID_DARK)
def add_trace(self, time_sec: List[float], values: List[float], label: str, metric: str = "memory"):
color = KeiroPalette.PRIMARY["dense"] if "Before" in label or "Dense" in label else KeiroPalette.PRIMARY["moe"]
linestyle = "-" if metric == "memory" else "--"
axis = self.ax if metric == "memory" else self.ax2
axis.plot(
time_sec, values, label=label,
color=color, linestyle=linestyle, linewidth=2.5, alpha=0.8
)
def finalize(self):
lines1, labels1 = self.ax.get_legend_handles_labels()
lines2, labels2 = self.ax2.get_legend_handles_labels()
self.ax2.legend(lines1 + lines2, labels1 + labels2, loc="best", framealpha=0.8)
plt.tight_layout()
class KeiroDashboard:
def __init__(self, figsize=(18, 12), dpi=150, theme="dark"):
_check_matplotlib()
self.dpi = dpi
self.theme = theme
self.fig, self.axes = plt.subplots(2, 2, figsize=figsize)
self._apply_theme()
def _apply_theme(self):
bg = KeiroPalette.BG_DARK if self.theme == "dark" else KeiroPalette.BG_LIGHT
grid = KeiroPalette.GRID_DARK if self.theme == "dark" else KeiroPalette.GRID_LIGHT
fg = 'white' if self.theme == "dark" else 'black'
self.fig.patch.set_facecolor(bg)
for ax in self.axes.flat:
ax.set_facecolor(bg)
ax.tick_params(colors=fg)
ax.xaxis.label.set_color(fg)
ax.yaxis.label.set_color(fg)
ax.title.set_color(fg)
ax.grid(True, alpha=0.2, color=grid)
for spine in ax.spines.values():
spine.set_color(grid)
def plot_memory_scaling(self, ax_idx=(0,0), seq_lens=None, data_dict=None):
ax = self.axes[ax_idx]
ax.set_title("Peak Memory vs Sequence Length", fontsize=14, fontweight='bold')
ax.set_xlabel("Sequence Length")
ax.set_ylabel("Memory (MB)")
if seq_lens and data_dict:
for k, v in data_dict.items():
color = KeiroPalette.PRIMARY["dense"] if "Dense" in k else KeiroPalette.PRIMARY["moe"]
ax.plot(seq_lens, v, label=k, color=color, marker='o', linewidth=2)
ax.legend()
def plot_throughput(self, ax_idx=(0,1), seq_lens=None, data_dict=None):
ax = self.axes[ax_idx]
ax.set_title("Inference Throughput (tokens/sec)", fontsize=14, fontweight='bold')
ax.set_xlabel("Sequence Length")
ax.set_ylabel("Throughput")
if seq_lens and data_dict:
for k, v in data_dict.items():
color = KeiroPalette.PRIMARY["dense"] if "Dense" in k else KeiroPalette.PRIMARY["moe"]
ax.plot(seq_lens, v, label=k, color=color, marker='s', linewidth=2)
ax.legend()
def plot_expert_load(self, ax_idx=(1,0), expert_distribution=None):
ax = self.axes[ax_idx]
ax.set_title("MoE Expert Load Balancing", fontsize=14, fontweight='bold')
ax.set_xlabel("Expert ID")
ax.set_ylabel("Tokens Assigned (%)")
if expert_distribution:
x = np.arange(len(expert_distribution))
colors = [KeiroPalette.EXPERTS[i % len(KeiroPalette.EXPERTS)] for i in x]
total = sum(expert_distribution)
pcts = [100.0 * c / total for c in expert_distribution] if total > 0 else expert_distribution
ax.bar(x, pcts, color=colors, alpha=0.8)
ax.set_xticks(x)
ax.set_xticklabels([f"E{i}" for i in x])
ax.axhline(100.0 / len(expert_distribution), color='gray', linestyle='--', label='Perfect Balance')
ax.legend()
def plot_speedup(self, ax_idx=(1,1), seq_lens=None, base_time=None, moe_time=None):
ax = self.axes[ax_idx]
ax.set_title("MoE Speedup vs Dense", fontsize=14, fontweight='bold')
ax.set_xlabel("Sequence Length")
ax.set_ylabel("Speedup (x)")
ax.axhline(1.0, color='gray', linestyle='--', alpha=0.5)
if seq_lens and base_time and moe_time:
speedups = [b/m if m > 0 else 0 for b, m in zip(base_time, moe_time)]
ax.plot(seq_lens, speedups, color=KeiroPalette.PRIMARY["expert_avg"], marker='D', linewidth=2, label="Speedup")
ax.legend()
def save(self, filepath: Union[str, Path]):
filepath = Path(filepath)
filepath.parent.mkdir(parents=True, exist_ok=True)
plt.tight_layout()
self.fig.savefig(filepath, dpi=self.dpi, bbox_inches="tight", facecolor=self.fig.get_facecolor())
def close(self):
plt.close(self.fig)
class ColorPalette(KeiroPalette):
pass
class DomainScorePlot(BasePlot):
def __init__(self, figsize=(10, 6), **kwargs):
super().__init__(figsize=figsize, title="Per-Domain Perplexity", **kwargs)
def plot_comparison(self, rows: List[Dict], include_dense: bool = False):
if not rows: return
domains = [r["domain"] for r in rows]
before = [r["ppl_before"] for r in rows]
after = [r["ppl_after"] for r in rows]
x = np.arange(len(domains))
width = 0.35 if not include_dense else 0.25
self.ax.bar(x - width/2, before, width, label='Before (Dense)', color=KeiroPalette.PRIMARY["dense"])
self.ax.bar(x + width/2, after, width, label='After (MoE)', color=KeiroPalette.PRIMARY["moe"])
if include_dense:
dense = [r.get("ppl_dense", 0) for r in rows]
self.ax.bar(x + 1.5*width, dense, width, label='Dense Baseline', color=KeiroPalette.PRIMARY["expert_avg"])
self.ax.set_xticks(x)
self.ax.set_xticklabels(domains, rotation=45, ha='right')
self.ax.set_ylabel("Perplexity (Lower is better)")
self.ax.legend()
self.fig.tight_layout()
class TrainingConvergencePlot(BasePlot):
def __init__(self, figsize=(10, 6), **kwargs):
super().__init__(figsize=figsize, title="Training Convergence", **kwargs)
def plot_history(self, history: Dict):
train_loss = history.get("train_loss", [])
val_loss = history.get("val_loss", [])
if train_loss:
self.ax.plot(train_loss, label="Train Loss", color=KeiroPalette.PRIMARY["dense"])
if val_loss:
if len(val_loss) < len(train_loss):
x_val = np.linspace(0, len(train_loss)-1, len(val_loss))
self.ax.plot(x_val, val_loss, label="Val Loss", marker='o', color=KeiroPalette.PRIMARY["moe"])
else:
self.ax.plot(val_loss, label="Val Loss", color=KeiroPalette.PRIMARY["moe"])
self.ax.set_xlabel("Steps (or Epochs)")
self.ax.set_ylabel("Cross Entropy Loss")
self.ax.legend()
class ExpertRoutingHeatmap(BasePlot):
def __init__(self, figsize=(12, 8), **kwargs):
super().__init__(figsize=figsize, title="Expert Routing by Domain", **kwargs)
def plot_routing(self, spec_dict: Dict):
affinity = spec_dict.get("affinity")
domains = spec_dict.get("domains")
labels = spec_dict.get("expert_labels")
if affinity is None or domains is None:
return
# Convert torch tensor to numpy
if hasattr(affinity, "cpu"):
matrix = affinity.cpu().numpy()
else:
matrix = np.array(affinity)
# Plotting
im = self.ax.imshow(matrix, aspect="auto", cmap="viridis")
self.ax.set_xticks(range(len(domains)))
self.ax.set_xticklabels(domains, rotation=45, ha='right')
# Only show individual expert labels if there aren't too many
if labels and len(labels) <= 64:
self.ax.set_yticks(range(len(labels)))
self.ax.set_yticklabels(labels, fontsize=6)
else:
self.ax.set_ylabel(f"{len(labels)} Layer-Experts")
self.ax.set_yticks([]) # Hide Y-axis labels for readability if too dense
self.fig.colorbar(im, ax=self.ax, fraction=0.046, pad=0.04)
self.fig.tight_layout()