vcruz305's picture
EXL3 2.20 bpw pack + Blackfrost derisk overlay
343ce31 verified
Raw History Blame Contribute Delete
13 kB
"""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)