Keiro / architecture /sparse_moe /analysis.py
iamrahulreddy's picture
add: sparse_moe architecture source
d8f717b verified
Raw
History Blame Contribute Delete
5.32 kB
from __future__ import annotations
from typing import Dict, Optional
import torch
import torch.nn as nn
from torch.utils.data import DataLoader
from .evaluation import _named_moe_layers
from .utils import build_labels
# Core: expert–domain affinity matrix
@torch.no_grad()
def compute_specialization_matrix(model: nn.Module, domain_loaders: Dict[str, DataLoader], device: Optional[torch.device] = None) -> Dict[str, object]:
device = device or next(model.parameters()).device
was_training = model.training
model.eval()
try:
moe_layers = _named_moe_layers(model)
if not moe_layers:
raise ValueError("No SparseMoELayer found in model")
domains = list(domain_loaders.keys())
num_domains = len(domains)
# Accumulate weighted expert fractions per domain, keeping layers separate.
layer_affinity = {
layer_name: torch.zeros(layer.num_experts, num_domains)
for layer_name, layer in moe_layers
}
layer_weight_per_domain = {
layer_name: torch.zeros(num_domains) for layer_name, _ in moe_layers
}
for d_idx, domain in enumerate(domains):
loader = domain_loaders[domain]
for batch in loader:
batch = {k: v.to(device) for k, v in batch.items()}
labels = build_labels(batch)
_ = model(**batch, labels=labels)
for layer_name, module in moe_layers:
stats = getattr(module, "_last_stats", None)
if stats is None:
continue
fracs = stats.get("expert_fractions", [])
if len(fracs) != module.num_experts:
continue
num_tokens = float(stats.get("num_tokens", 1))
frac_t = torch.tensor(fracs, dtype=torch.float32)
layer_affinity[layer_name][:, d_idx] += frac_t * num_tokens
layer_weight_per_domain[layer_name][d_idx] += num_tokens
# Normalise and flatten rows as unique layer-expert slots.
affinity_rows = []
expert_labels = []
for layer_name, module in moe_layers:
layer_matrix = layer_affinity[layer_name]
for d_idx in range(num_domains):
weight = layer_weight_per_domain[layer_name][d_idx]
if weight > 0:
layer_matrix[:, d_idx] /= weight
affinity_rows.append(layer_matrix)
for expert_idx in range(module.num_experts):
expert_labels.append(f"{layer_name}:E{expert_idx}")
affinity = torch.cat(affinity_rows, dim=0) if affinity_rows else torch.zeros(0, num_domains)
# Specialisation
row_max = affinity.max(dim=1).values
row_mean = affinity.mean(dim=1).clamp(min=1e-9)
specialisation = row_max / row_mean
# JS divergence
col_sums = affinity.sum(dim=0, keepdim=True).clamp(min=1e-9)
col_norm = affinity / col_sums
divergence = _pairwise_js(col_norm)
return {
"affinity": affinity,
"layer_affinity": layer_affinity,
"domains": domains,
"expert_labels": expert_labels,
"specialisation": specialisation,
"divergence": divergence,
}
finally:
model.train(was_training)
# Helpers
def _kl_divergence(P: torch.Tensor, q: torch.Tensor) -> torch.Tensor:
P_safe = P.clamp(min=1e-9)
q_safe = q.clamp(min=1e-9)
return (P_safe * (P_safe / q_safe).log()).sum()
def _pairwise_js(col_norm: torch.Tensor) -> torch.Tensor:
D = col_norm.shape[1]
js = torch.zeros(D, D)
for i in range(D):
for j in range(i + 1, D):
p = col_norm[:, i]
q = col_norm[:, j]
m = 0.5 * (p + q)
jsd = 0.5 * _kl_divergence(p, m) + 0.5 * _kl_divergence(q, m)
js[i, j] = js[j, i] = jsd.item()
return js
# Pretty-print
def print_specialization_report(result: Dict[str, object]):
affinity = result["affinity"]
domains = result["domains"]
expert_labels = result.get("expert_labels", [f"E{i}" for i in range(affinity.shape[0])])
spec = result["specialisation"]
div = result["divergence"]
E, D = affinity.shape
print("\n" + "=" * 62)
print(" Expert Specialisation Report")
print("=" * 62)
# Affinity matrix
header = " Layer-Expert " + "".join(f" {d:>8s}" for d in domains) + " Spec"
print(header)
print(" " + "-" * (len(header) - 2))
for e in range(E):
row = f" {expert_labels[e]:<28s}"
for d in range(D):
row += f" {affinity[e, d]:>8.3f}"
row += f" {spec[e]:>5.2f}"
print(row)
# Divergence
print(f"\n JS Divergence (domain routing distances):")
div_header = " " + "".join(f" {d:>8s}" for d in domains)
print(div_header)
for i, d_name in enumerate(domains):
row = f" {d_name:<8s}"
for j in range(D):
row += f" {div[i, j]:>8.4f}"
print(row)
print("=" * 62 + "\n")