From 3f6ff40628a82cd68060fbc9585da7149b2da2c3 Mon Sep 17 00:00:00 2001 From: continuousml Date: Fri, 24 Jul 2026 13:10:05 -0700 Subject: [PATCH] Clean up splash attention imports Import the general splash attention path from the tokamax dependency; the ring path keeps the adapted kernels via tokamax_ring_attention. --- src/maxtext/layers/attention_op.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/maxtext/layers/attention_op.py b/src/maxtext/layers/attention_op.py index cf076b4794..8c40613266 100644 --- a/src/maxtext/layers/attention_op.py +++ b/src/maxtext/layers/attention_op.py @@ -67,8 +67,6 @@ from maxtext.kernels.attention import tokamax_ring_attention from maxtext.kernels.attention.ragged_attention import ragged_gqa from maxtext.kernels.attention.ragged_attention import ragged_mha -from maxtext.kernels.tokamax_splash_attention import splash_attention_kernel as tokamax_splash_kernel -from maxtext.kernels.tokamax_splash_attention import splash_attention_mask as tokamax_splash_mask from maxtext.layers import nnx_wrappers from maxtext.layers.initializers import variable_to_logically_partitioned from maxtext.layers.quantizations import AqtQuantization as Quant @@ -77,6 +75,8 @@ import numpy as np from tokamax._src.ops.attention import base as tokamax_attention_base from tokamax._src.ops.attention import pallas_triton as tokamax_pallas_triton +from tokamax._src.ops.experimental.tpu.splash_attention import splash_attention_kernel as tokamax_splash_kernel +from tokamax._src.ops.experimental.tpu.splash_attention import splash_attention_mask as tokamax_splash_mask # pylint: disable=line-too-long, g-doc-args, g-doc-return-or-yield, bad-continuation, g-inconsistent-quotes # pytype: disable=attribute-error