Skip to content

Commit 7776861

Browse files
committed
Remove target_sparsity_ratio mode
Signed-off-by: Kai Xu <kaix@nvidia.com>
1 parent 84ae221 commit 7776861

7 files changed

Lines changed: 49 additions & 174 deletions

File tree

examples/llm_sparsity/kv_cache_sparsity/hf_triattention.py

Lines changed: 8 additions & 32 deletions
Original file line numberDiff line numberDiff line change
@@ -27,9 +27,6 @@
2727
# Fixed-size budget (retain top-K tokens per head)
2828
python hf_triattention.py --model Qwen/Qwen3-0.6B --budget 2048
2929
30-
# Percentile-based eviction (evict 70% at each prune step)
31-
python hf_triattention.py --model Qwen/Qwen3-0.6B --target-sparsity-ratio 0.7
32-
3330
# With custom budget and calibration length
3431
python hf_triattention.py --model Qwen/Qwen3-0.6B --budget 1024 --calib-seq-len 4096
3532
@@ -119,20 +116,11 @@ def main(args):
119116
print(f" Calibration complete in {elapsed:.1f}s")
120117

121118
# Step 2: Apply KV cache sparsity mode
122-
if args.target_sparsity_ratio is not None:
123-
print(
124-
f"\nApplying TriAttention mode (target_sparsity_ratio={args.target_sparsity_ratio})..."
125-
)
126-
triattention_config = TriAttentionConfig(
127-
target_sparsity_ratio=args.target_sparsity_ratio,
128-
prune_interval=args.prune_interval,
129-
)
130-
else:
131-
print(f"\nApplying TriAttention mode (budget={args.budget})...")
132-
triattention_config = TriAttentionConfig(
133-
budget=args.budget,
134-
prune_interval=args.prune_interval,
135-
)
119+
print(f"\nApplying TriAttention mode (budget={args.budget})...")
120+
triattention_config = TriAttentionConfig(
121+
budget=args.budget,
122+
prune_interval=args.prune_interval,
123+
)
136124
model = mtskv.sparsify(model, triattention_config)
137125
print(" Mode applied (no-op on weights).")
138126

@@ -184,22 +172,12 @@ def main(args):
184172
default="Qwen/Qwen3-0.6B",
185173
help="HuggingFace model name or local path.",
186174
)
187-
policy = parser.add_mutually_exclusive_group()
188-
policy.add_argument(
175+
parser.add_argument(
189176
"--budget",
190177
type=int,
191-
default=None,
178+
default=2048,
192179
help="KV token budget (tokens to retain per head). "
193-
"Mutually exclusive with --target-sparsity-ratio. "
194-
"Defaults to 2048 if neither is set.",
195-
)
196-
policy.add_argument(
197-
"--target-sparsity-ratio",
198-
type=float,
199-
default=None,
200-
help="Fraction of tokens to evict at each prune step, in (0, 1). "
201-
"Example: 0.7 evicts 70%% of tokens (keeps top 30%% by score). "
202-
"Mutually exclusive with --budget.",
180+
"Compression triggers after --prune-interval additional tokens.",
203181
)
204182
parser.add_argument(
205183
"--prune-interval",
@@ -229,6 +207,4 @@ def main(args):
229207
)
230208

231209
args = parser.parse_args()
232-
if args.budget is None and args.target_sparsity_ratio is None:
233-
args.budget = 2048
234210
main(args)

examples/speculative_decoding/eagle_utils.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -25,7 +25,6 @@
2525
import transformers
2626
from datasets import load_dataset
2727
from packaging.version import Version
28-
from scripts.ar_validate import validate_ar
2928
from transformers import Trainer, TrainerCallback
3029

3130
import modelopt
@@ -41,6 +40,7 @@
4140
ShardedDataset,
4241
VisionLanguageDataCollator,
4342
)
43+
from scripts.ar_validate import validate_ar
4444

4545
try:
4646
import wandb

modelopt/torch/sparsity/kv_cache/config.py

