Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
42 changes: 36 additions & 6 deletions modelopt/torch/sparsity/state_sparsity/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
"""Configuration and result schemas for DASC state sparsity."""

import math
from numbers import Real
from typing import Literal

from pydantic import ConfigDict, Field, field_validator, model_validator
Expand All @@ -31,6 +32,37 @@
]

_DecayParameterStorageDtype = Literal["float16", "bfloat16", "float32"]
_DEFAULT_EPSILON = 1e-3
_DEFAULT_STATIC_GATE_INPUT = -0.3


def _validate_analysis_arguments(
epsilon: object = _DEFAULT_EPSILON,
static_gate_input: object = _DEFAULT_STATIC_GATE_INPUT,
) -> None:
"""Reject decay-analysis arguments that cannot produce well-defined horizons."""
try:
epsilon_is_valid = (
isinstance(epsilon, Real)
and not isinstance(epsilon, bool)
and math.isfinite(epsilon)
and 0.0 < epsilon < 1.0
)
except OverflowError:
epsilon_is_valid = False
if not epsilon_is_valid:
raise ValueError("epsilon must be finite and in (0, 1)")

try:
static_gate_input_is_valid = (
isinstance(static_gate_input, Real)
and not isinstance(static_gate_input, bool)
and math.isfinite(static_gate_input)
)
except OverflowError:
static_gate_input_is_valid = False
if not static_gate_input_is_valid:
raise ValueError("static_gate_input must be finite")
Comment on lines +44 to +65

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

[IMPORTANT Compatibility] The normalized exception set misses RuntimeError, which is exactly what a non-scalar tensor raises — so the public ValueError contract this PR is consolidating now leaks for a whole input class that the previous version rejected cleanly.

What changed: the version removed from policy.py gated on isinstance(epsilon, float | int) first, so any non-float/int argument (including every torch.Tensor) took the ValueError path. This one drops the isinstance gate and relies on math.isfinite raising a normalizable exception instead. That holds for lists/strings/complex (TypeError), numpy arrays (TypeError), and oversized ints (OverflowError) — but math.isfinite reaches a tensor's __float__, and torch raises RuntimeError: a Tensor with N elements cannot be converted to Scalar for any tensor whose numel() != 1 (including empty tensors). RuntimeError is not in the except tuple, so it propagates unchanged out of mtss.analyze_gdn_decay and mtss.compute_gdn_decay_horizons.

Why it matters: the new test_analysis_accepts_tensor_scalar_arguments deliberately makes 0-d tensors a supported input, which makes epsilon=torch.tensor([1e-3, 2e-3]) (or a stray A_log-shaped tensor) a very plausible caller mistake rather than an exotic one. Callers that follow the documented contract and wrap these entry points in except ValueError will crash instead of surfacing the intended message — a regression relative to the base branch, on the exact axis this PR is about. 0.0 < epsilon < 1.0 and the if not epsilon_is_valid bool conversion have the same exposure for multi-element tensors.

Fix: add RuntimeError to both handlers (a targeted follow-up test with a 2-element tensor would lock it in):

Suggested change
try:
epsilon_is_valid = math.isfinite(epsilon) and 0.0 < epsilon < 1.0
except (TypeError, ValueError, OverflowError):
epsilon_is_valid = False
if not epsilon_is_valid:
raise ValueError("epsilon must be finite and in (0, 1)")
try:
static_gate_input_is_valid = math.isfinite(static_gate_input)
except (TypeError, ValueError, OverflowError):
static_gate_input_is_valid = False
if not static_gate_input_is_valid:
raise ValueError("static_gate_input must be finite")
try:
epsilon_is_valid = math.isfinite(epsilon) and 0.0 < epsilon < 1.0
except (TypeError, ValueError, OverflowError, RuntimeError):
epsilon_is_valid = False
if not epsilon_is_valid:
raise ValueError("epsilon must be finite and in (0, 1)")
try:
static_gate_input_is_valid = math.isfinite(static_gate_input)
except (TypeError, ValueError, OverflowError, RuntimeError):
static_gate_input_is_valid = False
if not static_gate_input_is_valid:
raise ValueError("static_gate_input must be finite")

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed in #2397 with a stricter scalar boundary rather than catching arbitrary framework RuntimeError. The annotated inputs now require numbers.Real; NumPy real scalars remain accepted, while scalar and multi-element tensors are rejected uniformly with ValueError before math.isfinite. Multi-element regressions cover both public entry points.



