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")