This is an automated email from the ASF dual-hosted git repository. tqchen pushed a commit to branch script/canonical-parser-df in repository https://gitbox.apache.org/repos/asf/tvm.git
commit 4f568487e2bb192454954d1bd17c75608d2535b2 Author: Tianqi Chen <[email protected]> AuthorDate: Wed Sep 23 07:31:29 2026 +0000 Keep tree attention reduction axes distinct from key storage --- python/tvm/relax/frontend/nn/llm/tree_attn.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/python/tvm/relax/frontend/nn/llm/tree_attn.py b/python/tvm/relax/frontend/nn/llm/tree_attn.py index 9cf0ccb678..237cebc917 100644 --- a/python/tvm/relax/frontend/nn/llm/tree_attn.py +++ b/python/tvm/relax/frontend/nn/llm/tree_attn.py @@ -416,10 +416,10 @@ def tree_attn(h_kv, h_q, d, dtype, rope_scaling: dict[str, Any], target: Target) with T.sblock(): for li, lj, lk in T.grid(tile_x, tile_z, tile_y): with T.sblock("S_gemm"): - i, j, k = T.axis.remap("SSR", [li, lj, lk]) + i, j, k_axis = T.axis.remap("SSR", [li, lj, lk]) with T.init(): S_local[i, j] = 0.0 - S_local[i, j] += T.cast(Q_smem[i, k], "float32") * T.cast(K_smem[j, k], "float32") * sm_scale * math.log2(math.exp(1)) + S_local[i, j] += T.cast(Q_smem[i, k_axis], "float32") * T.cast(K_smem[j, k_axis], "float32") * sm_scale * math.log2(math.exp(1)) T.tvm_storage_sync("shared") for li, lj in T.grid(tile_x, tile_z): with T.sblock("S_store"): @@ -482,10 +482,10 @@ def tree_attn(h_kv, h_q, d, dtype, rope_scaling: dict[str, Any], target: Target) with T.sblock(): for li, lj, lk in T.grid(tile_x, tile_y, tile_z): with T.sblock("O_gemm"): - i, j, k = T.axis.remap("SSR", [li, lj, lk]) + i, j, k_axis = T.axis.remap("SSR", [li, lj, lk]) with T.init(): O_local[i, j] *= T.exp2(m_prev_smem[i] - m_smem[i]) - O_local[i, j] += S_smem[i, k] * T.cast(V_smem[k, j], "float32") + O_local[i, j] += S_smem[i, k_axis] * T.cast(V_smem[k_axis, j], "float32") # Store O from smem to gmem for li, lj in T.grid(tile_x, tile_y):
