mirror of
https://github.com/wassname/evil_MoE.git
synced 2026-09-12 07:50:38 +08:00
add routeV_absorb_all: 100% absorption, no vector (H2 extreme control)
Route the whole gradient of every knob-on rollout into the quarantine; the deployed knob learns only from the knob-off exploration floor. Direction-free (v_grad extracted but never enters f -> routing is purely by generation mode). Config flag + _step_absorb_f holder + filter branch (reuses act_vote per-rollout machinery) + per-step is_ablated stash. just smoke-absorb passes (keep=0.25/ rout=0.75 = the floor/knob-on split). Queued s43 as job 29 (frac=0.25). Co-Authored-By: Claudypoo <288921227+claudypoo@users.noreply.github.com>
This commit is contained in:
+27
-1
@@ -212,6 +212,13 @@ class Config:
|
||||
# band each step. No pairs needed for threshold calibration -- direction only.
|
||||
online_stats_lo: float = 0.05 # lower quantile -> keep tail
|
||||
online_stats_hi: float = 0.95 # upper quantile -> route tail
|
||||
# 100%-absorption control (NO vector). Route the WHOLE gradient of every knob-on
|
||||
# rollout into the quarantine (f=1), keep only the knob-off exploration-floor rollouts
|
||||
# (is_ablated, f=0) in the deployed knob. The extreme of H2: the quarantine as a pure
|
||||
# gradient sink, routing by generation-mode not by any direction. v_grad is still
|
||||
# extracted (reuses the routeV path) but never touches f -- routing is direction-free.
|
||||
# Requires rollout_ablate_frac>0, else the deployed knob never updates (= base model).
|
||||
routeV_absorb_all: bool = False
|
||||
# Per-source cin diagnostic: split each prompt's backward into student-only
|
||||
# + teacher-only passes (~2x backward time). 1 = every step (default; full
|
||||
# signal); N>1 = only every Nth step (combined backward elsewhere, ~halves
|
||||
@@ -991,6 +998,7 @@ def main(cfg: Config) -> int:
|
||||
# modules (the global activation vote, computed post-backward before the per-module
|
||||
# routing). 1-element list so the filter closure reads the current step's value.
|
||||
_step_f_roll: list[torch.Tensor | None] = [None]
|
||||
_step_absorb_f: list[torch.Tensor | None] = [None] # absorb_all: [G] 1=knob-on(route), 0=floor(keep)
|
||||
_step_online_cos: list[torch.Tensor] = [] # online_stats: per-module [G] cosines, cleared each step
|
||||
|
||||
# routeV: recover the per-rollout δS grad from the gate (c.grad = δS * g_b),
|
||||
@@ -1024,7 +1032,20 @@ def main(cfg: Config) -> int:
|
||||
# per-token (routeV_per_token): one cos/f per token -- finer but noisier.
|
||||
lower, upper = route_band[name]
|
||||
band = max(upper - lower, 1e-6)
|
||||
if cfg.routeV_gate == "act_vote":
|
||||
if cfg.routeV_absorb_all:
|
||||
# NO vector: f is purely the generation-mode mask (1=knob-on -> route the
|
||||
# whole rollout, 0=knob-off floor -> keep). Direction-free 100% absorption;
|
||||
# v_grad/band above are computed but never enter f.
|
||||
cg = cg_full.sum(1) # [G, r] per-rollout δS*g
|
||||
g_b = torch.where(reliable, cg / dS_safe, torch.zeros_like(cg)) # [G, r]
|
||||
f = _step_absorb_f[0] # [G] 1=route, 0=keep
|
||||
routed = torch.where(reliable, (cg * f.unsqueeze(1)).sum(0) / dS_safe,
|
||||
torch.zeros_like(g))
|
||||
step_flagged.append(f.mean().item())
|
||||
_kn, _rn, _on, _ke, _re, _oe = _zone_stats(f, g_b.norm(dim=1))
|
||||
step_zkeep.append(_kn); step_zresid.append(_rn); step_zrout.append(_on)
|
||||
step_zkeepE.append(_ke); step_zresidE.append(_re); step_zroutE.append(_oe)
|
||||
elif cfg.routeV_gate == "act_vote":
|
||||
# Global gate: route every module's per-rollout grad by the SAME f_roll
|
||||
# (the activation vote, computed once for the step). Per-rollout granularity
|
||||
# by construction; per_token is ignored under act_vote.
|
||||
@@ -1458,6 +1479,11 @@ def main(cfg: Config) -> int:
|
||||
# routing (activations are cached on every layer from the loss forward).
|
||||
if is_routeV and cfg.routeV_gate == "act_vote":
|
||||
_step_f_roll[0] = _act_vote_f_roll(merged.shape[0], plen, mask)
|
||||
# absorb_all: per-rollout route mask = generation mode (knob-on -> 1 route,
|
||||
# knob-off floor -> 0 keep). Same row order as merged (students then teachers).
|
||||
if is_routeV and cfg.routeV_absorb_all:
|
||||
_step_absorb_f[0] = torch.tensor(
|
||||
[0.0 if ab else 1.0 for ab in is_ablated], device=device)
|
||||
for name, info in wrappers.items():
|
||||
g = info["delta_S"].grad
|
||||
if g is None:
|
||||
|
||||
Reference in New Issue
Block a user