Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
119 changes: 118 additions & 1 deletion benchmarks/python/sdpa_bench.py
Original file line number Diff line number Diff line change
Expand Up @@ -155,13 +155,100 @@ def bench_shape(
return time_mlx_fused, time_mlx_unfused


def set_vjp_fallback(enabled):
os.environ["MLX_SDPA_VJP_FALLBACK"] = "1" if enabled else "0"


def mlx_fused_attn_grads(q, k, v, scale, cotan, mask=None, transpose=False):
def f(q_, k_, v_):
return do_attention(mlx_fused_attn, q_, k_, v_, scale, mask, transpose)

_, grads = mx.vjp(f, [q, k, v], [cotan])
return grads


def do_attention_vjp_bench(q, k, v, scale, cotan, mask=None, transpose=False):
def f(q_, k_, v_):
return do_attention(mlx_fused_attn, q_, k_, v_, scale, mask, transpose)

dq = q

for i in range(N_iter_func):
_, (dq, dk, dv) = mx.vjp(f, [dq, k, v], [cotan])

mx.eval([dq, dk, dv])
return dq


def peak_mem_vjp(q, k, v, scale, cotan, mask=None, transpose=False):
mx.clear_cache()
mx.reset_peak_memory()
grads = mlx_fused_attn_grads(q, k, v, scale, cotan, mask, transpose)
mx.eval(grads)
peak = mx.get_peak_memory() / float(1024.0**3)
del grads
mx.clear_cache()
return peak


def max_rel_diff(a_list, b_list):
rel = 0.0
for a, b in zip(a_list, b_list):
denom = mx.maximum(mx.max(mx.abs(a)), mx.array(1e-6, mx.float32))
rel = max(rel, (mx.max(mx.abs(a - b)) / denom).item())
return rel


def bench_shape_vjp(
B, qsl, ksl, head_dim, n_q_heads, n_kv_heads, dtype, transpose=True, mask_in=None
):
q_mx, k_mx, v_mx, scale, mask = prepare_inputs(
B, qsl, ksl, head_dim, n_q_heads, n_kv_heads, mask_in, transpose, dtype
)
cotan = mx.array(np.random.normal(0.0, 1.0, q_mx.shape).astype(getattr(np, dtype)))
mx.eval(cotan)

set_vjp_fallback(True)
time_fallback = bench(
do_attention_vjp_bench, q_mx, k_mx, v_mx, scale, cotan, mask, transpose
)
mem_fallback = peak_mem_vjp(q_mx, k_mx, v_mx, scale, cotan, mask, transpose)
g_fallback = mlx_fused_attn_grads(q_mx, k_mx, v_mx, scale, cotan, mask, transpose)
mx.eval(g_fallback)

set_vjp_fallback(False)
time_vjp = bench(
do_attention_vjp_bench, q_mx, k_mx, v_mx, scale, cotan, mask, transpose
)
mem_vjp = peak_mem_vjp(q_mx, k_mx, v_mx, scale, cotan, mask, transpose)
g_vjp = mlx_fused_attn_grads(q_mx, k_mx, v_mx, scale, cotan, mask, transpose)
mx.eval(g_vjp)

rel = max_rel_diff(g_fallback, g_vjp)
# nax truncates so the accuracy is not full fp32
atol = 5e-3

if rel > atol:
print(
f"Failed at (B: {B}, qsl: {qsl}, ksl: {ksl}, head_dim: {head_dim}, n_qh: {n_q_heads}, n_kvh: {n_kv_heads}, mask: {mask_in}) [tpose = {transpose}] with max rel = {rel:3.2e}"
)

return time_vjp, time_fallback, mem_vjp, mem_fallback, rel


def get_gflop_count(B, M, N, K):
return float(2.0 * N_iter_bench * N_iter_func * B * M * N * K) / float(1024.0**3)


if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Run gemm benchmarks")

parser.add_argument(
"-bw",
"--backward",
action="store_true",
help="benchmark the vjp against the fallback",
)
args = parser.parse_args()
dtypes = ("float16", "float32")[:1]
transposes = (False,)

Expand Down Expand Up @@ -228,6 +315,36 @@ def get_gflop_count(B, M, N, K):

shapes = shapes_64 + shapes_72 + shapes_80 + shapes_96 + shapes_128 + shapes_256

if args.backward:
masks = [None, "causal"]

print(
" B, qsl, ksl, hdim, n_qh, n_kvh, t, dtype, mask, t_fall, t_vjp, speedup, m_fall, m_vjp, mem_x, rel"
)

for dtype in ["float32"]:
for transpose in transposes:
for B, qsl, ksl, head_dim, n_q_heads, n_kv_heads in shapes:
for mask_in in masks:
t_vjp, t_fall, m_vjp, m_fall, rel = bench_shape_vjp(
B,
qsl,
ksl,
head_dim,
n_q_heads,
n_kv_heads,
dtype,
transpose,
mask_in,
)
speedup = t_fall / t_vjp
mem_x = m_fall / max(m_vjp, 1e-9)
t_str = 1 if transpose else 0
print(
f"{B:3d}, {qsl:5d}, {ksl:5d}, {head_dim:4d}, {n_q_heads:4d}, {n_kv_heads:5d}, {t_str:1d}, {dtype}, {str(mask_in):>8}, {t_fall: 2.3f}, {t_vjp: 2.3f}, {speedup:6.2f}x, {m_fall:6.2f}, {m_vjp:6.2f}, {mem_x:5.2f}x, {rel:3.1e}"
)
exit(0)

masks = [None, "bool", "causal"]

print(
Expand Down
1 change: 1 addition & 0 deletions mlx/backend/metal/kernels/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -56,6 +56,7 @@ build_kernel(random)
build_kernel(rms_norm)
build_kernel(rope)
build_kernel(scaled_dot_product_attention sdpa_vector.h)
build_kernel(scaled_dot_product_attention_vjp)
build_kernel(gated_delta_update gated_delta_update.h)
if(MLX_METAL_VERSION GREATER_EQUAL 320)
build_kernel(fence)
Expand Down
154 changes: 154 additions & 0 deletions mlx/backend/metal/kernels/scaled_dot_product_attention_vjp.metal
Original file line number Diff line number Diff line change
@@ -0,0 +1,154 @@
// Copyright © 2026 Apple Inc.

#include "mlx/backend/metal/kernels/utils.h"

using namespace metal;

template <typename InT, int Dv>
[[kernel]] void sdpa_vjp_odo(
const device InT* o [[buffer(0)]], // [B, H, qL, Dv]
const device InT* cot_o [[buffer(1)]], // [B, H, qL, Dv]
device float* odo [[buffer(2)]], // [B, H, qL]
constant int& qL [[buffer(3)]],
uint3 tpg [[thread_position_in_grid]],
uint simd_lane_id [[thread_index_in_simdgroup]]) {
const uint bh = tpg.z;
const int i = int(tpg.y);

if (i >= qL) {
return;
}

const size_t row = (size_t)bh * qL + i;
auto o_i = o + row * Dv;
auto co_i = cot_o + row * Dv;

float acc = 0.0f;
for (int d = int(simd_lane_id); d < Dv; d += 32) {
acc += float(o_i[d]) * float(co_i[d]);
}
acc = simd_sum(acc);

if (simd_lane_id == 0) {
odo[row] = acc;
}
}

struct SDPAVJPTileParams {
int bq;
int bk;
int qL;
int i0;
int j0;
float scale;
int diag_off;
int causal;
};

template <typename T>
[[kernel]] void sdpa_vjp_ds(
const device T* S [[buffer(0)]], // [BH, bq, bk]
const device T* dP [[buffer(1)]], // [BH, bq, bk]
const device float* lse [[buffer(2)]], // [B, H, qL]
const device float* odo [[buffer(3)]], // [B, H, qL]
device T* dS [[buffer(4)]], // [BH, bq, bk]
device T* P [[buffer(5)]], // [BH, bk, bq]
const constant SDPAVJPTileParams& p [[buffer(6)]],
uint3 gid [[thread_position_in_grid]]) {
int col = int(gid.x);
int row = int(gid.y);
int bh = int(gid.z);
if (row >= p.bq || col >= p.bk) {
return;
}

size_t sbase = (size_t(bh) * size_t(p.bq) + size_t(row)) * size_t(p.bk);
size_t qi = size_t(bh) * size_t(p.qL) + size_t(p.i0 + row);

float l = lse[qi];
float dlt = odo[qi];
// Fully masked rows carry lse == -inf; exp(-inf - -inf) would be NaN.
bool dead = (l < 0.0f) && metal::isinf(l);
int lim = p.i0 + row + p.diag_off;

float pv;
if (dead || (p.causal != 0 && (p.j0 + col) > lim)) {
pv = 0.0f;
} else {
pv = metal::fast::exp(static_cast<float>(S[sbase + col]) * p.scale - l);
}
float dsv = pv * (static_cast<float>(dP[sbase + col]) - dlt) * p.scale;

dS[sbase + col] = static_cast<T>(dsv);
P[sbase + col] = static_cast<T>(pv);
}

template <typename T, bool Accum>
[[kernel]] void sdpa_vjp_reduce(
const device T* src [[buffer(0)]],
device float* acc [[buffer(1)]],
device T* out [[buffer(2)]],
const constant int& rows [[buffer(3)]],
const constant int& dim [[buffer(4)]],
const constant int& group [[buffer(5)]],
const constant int& acc_rows [[buffer(6)]],
const constant int& row_off [[buffer(7)]],
uint3 gid [[thread_position_in_grid]]) {
int c = int(gid.x);
int row = int(gid.y);
int bh = int(gid.z);
if (c >= dim || row >= rows) {
return;
}

size_t sbase = size_t(bh) * size_t(group) * size_t(rows) * size_t(dim);
float sum = 0.0f;
for (int g = 0; g < group; ++g) {
size_t si = sbase +
(size_t(g) * size_t(rows) + size_t(row)) * size_t(dim) + size_t(c);
sum += static_cast<float>(src[si]);
}

size_t ai =
(size_t(bh) * size_t(acc_rows) + size_t(row_off + row)) * size_t(dim) +
size_t(c);
float total = Accum ? acc[ai] + sum : sum;
acc[ai] = total;
out[ai] = static_cast<T>(total);
}

// Instantiations
#define instantiate_odo(in_type, dv) \
instantiate_kernel( \
"sdpa_vjp_odo_" #in_type "_" #dv, sdpa_vjp_odo, in_type, dv)

#define instantiate_odo_shapes(in_type) \
instantiate_odo(in_type, 64); \
instantiate_odo(in_type, 72); \
instantiate_odo(in_type, 80); \
instantiate_odo(in_type, 96); \
instantiate_odo(in_type, 128); \
instantiate_odo(in_type, 192); \
instantiate_odo(in_type, 256); \
instantiate_odo(in_type, 512); \

instantiate_odo_shapes(bfloat16_t);
instantiate_odo_shapes(float16_t);
instantiate_odo_shapes(float);

#define instantiate_sdpa_vjp_ds(tname, type) \
instantiate_kernel("sdpa_vjp_ds_" #tname, sdpa_vjp_ds, type)

instantiate_sdpa_vjp_ds(float, float);
instantiate_sdpa_vjp_ds(float16_t, float16_t);
instantiate_sdpa_vjp_ds(bfloat16_t, bfloat16_t);

#define instantiate_sdpa_vjp_reduce(tname, type) \
instantiate_kernel( \
"sdpa_vjp_reduce_add_" #tname, sdpa_vjp_reduce, type, true) \
instantiate_kernel( \
"sdpa_vjp_reduce_set_" #tname, sdpa_vjp_reduce, type, false)

instantiate_sdpa_vjp_reduce(float, float);
instantiate_sdpa_vjp_reduce(float16_t, float16_t);
instantiate_sdpa_vjp_reduce(bfloat16_t, bfloat16_t);
24 changes: 23 additions & 1 deletion mlx/backend/metal/kernels/steel/attn/kernels/steel_attention.h
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@ constant bool align_K [[function_constant(201)]];
constant bool has_mask [[function_constant(300)]];
constant bool do_causal [[function_constant(301)]];
constant bool has_sinks [[function_constant(302)]];
constant bool save_lse [[function_constant(303)]];

struct MaxOp {
template <typename T>
Expand Down Expand Up @@ -77,6 +78,7 @@ template <
const constant AttnMaskParams* mask_params [[buffer(5), function_constant(has_mask)]],
const device MaskType* mask [[buffer(6), function_constant(has_mask)]],
const device T* sinks [[buffer(7), function_constant(has_sinks)]],
device float* lse [[buffer(8), function_constant(save_lse)]],
uint simd_lane_id [[thread_index_in_simdgroup]],
uint simd_group_id [[simdgroup_index_in_threadgroup]],
uint3 tid [[threadgroup_position_in_grid]],
Expand All @@ -103,6 +105,12 @@ template <
tidl.y * params->O_strides[1] + // Head
tidl.x * BQ * params->O_strides[2]; // Sequence

if (save_lse) {
lse += tidl.z * params->H * params->qL + // Batch
tidl.y * params->qL + // Head
tidl.x * BQ; // Sequence
}

if (has_mask) {
mask += tidl.z * mask_params->M_strides[0] + // Batch
tidl.y * mask_params->M_strides[1]; // Head
Expand Down Expand Up @@ -532,6 +540,20 @@ template <
Otile.template row_bin_op<DivOp>(sum_score);
threadgroup_barrier(mem_flags::mem_none);

// Store the logsumexp for the backward pass
if (save_lse && sn == 0) {
using stile_t = decltype(Stile);
const bool is_last_q = int(tid.x) == (params->NQ_aligned);

STEEL_PRAGMA_UNROLL
for (short i = 0; i < kRowsPT; ++i) {
const short r = tm + sm + i * stile_t::kFragRows;
if (align_Q || !is_last_q || r < params->qL_rem) {
lse[r] = M_LN2_F * (max_score[i] + metal::log2(sum_score[i]));
}
}
}

// Store results
O += (tm + sm) * params->O_strides[2] + d_half * BVh + sn;

Expand All @@ -545,4 +567,4 @@ template <
} else {
Otile.template store<T, 1, 1>(O, params->O_strides[2]);
}
}
}
Loading
Loading