class DASCQualityMeasurement(ModeloptBaseConfig):
Expand Down Expand Up @@ -81,11 +113,11 @@ class DASCConfig(ModeloptBaseConfig):
description="Use zero recovery (DASC-NR) or suffix replay recovery (DASC-WR).",
)
epsilon: float = ModeloptField(
default=1e-3,
default=_DEFAULT_EPSILON,
description="Retained contribution threshold used to derive static decay horizons.",
)
static_gate_input: float = ModeloptField(
default=-0.3,
default=_DEFAULT_STATIC_GATE_INPUT,
description="Static gate input added to each GDN head's dt_bias.",
)
decay_parameter_storage_dtype: _DecayParameterStorageDtype = ModeloptField(
Expand Down Expand Up @@ -118,16 +150,14 @@ class DASCConfig(ModeloptBaseConfig):
@classmethod
def validate_epsilon(cls, epsilon: float) -> float:
"""Require a finite decay threshold strictly between zero and one."""
if not math.isfinite(epsilon) or not 0.0 < epsilon < 1.0:
raise ValueError("epsilon must be finite and in (0, 1)")
_validate_analysis_arguments(epsilon=epsilon)
return epsilon

@field_validator("static_gate_input")
@classmethod
def validate_static_gate_input(cls, value: float) -> float:
"""Require a finite representative gate input."""
if not math.isfinite(value):
raise ValueError("static_gate_input must be finite")
_validate_analysis_arguments(static_gate_input=value)
return value

@field_validator("wmax_candidates", mode="before")
Expand Down
13 changes: 3 additions & 10 deletions modelopt/torch/sparsity/state_sparsity/policy.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,7 @@
DASCLayerPolicy,
DASCPolicy,
_DecayParameterStorageDtype,
_validate_analysis_arguments,
)

__all__ = ["analyze_gdn_decay", "compute_gdn_decay_horizons"]
Expand Down Expand Up @@ -65,14 +66,6 @@ class _DASCDecayParametersUnavailableError(_DASCRecoverableStalenessError):
"""Identify supported GDN modules whose decay tensors are temporarily unavailable."""


def _validate_analysis_arguments(epsilon: float, static_gate_input: float) -> None:
"""Normalize invalid public analysis arguments to the ValueError contract."""
if not isinstance(epsilon, float | int) or not (math.isfinite(epsilon) and 0.0 < epsilon < 1.0):
raise ValueError("epsilon must be finite and in (0, 1)")
if not isinstance(static_gate_input, float | int) or not math.isfinite(static_gate_input):
raise ValueError("static_gate_input must be finite")


def _validate_gdn_decay_tensors(a_log: torch.Tensor, dt_bias: torch.Tensor) -> None:
"""Reject decay tensors that cannot produce well-defined horizons."""
if a_log.ndim != 1 or dt_bias.ndim != 1 or a_log.shape != dt_bias.shape or not a_log.numel():
Expand Down Expand Up @@ -129,7 +122,7 @@ def compute_gdn_decay_horizons(
dt_bias_cpu = dt_bias.detach().to(device="cpu", dtype=torch.float64)

decay = -torch.exp(a_log_cpu) * F.softplus(dt_bias_cpu + static_gate_input)
horizons = torch.log(torch.tensor(epsilon, dtype=torch.float64)) / decay
horizons = math.log(epsilon) / decay
if not torch.isfinite(horizons).all() or not torch.all(horizons > 0):
raise ValueError("GDN decay parameters produced non-finite or non-positive horizons")
return horizons
Expand Down Expand Up @@ -527,7 +520,7 @@ def validate_dasc_decay_parameters(model: nn.Module, policy: DASCPolicy) -> None
"DASC policy head mask does not match current decay parameters in layer "
f"{name!r}"
)
stored = torch.tensor(layer.static_horizons, dtype=torch.float64)
stored = torch.tensor(layer.static_horizons, device="cpu", dtype=torch.float64)
numerical_slack = 32.0 * torch.finfo(torch.float64).eps
if torch.any(stored < lower * (1.0 - numerical_slack)) or torch.any(
stored > upper * (1.0 + numerical_slack)
Expand Down
34 changes: 32 additions & 2 deletions tests/unit/torch/sparsity/state_sparsity/test_dasc.py
Original file line number Diff line number Diff line change
Expand Up @@ -454,14 +454,16 @@ def test_analysis_arguments_fail_at_the_public_boundary(invalid_storage_dtype):
)


