Use H3 cache constants in tests
This commit is contained in:
@@ -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():
|
||||
|
||||
Reference in New Issue
Block a user