Use H3 cache constants in tests

This commit is contained in:
xmarre
2026-08-05 21:39:01 +02:00
parent 94da0a6c57
commit f7bacf6ccd
+5 -5
View File
@@ -88,9 +88,9 @@ def runtime_for(config, inner, *, packed_len=6, beta=0.7, selected=None):
"sage_lse_unavailable": None,
"debug_state": None,
"debug_topology": (packed_len,),
"attention_bias_cache": {},
"text_index_cache": {},
"scaled_rope_cache": {},
h3._ATTENTION_BIAS_CACHE_KEY: {},
h3._TEXT_INDEX_CACHE_KEY: {},
h3._ROPE_CACHE_KEY: {},
}
@@ -644,7 +644,7 @@ def test_scaled_rope_is_computed_once_per_forward_and_invalidated_by_tensor(monk
for replacement in replacements:
replacement(block_args(runtime, rope), {"original_block": original})
assert count == 1
cached_source, cached_scaled = next(iter(runtime["scaled_rope_cache"].values()))
cached_source, cached_scaled = next(iter(runtime[h3._ROPE_CACHE_KEY].values()))
assert cached_source is rope
assert torch.is_tensor(cached_scaled)
replacement = replacements[0]
@@ -754,7 +754,7 @@ def test_runtime_masks_are_isolated_between_configured_clones():
mask_b = h3._compact_text_bias(runtime_b, q)
assert mask_a is not mask_b
assert not torch.equal(mask_a, mask_b)
assert runtime_a["attention_bias_cache"] is not runtime_b["attention_bias_cache"]
assert runtime_a[h3._ATTENTION_BIAS_CACHE_KEY] is not runtime_b[h3._ATTENTION_BIAS_CACHE_KEY]
def test_interrupted_attention_call_does_not_poison_next_runtime():