"""Blackfrost MiMo MoE-output intervention -- port to the EXL3 (exllamav3) runtime. Faithful port of the patched-sglang implementation (schemas blackfrost.mimo_mlp_output_intervention.v2 / .v3): h <- h - alpha * (h . d_hat) * d_hat (fp32 math, d_hat orthonormalized) Weights and routing are NOT modified: the effect is a pure runtime activation edit applied to the MoE / MLP output of the target layers, immediately before that output joins the residual stream. Directions are orthonormalized ONCE per direction set, in serialized manifest order (the reference build does this per layer module; the vector set is the same for every target layer, so it is computed once here). Enable with the same env var the sglang runtime uses: BLACKFROST_MIMO_MOE_INTERVENTION=/path/to/manifest.json Unset -> no intervention, baseline behaviour byte-for-byte. """ from __future__ import annotations import hashlib import json import math import os from pathlib import Path import torch ENV_VAR = "BLACKFROST_MIMO_MOE_INTERVENTION" DIRECTION_SCHEMA = "blackfrost.mimo_reasoning_boundary_direction.v1" ORIENTATION = "refusal_minus_control" SCHEMA_V1 = "blackfrost.mimo_moe_output_intervention.v1" SCHEMA_V2 = "blackfrost.mimo_mlp_output_intervention.v2" SCHEMA_V3 = "blackfrost.mimo_mlp_output_intervention.v3" class InterventionError(RuntimeError): """Raised when a manifest or direction file fails validation.""" def sha256_file(path: Path) -> str: digest = hashlib.sha256() with open(path, "rb") as handle: for block in iter(lambda: handle.read(1024 * 1024), b""): digest.update(block) return digest.hexdigest() class MoeOutputIntervention: """Loads one manifest and applies its directions to MoE outputs.""" def __init__(self, manifest_path, hidden_size: int, enabled: bool = True, logger=None): self.manifest_path = Path(manifest_path).resolve() self.hidden_size = int(hidden_size) self.enabled = bool(enabled) self.logger = logger self.schema = None self.experiment = None self.target_layers: tuple[int, ...] = () self._base_directions: list[torch.Tensor] = [] self._alpha: list[float] = [] self._device: list[torch.Tensor | None] = [] self._meta: list[dict] = [] self._v1_layers: dict[int, dict] = {} self._load() # ---------------------------------------------------------------- loading def _load(self) -> None: payload = json.loads(self.manifest_path.read_text(encoding="utf-8")) schema = payload.get("schema") self.schema = schema self.experiment = payload.get("experiment") if payload.get("weight_mutation"): raise InterventionError("manifest declares weight mutation; refusing") if payload.get("router_mutation"): raise InterventionError("manifest declares router mutation; refusing") if schema == SCHEMA_V1: if payload.get("surface") != "moe_outputs": raise InterventionError("v1 intervention must target moe_outputs") if payload.get("orientation") != ORIENTATION: raise InterventionError("intervention direction orientation mismatch") layers = payload.get("layers") if not isinstance(layers, dict) or not layers: raise InterventionError("v1 intervention has no selected layers") self._v1_layers = {int(k): v for k, v in layers.items()} self.target_layers = tuple(sorted(self._v1_layers)) return if schema not in (SCHEMA_V2, SCHEMA_V3): raise InterventionError(f"unsupported intervention schema: {schema!r}") if payload.get("application_surface") != "mlp_outputs_after_down_proj": raise InterventionError("intervention must target mlp_outputs_after_down_proj") if payload.get("orientation") != ORIENTATION: raise InterventionError("intervention direction orientation mismatch") target_layers = payload.get("target_layers") if not isinstance(target_layers, list) or not target_layers: raise InterventionError("intervention target_layers are invalid") self.target_layers = tuple(int(x) for x in target_layers) if schema == SCHEMA_V3: directions = payload.get("directions") if not isinstance(directions, list) or len(directions) < 2: raise InterventionError("v3 intervention requires at least two directions") entries = [(item, item.get("alpha")) for item in directions] else: alpha = payload.get("alpha") if not isinstance(alpha, (int, float)) or not math.isfinite(alpha) or alpha <= 0: raise InterventionError("v2 intervention alpha is invalid") entries = [(payload["direction"], float(alpha))] self._base_directions = self._orthonormalize(entries) def _orthonormalize(self, entries) -> list[torch.Tensor]: """Normalize, then Gram-Schmidt against prior directions, in serialized order.""" base: list[torch.Tensor] = [] for direction_config, alpha in entries: if not isinstance(alpha, (int, float)) or not math.isfinite(alpha) or alpha <= 0: raise InterventionError("direction alpha is invalid") rel = direction_config.get("direction", direction_config.get("path")) direction_path = (self.manifest_path.parent / rel).resolve() expected = direction_config.get("sha256") if expected is None: raise InterventionError(f"{direction_path.name}: no sha256 in manifest") if sha256_file(direction_path) != expected: raise InterventionError(f"{direction_path.name}: direction hash mismatch") direction_payload = torch.load(direction_path, map_location="cpu", weights_only=False) source_layer = direction_config.get("source_layer") source_tensor = direction_config.get("tensor") bank_sha256 = direction_config.get("bank_sha256") if ( direction_payload.get("schema") != DIRECTION_SCHEMA or (source_layer is not None and int(direction_payload.get("layer", -1)) != int(source_layer)) or (source_tensor is not None and direction_payload.get("tensor") != source_tensor) or direction_payload.get("orientation") != ORIENTATION or (bank_sha256 is not None and direction_payload.get("bank_sha256") != bank_sha256) ): raise InterventionError(f"{direction_path.name}: direction metadata mismatch") direction = direction_payload["direction"].float().contiguous() if direction.numel() != self.hidden_size: raise InterventionError(f"{direction_path.name}: direction width mismatch") direction = direction / direction.norm().clamp_min(1.0e-12) max_overlap = 0.0 for prior in base: signed = torch.dot(direction, prior) max_overlap = max(max_overlap, float(signed.abs())) direction.add_(prior, alpha=-float(signed)) retained_norm = float(direction.norm()) if retained_norm <= 1.0e-6: raise InterventionError(f"{direction_path.name}: linearly dependent direction") direction = direction / direction.norm().clamp_min(1.0e-12) for prior in base: if float(torch.dot(direction, prior).abs()) > 1.0e-5: raise InterventionError(f"{direction_path.name}: runtime orthogonalization failed") base.append(direction) self._meta.append( { "file": direction_path.name, "source_layer": source_layer, "tensor": source_tensor, "alpha": float(alpha), "serialized_overlap": max_overlap, "retained_norm": retained_norm, } ) self._alpha.append(float(alpha)) if self.logger is not None: for meta in self._meta: self.logger.info( "BLACKFROST_MIMO_MOE_INTERVENTION schema=%s experiment=%s layers=%s direction=%s " "source_layer=%s alpha=%.8f serialized_overlap=%.8f retained_norm=%.8f", self.schema, self.experiment, list(self.target_layers), meta["file"], meta["source_layer"], meta["alpha"], meta["serialized_overlap"], meta["retained_norm"], ) return base # ------------------------------------------------------------- properties @property def directions(self) -> list[torch.Tensor]: """Orthonormalized fp32 unit directions, in serialized order.""" return self._base_directions @property def alphas(self) -> list[float]: return self._alpha def handles_layer(self, layer_idx: int) -> bool: return self.enabled and int(layer_idx) in self.target_layers def describe(self) -> dict: return { "schema": self.schema, "experiment": self.experiment, "target_layers": list(self.target_layers), "directions": self._meta, "alphas": self._alpha, "enabled": self.enabled, "manifest": str(self.manifest_path), } # -------------------------------------------------------------- inference def apply(self, hidden_states: torch.Tensor) -> torch.Tensor: """h <- h - alpha * (h . d_hat) * d_hat, fp32 math, input dtype restored.""" if not self.enabled or not self._base_directions or hidden_states.numel() == 0: return hidden_states values = hidden_states.float() if values.data_ptr() == hidden_states.data_ptr(): # fp32 input: .float() is a no-op view, so clone before the in-place # update. Numerics are unchanged; this only avoids mutating the # caller's buffer (the reference mutates in place). values = values.clone() for index, (direction_cpu, alpha) in enumerate(zip(self._base_directions, self._alpha)): direction = self._device[index] if index < len(self._device) else None if direction is None or direction.device != hidden_states.device: direction = direction_cpu.to(device=hidden_states.device, non_blocking=False) while len(self._device) <= index: self._device.append(None) self._device[index] = direction projection = torch.matmul(values, direction) values.add_(projection.unsqueeze(-1) * direction, alpha=-alpha) return values.to(dtype=hidden_states.dtype) def load_from_env(hidden_size: int, logger=None) -> MoeOutputIntervention | None: """Build the intervention from BLACKFROST_MIMO_MOE_INTERVENTION, or None if unset.""" manifest = os.environ.get(ENV_VAR) if not manifest: return None if os.environ.get("BLACKFROST_MIMO_MOE_INTERVENTION_ENABLED", "1").strip() in ("0", "false", "False"): return None return MoeOutputIntervention(manifest, hidden_size, enabled=True, logger=logger) # -------------------------------------------------------------------------- # Fork integration # -------------------------------------------------------------------------- _CACHE: dict = {} def get_cached(hidden_size: int, logger=None): """Load once per process (per hidden size). Returns None when inactive.""" if hidden_size not in _CACHE: _CACHE[hidden_size] = load_from_env(hidden_size, logger=logger) return _CACHE[hidden_size] def apply_layer(layer_idx, hidden_states, params=None): """Apply the intervention to one layer's MLP/MoE output. No-op (returns the input unchanged) when the env var is unset, when the layer is not a target layer, or when the caller passes params["blackfrost_derisk"] = False for a per-request disable. """ if hidden_states is None or hidden_states.numel() == 0: return hidden_states if params is not None and params.get("blackfrost_derisk") is False: return hidden_states size = hidden_states.shape[-1] if size not in _CACHE: _CACHE[size] = load_from_env(size) loaded = _CACHE[size] if loaded is not None and not getattr(loaded, "_announced", False): loaded._announced = True print( "BLACKFROST_DERISK active: layers=" + str(list(loaded.target_layers)) + " alphas=" + str(loaded.alphas) + " manifest=" + str(loaded.manifest_path), flush=True, ) intervention = _CACHE[size] if intervention is None or not intervention.handles_layer(layer_idx): return hidden_states return intervention.apply(hidden_states)