mirror of https://github.com/razor-ai/soup.git
feat(shrink): pure verdict half (decide_shrink) (v0.71.29)
This commit is contained in:
parent
0141d6267e
commit
1ca3f7d01a
|
|
@ -0,0 +1,133 @@
|
|||
"""soup shrink — depth-prune + distill-heal (v0.71.29, arXiv:2403.17887).
|
||||
|
||||
"The Unreasonable Ineffectiveness of the Deeper Layers" (Gromov et al.): rank a
|
||||
model's decoder layers by the angular distance of the residual stream across a
|
||||
contiguous block over a calibration set, drop the least-important block, then
|
||||
optionally *heal* by distilling the original model into the pruned student.
|
||||
|
||||
This module has two halves:
|
||||
|
||||
* a **pure** verdict half (frozen dataclasses + ``decide_shrink`` +
|
||||
``render_shrink_panel`` + ``shrink_verdict_to_dict``) with NO top-level torch
|
||||
import, so it is fully CPU-testable and cheap to import; and
|
||||
* a **torch-lazy** prune / importance half (``compute_layer_importance``,
|
||||
``select_drop_block``, ``prune_model_layers``, arch allowlist) whose heavy
|
||||
imports happen inside the functions.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from dataclasses import asdict, dataclass
|
||||
|
||||
from rich.panel import Panel
|
||||
|
||||
from soup_cli import __version__
|
||||
|
||||
DECISION_SHIP = "SHIP"
|
||||
DECISION_DONT_SHIP = "DON'T SHIP"
|
||||
DEFAULT_TOLERANCE = 0.10
|
||||
MAX_TOLERANCE = 5.0
|
||||
|
||||
# Verdict ratio epsilon so an exact-boundary drop (ratio-1 == tolerance) SHIPs.
|
||||
_RATIO_EPS = 1e-9
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Frozen dataclasses (pure)
|
||||
# ---------------------------------------------------------------------------
|
||||
@dataclass(frozen=True)
|
||||
class LayerImportance:
|
||||
"""One candidate contiguous block, ranked by residual angular distance."""
|
||||
|
||||
start: int # first dropped decoder layer (0-indexed)
|
||||
block_size: int # number of layers in the block
|
||||
angular_distance: float # mean per-token angular distance (lower = safer to drop)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ShrinkVerdict:
|
||||
"""The binary shrink decision plus the evidence that produced it."""
|
||||
|
||||
decision: str # DECISION_SHIP | DECISION_DONT_SHIP
|
||||
ppl_original: float
|
||||
ppl_final: float
|
||||
ppl_ratio: float # ppl_final / ppl_original
|
||||
tolerance: float
|
||||
layers_before: int
|
||||
layers_after: int
|
||||
params_saved_pct: float
|
||||
healed: bool
|
||||
soup_version: str
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Verdict (pure)
|
||||
# ---------------------------------------------------------------------------
|
||||
def _finite_positive(value: object, name: str) -> float:
|
||||
"""Coerce ``value`` to a finite, strictly-positive float."""
|
||||
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
||||
raise ValueError(f"{name} must be a number, got {type(value).__name__}")
|
||||
out = float(value)
|
||||
if not math.isfinite(out) or out <= 0.0:
|
||||
raise ValueError(f"{name} must be a finite positive number")
|
||||
return out
|
||||
|
||||
|
||||
def decide_shrink(
|
||||
ppl_original: object,
|
||||
ppl_final: object,
|
||||
*,
|
||||
tolerance: float = DEFAULT_TOLERANCE,
|
||||
layers_before: int,
|
||||
layers_after: int,
|
||||
params_saved_pct: float = 0.0,
|
||||
healed: bool = False,
|
||||
soup_version: str = __version__,
|
||||
) -> ShrinkVerdict:
|
||||
"""SHIP iff ``ppl_final / ppl_original - 1 <= tolerance``.
|
||||
|
||||
``decide_ship`` (soup ship) would trivially reject every shrink because
|
||||
pruning always raises perplexity — so shrink has its own rule: the pruned
|
||||
(and optionally healed) model ships when its perplexity regression stays
|
||||
within ``tolerance`` (absolute ratio, default 10 %).
|
||||
"""
|
||||
orig = _finite_positive(ppl_original, "ppl_original")
|
||||
final = _finite_positive(ppl_final, "ppl_final")
|
||||
if isinstance(tolerance, bool) or not isinstance(tolerance, (int, float)):
|
||||
raise ValueError("tolerance must be a number")
|
||||
tol = float(tolerance)
|
||||
if not math.isfinite(tol) or not (0.0 <= tol <= MAX_TOLERANCE):
|
||||
raise ValueError(f"tolerance must be in [0.0, {MAX_TOLERANCE}]")
|
||||
ratio = final / orig
|
||||
decision = DECISION_SHIP if (ratio - 1.0) <= tol + _RATIO_EPS else DECISION_DONT_SHIP
|
||||
return ShrinkVerdict(
|
||||
decision=decision,
|
||||
ppl_original=round(orig, 4),
|
||||
ppl_final=round(final, 4),
|
||||
ppl_ratio=round(ratio, 4),
|
||||
tolerance=tol,
|
||||
layers_before=int(layers_before),
|
||||
layers_after=int(layers_after),
|
||||
params_saved_pct=round(float(params_saved_pct), 2),
|
||||
healed=bool(healed),
|
||||
soup_version=str(soup_version),
|
||||
)
|
||||
|
||||
|
||||
def shrink_verdict_to_dict(verdict: ShrinkVerdict) -> dict:
|
||||
"""Plain-dict view of a ``ShrinkVerdict`` (JSON-serialisable)."""
|
||||
return asdict(verdict)
|
||||
|
||||
|
||||
def render_shrink_panel(verdict: ShrinkVerdict) -> Panel:
|
||||
"""One-screen Rich panel summarising the shrink verdict."""
|
||||
color = "green" if verdict.decision == DECISION_SHIP else "red"
|
||||
body = (
|
||||
f"[bold]{verdict.decision}[/]\n\n"
|
||||
f"Layers: {verdict.layers_before} -> {verdict.layers_after} "
|
||||
f"(params saved {verdict.params_saved_pct:.1f}%)\n"
|
||||
f"Perplexity: {verdict.ppl_original:.3f} -> {verdict.ppl_final:.3f} "
|
||||
f"(x{verdict.ppl_ratio:.3f}, tolerance {verdict.tolerance:.0%})\n"
|
||||
f"Healed: {'yes' if verdict.healed else 'no'}"
|
||||
)
|
||||
return Panel(body, title="soup shrink", border_style=color)
|
||||
|
|
@ -0,0 +1,118 @@
|
|||
"""v0.71.29 — `soup shrink`: depth-prune + distill-heal (arXiv:2403.17887).
|
||||
|
||||
Tests the pure verdict half, the torch-lazy prune/importance half, the CLI
|
||||
orchestration, the subprocess distill-heal wiring, and registry attach.
|
||||
"""
|
||||
import ast
|
||||
import math
|
||||
import pathlib
|
||||
from io import StringIO
|
||||
|
||||
import pytest
|
||||
from rich.console import Console
|
||||
|
||||
from soup_cli.utils.shrink import (
|
||||
DECISION_DONT_SHIP,
|
||||
DECISION_SHIP,
|
||||
LayerImportance,
|
||||
decide_shrink,
|
||||
render_shrink_panel,
|
||||
shrink_verdict_to_dict,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Task 1 — pure verdict half
|
||||
# ---------------------------------------------------------------------------
|
||||
class TestDecideShrink:
|
||||
def test_within_tolerance_ships(self):
|
||||
v = decide_shrink(10.0, 10.5, tolerance=0.10, layers_before=30, layers_after=24)
|
||||
assert v.decision == DECISION_SHIP
|
||||
assert math.isclose(v.ppl_ratio, 1.05)
|
||||
|
||||
def test_exceeds_tolerance_dont_ship(self):
|
||||
v = decide_shrink(10.0, 12.0, tolerance=0.10, layers_before=30, layers_after=24)
|
||||
assert v.decision == DECISION_DONT_SHIP
|
||||
|
||||
def test_boundary_exactly_at_tolerance_ships(self):
|
||||
v = decide_shrink(10.0, 11.0, tolerance=0.10, layers_before=30, layers_after=24)
|
||||
assert v.decision == DECISION_SHIP # ratio-1 == tolerance -> <=, SHIP
|
||||
|
||||
def test_improved_ppl_ships(self):
|
||||
v = decide_shrink(10.0, 9.5, tolerance=0.10, layers_before=30, layers_after=24)
|
||||
assert v.decision == DECISION_SHIP
|
||||
|
||||
def test_rejects_nonpositive_ppl(self):
|
||||
with pytest.raises(ValueError):
|
||||
decide_shrink(0.0, 5.0, layers_before=30, layers_after=24)
|
||||
with pytest.raises(ValueError):
|
||||
decide_shrink(5.0, -1.0, layers_before=30, layers_after=24)
|
||||
|
||||
def test_rejects_nonfinite(self):
|
||||
with pytest.raises(ValueError):
|
||||
decide_shrink(10.0, float("inf"), layers_before=30, layers_after=24)
|
||||
with pytest.raises(ValueError):
|
||||
decide_shrink(float("nan"), 5.0, layers_before=30, layers_after=24)
|
||||
|
||||
def test_rejects_bool_ppl(self):
|
||||
with pytest.raises(ValueError):
|
||||
decide_shrink(True, 5.0, layers_before=30, layers_after=24)
|
||||
|
||||
def test_rejects_bad_tolerance(self):
|
||||
with pytest.raises(ValueError):
|
||||
decide_shrink(10.0, 10.0, tolerance=-0.1, layers_before=30, layers_after=24)
|
||||
with pytest.raises(ValueError):
|
||||
decide_shrink(10.0, 10.0, tolerance=6.0, layers_before=30, layers_after=24)
|
||||
with pytest.raises(ValueError):
|
||||
decide_shrink(10.0, 10.0, tolerance=True, layers_before=30, layers_after=24)
|
||||
|
||||
def test_frozen(self):
|
||||
v = decide_shrink(10.0, 10.5, layers_before=30, layers_after=24)
|
||||
with pytest.raises(Exception):
|
||||
v.decision = "x" # type: ignore[misc]
|
||||
|
||||
def test_to_dict_roundtrip(self):
|
||||
v = decide_shrink(
|
||||
10.0, 10.5, layers_before=30, layers_after=24, params_saved_pct=20.0, healed=True
|
||||
)
|
||||
d = shrink_verdict_to_dict(v)
|
||||
assert d["decision"] == v.decision
|
||||
assert d["healed"] is True
|
||||
assert set(d) >= {
|
||||
"decision", "ppl_original", "ppl_final", "ppl_ratio", "tolerance",
|
||||
"layers_before", "layers_after", "params_saved_pct", "healed", "soup_version",
|
||||
}
|
||||
|
||||
def test_render_panel_names_decision(self):
|
||||
v = decide_shrink(10.0, 12.0, layers_before=30, layers_after=24)
|
||||
buf = StringIO()
|
||||
Console(file=buf, width=100).print(render_shrink_panel(v))
|
||||
assert "DON'T SHIP" in buf.getvalue()
|
||||
|
||||
def test_render_panel_ship(self):
|
||||
v = decide_shrink(10.0, 10.2, layers_before=30, layers_after=24)
|
||||
buf = StringIO()
|
||||
Console(file=buf, width=100).print(render_shrink_panel(v))
|
||||
out = buf.getvalue()
|
||||
assert "SHIP" in out and "DON'T SHIP" not in out
|
||||
|
||||
def test_layer_importance_frozen(self):
|
||||
li = LayerImportance(start=5, block_size=8, angular_distance=0.12)
|
||||
assert (li.start, li.block_size) == (5, 8)
|
||||
with pytest.raises(Exception):
|
||||
li.start = 1 # type: ignore[misc]
|
||||
|
||||
|
||||
class TestNoTopLevelTorch:
|
||||
def test_shrink_module_has_no_top_level_heavy_import(self):
|
||||
src = pathlib.Path("src/soup_cli/utils/shrink.py").read_text(encoding="utf-8")
|
||||
tree = ast.parse(src)
|
||||
names: list[str] = []
|
||||
for node in tree.body:
|
||||
if isinstance(node, ast.Import):
|
||||
names += [a.name for a in node.names]
|
||||
elif isinstance(node, ast.ImportFrom):
|
||||
names.append(node.module or "")
|
||||
assert not any(
|
||||
m.split(".")[0] in {"torch", "transformers", "peft"} for m in names
|
||||
), names
|
||||
Loading…
Reference in New Issue