From a0a5d08d3d5194c5326a29e8a443c5bf592dfc2e Mon Sep 17 00:00:00 2001 From: Kai Xu Date: Thu, 10 Sep 2026 18:12:04 -0700 Subject: [PATCH] Allow DASC quality gates above parity Signed-off-by: Kai Xu --- modelopt/torch/sparsity/state_sparsity/config.py | 6 +++--- .../torch/sparsity/state_sparsity/test_dasc.py | 14 +++++++++++--- 2 files changed, 14 insertions(+), 6 deletions(-) diff --git a/modelopt/torch/sparsity/state_sparsity/config.py b/modelopt/torch/sparsity/state_sparsity/config.py index 8ae62cddff6..26c52e9dcf7 100644 --- a/modelopt/torch/sparsity/state_sparsity/config.py +++ b/modelopt/torch/sparsity/state_sparsity/config.py @@ -142,9 +142,9 @@ def validate_wmax_candidates(cls, candidates: object) -> object: @field_validator("min_perplexity_retention") @classmethod def validate_perplexity_gate(cls, value: float) -> float: - """Require a finite retention gate in (0, 1].""" - if not math.isfinite(value) or not 0.0 < value <= 1.0: - raise ValueError("min_perplexity_retention must be finite and in (0, 1]") + """Require a finite positive retention gate.""" + if not math.isfinite(value) or value <= 0.0: + raise ValueError("min_perplexity_retention must be finite and positive") return value @field_validator("min_top1_agreement") diff --git a/tests/unit/torch/sparsity/state_sparsity/test_dasc.py b/tests/unit/torch/sparsity/state_sparsity/test_dasc.py index 0ae651863c0..570e65fc24e 100644 --- a/tests/unit/torch/sparsity/state_sparsity/test_dasc.py +++ b/tests/unit/torch/sparsity/state_sparsity/test_dasc.py @@ -160,12 +160,20 @@ def test_config_fails_closed(override): def test_perplexity_retention_accepts_parity_improvements(): - """Allow an observed DASC perplexity improvement instead of requiring clamping.""" - measurement = mtss.DASCCalibrationMeasurement(**_candidate(7)) - measurement.quality[0].perplexity_retention = 1.0004 + """Allow improvement measurements and thresholds instead of requiring clamping.""" + candidate = _candidate(7) + candidate["quality"][0]["perplexity_retention"] = 1.0004 + measurement = mtss.DASCCalibrationMeasurement(**candidate) assert measurement.quality[0].perplexity_retention == 1.0004 + model = mtss.calibrate( + TinyGatedDeltaNetForCausalLM(), + _config(wmax_candidates=[7], min_perplexity_retention=1.0002), + [candidate], + ) + assert mtss.export_policy(model)["selected_wmax"] == 7 + def test_calibration_fails_closed_on_measurements_and_model_mismatch(): """Reject incomplete evidence, failing gates, wrong geometry, and unsupported layers."""