| 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
|
|
|
|
|
|
|
| @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)
|
|
|
|
|
| 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
|
|
|
|
|
| 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)
|
|
|
|
|
| row_max = affinity.max(dim=1).values
|
| row_mean = affinity.mean(dim=1).clamp(min=1e-9)
|
| specialisation = row_max / row_mean
|
|
|
|
|
| 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)
|
|
|
|
|
|
|
| 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
|
|
|
|
|
|
|
| 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)
|
|
|
|
|
| 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)
|
|
|
|
|
| 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")
|
|
|