@pytest.mark.parametrize("epsilon", [1.0, []])
@pytest.mark.parametrize("epsilon", [1.0, True, [], 10**1000, torch.tensor([1e-3, 2e-3])])
def test_analysis_rejects_invalid_epsilon_at_the_public_boundary(epsilon):
"""Normalize invalid epsilon values to the public ValueError contract."""
with pytest.raises(ValueError, match=r"epsilon must be finite and in \(0, 1\)"):
mtss.analyze_gdn_decay(TinyGatedDeltaNetForCausalLM(), epsilon=epsilon) # type: ignore[arg-type]


@pytest.mark.parametrize("static_gate_input", [torch.nan, []])
@pytest.mark.parametrize(
"static_gate_input", [torch.nan, True, [], 10**1000, torch.tensor([-0.3, -0.2])]
)
def test_analysis_rejects_invalid_static_gate_input_at_the_public_boundary(static_gate_input):
"""Normalize invalid static gate values to the public ValueError contract."""
with pytest.raises(ValueError, match="static_gate_input must be finite"):
Expand All @@ -475,7 +477,15 @@ def test_analysis_rejects_invalid_static_gate_input_at_the_public_boundary(stati
("argument", "value", "message"),
[
("epsilon", [], r"epsilon must be finite and in \(0, 1\)"),
("epsilon", True, r"epsilon must be finite and in \(0, 1\)"),
("epsilon", torch.nan, r"epsilon must be finite and in \(0, 1\)"),
("epsilon", 10**1000, r"epsilon must be finite and in \(0, 1\)"),
("epsilon", torch.tensor([1e-3, 2e-3]), r"epsilon must be finite and in \(0, 1\)"),
("static_gate_input", [], "static_gate_input must be finite"),
("static_gate_input", True, "static_gate_input must be finite"),
("static_gate_input", torch.nan, "static_gate_input must be finite"),
("static_gate_input", 10**1000, "static_gate_input must be finite"),
("static_gate_input", torch.tensor([-0.3, -0.2]), "static_gate_input must be finite"),
],
)
def test_horizon_computation_rejects_invalid_public_arguments(argument, value, message):
Expand All @@ -489,6 +499,26 @@ def test_horizon_computation_rejects_invalid_public_arguments(argument, value, m
)


def test_horizon_computation_ignores_the_default_device():
"""Keep CPU horizon analysis independent of PyTorch's ambient allocation device."""
a_log = torch.tensor([0.0])
dt_bias = torch.tensor([0.0])
with torch.device("meta"):
horizons = mtss.compute_gdn_decay_horizons(a_log, dt_bias)
assert horizons.device.type == "cpu"


def test_policy_lifecycle_ignores_the_default_device():
"""Keep calibration and checkpoint metadata validation on their declared CPU path."""
model = TinyGatedDeltaNetForCausalLM()
with torch.device("meta"):
calibrated = mtss.calibrate(model, _config(wmax_candidates=[7]), [_candidate(7)])
state = mto.modelopt_state(calibrated)
policy = mtss.export_policy(calibrated)
assert state["modelopt_state_dict"][0][0] == "dasc"
assert policy["layers"]["linear_attn"]["static_horizons"]


def test_bf16_storage_round_trip_loaded_in_fp32_preserves_policy():
"""Accept BF16-rounded values after a checkpoint loader materializes FP32 tensors."""
model = mtss.calibrate(
Expand Down
Loading