Lines changed: 13 additions & 29 deletions
Original file line numberDiff line numberDiff line change
@@ -33,29 +33,19 @@ class TriAttentionConfig(ModeloptBaseConfig):
3333
pre-RoPE Q/K concentration. Calibration computes per-head frequency statistics;
3434
at runtime, the serving engine scores and evicts tokens periodically.
3535
36-
Exactly one of ``budget`` or ``target_sparsity_ratio`` must be set:
37-
38-
- ``budget``: absolute token count to retain per head (fixed-size cache).
39-
- ``target_sparsity_ratio``: fraction of tokens to evict at each pruning step.
40-
Cache size auto-scales with generation length. Value in (0, 1).
36+
``budget`` is the absolute token count to retain per head after each pruning
37+
round. Runtime compression follows slack/sawtooth semantics: grow until
38+
``budget + prune_interval`` tokens are present, then evict back to
39+
``budget``.
4140
"""
4241

43-
# Eviction policy (exactly one must be set)
42+
# Eviction policy
4443
budget: int | None = ModeloptField(
4544
default=None,
4645
title="KV token budget (absolute).",
4746
description=(
48-
"Number of KV tokens to retain per head after pruning. "
49-
"Mutually exclusive with target_sparsity_ratio."
50-
),
51-
)
52-
target_sparsity_ratio: float | None = ModeloptField(
53-
default=None,
54-
title="Target sparsity ratio (percentile-based).",
55-
description=(
56-
"Fraction of tokens to evict at each pruning step, in (0, 1). "
57-
"Example: 0.7 means evict 70% of tokens (keep top 30% by score). "
58-
"Mutually exclusive with budget."
47+
"Number of KV tokens to retain per head after pruning. Runtime "
48+
"compression triggers after prune_interval additional tokens."
5949
),
6050
)
6151

@@ -127,16 +117,10 @@ def validate_score_aggregation(cls, v: str) -> str:
127117
return v
128118

129119
@model_validator(mode="after")
130-
def validate_budget_or_sparsity(self) -> TriAttentionConfig:
131-
"""Exactly one of budget or target_sparsity_ratio must be set."""
132-
budget_set = self.budget is not None
133-
sparsity_set = self.target_sparsity_ratio is not None
134-
if not budget_set and not sparsity_set:
135-
raise ValueError("Must set exactly one of 'budget' or 'target_sparsity_ratio'")
136-
if budget_set and sparsity_set:
137-
raise ValueError("Cannot set both 'budget' and 'target_sparsity_ratio'; pick one")
138-
if sparsity_set and not (0.0 < self.target_sparsity_ratio < 1.0):
139-
raise ValueError(
140-
f"target_sparsity_ratio must be in (0, 1), got {self.target_sparsity_ratio}"
141-
)
120+
def validate_budget(self) -> TriAttentionConfig:
121+
"""Validate the fixed KV token budget."""
122+
if self.budget is None:
123+
raise ValueError("TriAttention requires 'budget' to be set")
124+
if self.budget <= 0:
125+
raise ValueError(f"budget must be positive, got {self.budget}")
142126
return self

modelopt/torch/sparsity/kv_cache/triattention/scoring.py

Lines changed: 8 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -55,10 +55,11 @@ def compute_frequency_statistics_from_means(
5555
5656
Args:
5757
q_mean_complex: Mean of Q in complex frequency domain, shape (freq_count,).
58-
q_abs_mean: Mean of |Q| in frequency domain, shape (freq_count,).
58+
q_abs_mean: Mean of ``|Q|`` in frequency domain, shape (freq_count,).
5959
k_unrot: Unrotated key vectors, shape (num_keys, head_dim).
6060
style: RoPE pairing style.
61-
disable_mlr: If True, use q_abs_mean directly instead of (q_abs_mean - |q_mean|).
61+
disable_mlr: If True, use ``q_abs_mean`` directly instead of
62+
``q_abs_mean - |q_mean|``.
6263
6364
Returns:
6465
amp: Amplitude, shape (num_keys, freq_count).
@@ -135,37 +136,27 @@ def select_keys_to_keep(
135136
scores: torch.Tensor,
136137
*,
137138
kv_budget: int | None = None,
138-
target_sparsity_ratio: float | None = None,
139139
) -> torch.Tensor:
140140
"""Select which keys to retain based on importance scores.
141141
142-
Exactly one of ``kv_budget`` or ``target_sparsity_ratio`` must be provided.
143-
144142
Args:
145143
scores: Importance scores, shape (num_keys,). Higher = more important.
146144
kv_budget: Absolute number of tokens to retain. Keeps top-K.
147145
If budget >= num_keys, keeps all.
148-
target_sparsity_ratio: Fraction of tokens to evict, in (0, 1).
149-
Keeps top (1 - ratio) fraction. Example: 0.7 → keep top 30%.
150146
151147
Returns:
152148
Boolean mask, shape (num_keys,). True = keep, False = evict.
153149
"""
154-
budget_set = kv_budget is not None
155-
sparsity_set = target_sparsity_ratio is not None
156-
if budget_set == sparsity_set:
157-
raise ValueError(
158-
"select_keys_to_keep requires exactly one of kv_budget or target_sparsity_ratio"
159-
)
150+
if kv_budget is None:
151+
raise ValueError("select_keys_to_keep requires kv_budget")
152+
if kv_budget <= 0:
153+
raise ValueError(f"kv_budget must be positive, got {kv_budget}")
160154

161155
num_keys = scores.shape[0]
162156
if num_keys == 0:
163157
return torch.zeros(0, dtype=torch.bool, device=scores.device)
164158

165-
if budget_set:
166-
k = min(kv_budget, num_keys)
167-
else:
168-
k = max(1, round(num_keys * (1.0 - target_sparsity_ratio)))
159+
k = min(kv_budget, num_keys)
169160

170161
if k >= num_keys:
171162
return torch.ones(num_keys, dtype=torch.bool, device=scores.device)

tests/unit/torch/sparsity/kv_cache/test_config.py

Lines changed: 10 additions & 41 deletions
Original file line numberDiff line numberDiff line change
@@ -25,44 +25,24 @@ def test_budget_only():
2525
"""Setting only budget is valid."""
2626
config = TriAttentionConfig(budget=2048)
2727
assert config.budget == 2048
28-
assert config.target_sparsity_ratio is None
2928

3029

31-
def test_target_sparsity_only():
32-
"""Setting only target_sparsity_ratio is valid."""
33-
config = TriAttentionConfig(target_sparsity_ratio=0.7)
34-
assert config.budget is None
35-
assert config.target_sparsity_ratio == 0.7
36-
37-
38-
def test_both_budget_and_sparsity_raises():
39-
"""Setting both budget and target_sparsity_ratio raises."""
40-
with pytest.raises(ValidationError, match="Cannot set both"):
30+
def test_target_sparsity_ratio_is_not_supported():
31+
"""Ratio-based eviction is not part of the TriAttention API."""
32+
with pytest.raises(ValidationError):
4133
TriAttentionConfig(budget=2048, target_sparsity_ratio=0.7)
4234

4335

44-
def test_neither_budget_nor_sparsity_raises():
45-
"""Setting neither budget nor target_sparsity_ratio raises."""
46-
with pytest.raises(ValidationError, match="Must set exactly one"):
36+
def test_missing_budget_raises():
37+
"""TriAttention requires an explicit budget."""
38+
with pytest.raises(ValidationError, match="requires 'budget'"):
4739
TriAttentionConfig()
4840

4941

50-
def test_target_sparsity_out_of_range_low():
51-
"""target_sparsity_ratio <= 0 raises."""
52-
with pytest.raises(ValidationError, match="must be in"):
53-
TriAttentionConfig(target_sparsity_ratio=0.0)
54-
55-
56-
def test_target_sparsity_out_of_range_high():
57-
"""target_sparsity_ratio >= 1 raises."""
58-
with pytest.raises(ValidationError, match="must be in"):
59-
TriAttentionConfig(target_sparsity_ratio=1.0)
60-
61-
62-
def test_target_sparsity_negative():
63-
"""Negative target_sparsity_ratio raises."""
64-
with pytest.raises(ValidationError):
65-
TriAttentionConfig(target_sparsity_ratio=-0.1)
42+
def test_non_positive_budget_raises():
43+
"""Budget must be positive."""
44+
with pytest.raises(ValidationError, match="budget must be positive"):
45+
TriAttentionConfig(budget=0)
6646

6747

6848
def test_config_custom_values():
@@ -91,17 +71,6 @@ def test_config_serialization_roundtrip_budget():
9171
data = config.model_dump()
9272
restored = TriAttentionConfig(**data)
9373
assert restored.budget == 1024
94-
assert restored.target_sparsity_ratio is None
95-
assert restored.prune_interval == 64
96-
97-
98-
def test_config_serialization_roundtrip_sparsity():
99-
"""Config with target_sparsity_ratio survives serialization roundtrip."""
100-
config = TriAttentionConfig(target_sparsity_ratio=0.5, prune_interval=64)
101-
data = config.model_dump()
102-
restored = TriAttentionConfig(**data)
103-
assert restored.budget is None
104-
assert restored.target_sparsity_ratio == 0.5
10574
assert restored.prune_interval == 64
10675

10776

tests/unit/torch/sparsity/kv_cache/test_conversion.py

Lines changed: 1 addition & 26 deletions
Original file line numberDiff line numberDiff line change
@@ -92,37 +92,12 @@ def test_update_metadata():
9292
assert metadata["triattention_config"]["budget"] == 512
9393

9494

95-
def test_convert_metadata_with_sparsity_ratio():
96-
"""Metadata serializes target_sparsity_ratio when set."""
97-
model = nn.Linear(16, 16)
98-
config = TriAttentionConfig(target_sparsity_ratio=0.7)
99-
100-
_, metadata = convert_triattention(model, config)
101-
102-
serialized = metadata["triattention_config"]
103-
assert serialized["target_sparsity_ratio"] == 0.7
104-
assert serialized["budget"] is None
105-
106-
10795
def test_convert_metadata_with_budget():
108-
"""Metadata has budget set and target_sparsity_ratio None."""
96+
"""Metadata has budget set."""
10997
model = nn.Linear(16, 16)
11098
config = TriAttentionConfig(budget=1024)
11199

112100
_, metadata = convert_triattention(model, config)
113101

114102
serialized = metadata["triattention_config"]
115103
assert serialized["budget"] == 1024
116-
assert serialized["target_sparsity_ratio"] is None
117-
118-
119-
def test_update_metadata_with_sparsity_ratio():
120-
"""update_triattention_metadata serializes target_sparsity_ratio."""
121-
model = nn.Linear(16, 16)
122-
config = TriAttentionConfig(target_sparsity_ratio=0.5)
123-
metadata = {}
124-
125-
update_triattention_metadata(model, config, metadata)
126-
127-
assert metadata["triattention_config"]["target_sparsity_ratio"] == 0.5
128-
assert metadata["triattention_config"]["budget"] is None

tests/unit/torch/sparsity/kv_cache/test_scoring.py

Lines changed: 8 additions & 28 deletions
Original file line numberDiff line numberDiff line change
@@ -207,38 +207,18 @@ def test_select_keys_top_k_exceeds_size():
207207
assert mask.shape == scores.shape
208208

209209

210-
def test_select_keys_percentile_basic():
211-
"""Percentile selection evicts target fraction."""
212-
scores = torch.arange(10, dtype=torch.float32)
213-
# sparsity=0.7 → evict 70%, keep top 30% (3 tokens)
214-
mask = select_keys_to_keep(scores, target_sparsity_ratio=0.7)
215-
assert mask.dtype == torch.bool
216-
assert mask.sum().item() == 3
217-
# Top 3 by score are indices 7, 8, 9
218-
assert mask[7].item() is True
219-
assert mask[8].item() is True
220-
assert mask[9].item() is True
221-
222-
223-
def test_select_keys_percentile_half():
224-
"""50% sparsity keeps half the tokens."""
225-
scores = torch.arange(20, dtype=torch.float32)
226-
mask = select_keys_to_keep(scores, target_sparsity_ratio=0.5)
227-
assert mask.sum().item() == 10
228-
229-
230-
def test_select_keys_both_raises():
231-
"""Setting both budget and target_sparsity_ratio raises."""
210+
def test_select_keys_missing_budget_raises():
211+
"""Budget is required for selection."""
232212
scores = torch.rand(10)
233-
with pytest.raises(ValueError, match="exactly one"):
234-
select_keys_to_keep(scores, kv_budget=5, target_sparsity_ratio=0.5)
213+
with pytest.raises(ValueError, match="requires kv_budget"):
214+
select_keys_to_keep(scores)
235215

236216

237-
def test_select_keys_neither_raises():
238-
"""Setting neither budget nor target_sparsity_ratio raises."""
217+
def test_select_keys_non_positive_budget_raises():
218+
"""Budget must be positive."""
239219
scores = torch.rand(10)
240-
with pytest.raises(ValueError, match="exactly one"):
241-
select_keys_to_keep(scores)
220+
with pytest.raises(ValueError, match="must be positive"):
221+
select_keys_to_keep(scores, kv_budget=0)
242222

243223

244224
def test_select_keys_empty_scores():

0 commit comments

Comments
 (0)