[KDA] Fix missing beta sigmoid in PTX prefill (#40685)

Co-authored-by: Mohammad Angkad <mohammad.angkad@radixark.ai>
This commit is contained in:
Mohammad Miadh Angkad
2026-09-21 23:24:38 -07:00
committed by GitHub
co-authored by Mohammad Angkad
parent bc22e1de9e
commit 4c81cd1b09
3 changed files with 56 additions and 19 deletions
@@ -95,8 +95,10 @@ class TestPtxKDATrackRouting(CustomTestCase):
kernel = self._make_kernel()
kernel._triton = _RejectTriton()
h = torch.zeros(3, 2, 128, 128, dtype=torch.float32)
beta_flags = []
def fake_fwd(*args, **kwargs):
beta_flags.append(kwargs["use_beta_sigmoid_in_kernel"])
return [
args[2].clone(), # out == v
kwargs["initial_state"].clone(), # final_state
@@ -107,24 +109,16 @@ class TestPtxKDATrackRouting(CustomTestCase):
kernel._fwd = fake_fwd
x = self._inputs()
out, h_out = kernel.extend(
x["q"],
x["k"],
x["v"],
x["g"],
x["beta"],
ssm_states=x["ssm_states"],
cache_indices=x["cache_indices"],
query_start_loc=x["query_start_loc"],
A_log=x["A_log"],
dt_bias=x["dt_bias"],
return_intermediate_states=True,
track_ssm_h_src=torch.empty(0, dtype=torch.long),
extend_seq_lens_cpu=x["extend_seq_lens_cpu"],
)
self.assertEqual(tuple(out.shape), (1, 164, 2, 128))
self.assertIs(h_out, h)
for beta_kwargs in ({"beta_is_raw": True}, {}, {"beta_is_raw": False}):
out, h_out = kernel.extend(
**x,
**beta_kwargs,
return_intermediate_states=True,
track_ssm_h_src=torch.empty(0, dtype=torch.long),
)
self.assertEqual(tuple(out.shape), (1, 164, 2, 128))
self.assertIs(h_out, h)
self.assertEqual(beta_flags, [True, False, False])
if __name__ == "__main__":