From 693639a7d34c7df693fe8fda22e9ee03a2bb9ee4 Mon Sep 17 00:00:00 2001 From: Disty0 Date: Sun, 28 Jun 2026 22:07:37 +0300 Subject: [PATCH] SDNQ atten fix enable_gqa + is_causal --- modules/sdnq/kernels/triton_atten.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/modules/sdnq/kernels/triton_atten.py b/modules/sdnq/kernels/triton_atten.py index 8b2895cbf..72a3afb59 100644 --- a/modules/sdnq/kernels/triton_atten.py +++ b/modules/sdnq/kernels/triton_atten.py @@ -116,7 +116,7 @@ def sdnq_attn_kernel( lo = 0 if is_causal: - hi = tl.minimum(KN, (start_m + 1) * BLOCK_SIZE_M + KN - QN) + hi = tl.minimum(KN, (start_m + 1) * BLOCK_SIZE_M) else: hi = KN @@ -141,7 +141,7 @@ def sdnq_attn_kernel( qk = tl.dot(q, k, out_dtype=tl.float32) if is_causal: - causal_mask = (offs_m[:, None] + (KN - QN)) >= (start_n + offs_n[None, :]) + causal_mask = offs_m[:, None] >= (start_n + offs_n[None, :]) qk = tl.where(causal_mask, qk, -float('inf')) if mask_ptr is not None: