Download exllamav3/blackfrost_derisk.py from vcruz305/MiMo-V2.6-Flash-MOPD-DERISKED-EXL3-2.20bpw: direct link, hf CLI and curl.
- Browser
- Download file 13 kB
-
https://huggingface.co/vcruz305/MiMo-V2.6-Flash-MOPD-DERISKED-EXL3-2.20bpw/resolve/main/exllamav3/blackfrost_derisk.py
- Command line
-
hf download hf://vcruz305/MiMo-V2.6-Flash-MOPD-DERISKED-EXL3-2.20bpw/exllamav3/blackfrost_derisk.py
-
curl -L -o blackfrost_derisk.py https://huggingface.co/vcruz305/MiMo-V2.6-Flash-MOPD-DERISKED-EXL3-2.20bpw/resolve/main/exllamav3/blackfrost_derisk.py
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 | |
| def directions(self) -> list[torch.Tensor]: | |
| """Orthonormalized fp32 unit directions, in serialized order.""" | |
| return self._base_directions | |
| 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) | |