Skip to content
Merged
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
101 changes: 75 additions & 26 deletions mlx/backend/cuda/scaled_dot_product_attention.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -75,6 +75,8 @@ struct SDPACacheKey {
std::array<int64_t, QKV_NDIM> k_strides;
std::array<int64_t, QKV_NDIM> v_strides;
bool do_causal;
std::array<int, QKV_NDIM> mask_shape;
std::array<int64_t, QKV_NDIM> mask_strides;
bool output_logsumexp;
};

Expand All @@ -84,6 +86,7 @@ inline BytesKey<SDPACacheKey> build_sdpa_cache_key(
const array& k,
const array& v,
bool do_causal,
const std::optional<array>& mask_arr,
bool output_logsumexp = true) {
BytesKey<SDPACacheKey> cache_key;
cache_key.pod = {
Expand All @@ -96,8 +99,14 @@ inline BytesKey<SDPACacheKey> build_sdpa_cache_key(
vector_key<QKV_NDIM>(k.strides()),
vector_key<QKV_NDIM>(v.strides()),
do_causal,
{},
{},
output_logsumexp,
};
if (mask_arr) {
cache_key.pod.mask_shape = vector_key<QKV_NDIM>(mask_arr->shape());
cache_key.pod.mask_strides = vector_key<QKV_NDIM>(mask_arr->strides());
}
return cache_key;
}

Expand All @@ -118,6 +127,7 @@ enum UIDS {
K,
V,
SCALE,
BIAS,
O,
STATS,
// Backward graph:
Expand All @@ -133,6 +143,7 @@ fe::graph::Graph build_sdpa_graph(
const array& k,
const array& v,
bool do_causal,
const std::optional<array>& mask_arr,
bool output_logsumexp,
const array& o,
const array& stats) {
Expand Down Expand Up @@ -166,6 +177,11 @@ fe::graph::Graph build_sdpa_graph(
.set_attn_scale(scale)
.set_causal_mask(do_causal)
.set_generate_stats(output_logsumexp);
if (mask_arr) {
auto bias_ = graph.tensor(fe::graph::Tensor_attributes().set_name("BIAS"));
set_tensor_attrs(bias_, BIAS, *mask_arr);
options.set_bias(bias_);
}

auto [o_, stats_] = graph.sdpa(q_, k_, v_, options);
o_->set_output(true);
Expand All @@ -192,6 +208,7 @@ fe::graph::Graph build_sdpa_backward_graph(
const array& k,
const array& v,
bool do_causal,
const std::optional<array>& mask_arr,
const array& o,
const array& d_o,
const array& stats,
Expand Down Expand Up @@ -234,6 +251,11 @@ fe::graph::Graph build_sdpa_backward_graph(
.set_name("sdpa_backward_cudnn")
.set_attn_scale(scale)
.set_causal_mask(do_causal);
if (mask_arr) {
auto bias_ = graph.tensor(fe::graph::Tensor_attributes().set_name("BIAS"));
set_tensor_attrs(bias_, BIAS, *mask_arr);
options.set_bias(bias_);
}

auto [d_q_, d_k_, d_v_] =
graph.sdpa_backward(q_, k_, v_, o_, d_o_, stats_, options);
Expand Down Expand Up @@ -286,8 +308,6 @@ bool supports_sdpa_cudnn(
const array& q,
const array& k,
const array& v,
bool has_mask,
bool do_causal,
Stream s) {
static bool enabled = env::get_var("MLX_CUDA_USE_CUDNN_SPDA", 1);
if (!enabled) {
Expand All @@ -299,17 +319,6 @@ bool supports_sdpa_cudnn(
return false;
}

if (has_mask) {
// TODO: Support array masks.
if (!do_causal) {
return false;
}
// FIXME: Causal mask generates wrong results when L_Q != L_K.
if (q.shape(2) != k.shape(2)) {
return false;
}
}

// Only use cuDNN for prefilling and training.
if (q.shape(2) != k.shape(2)) {
return false;
Expand All @@ -333,6 +342,7 @@ void sdpa_cudnn(
array& o,
array& stats,
bool do_causal,
const std::optional<array>& mask_arr,
bool output_logsumexp,
Stream s) {
auto& encoder = cu::get_command_encoder(s);
Expand All @@ -346,19 +356,21 @@ void sdpa_cudnn(
encoder.set_input_array(k);
encoder.set_input_array(v);
encoder.set_output_array(o);

if (mask_arr) {
encoder.set_input_array(*mask_arr);
}
if (output_logsumexp) {
stats.set_data(cu::malloc_async(stats.nbytes(), encoder));
encoder.set_output_array(stats);
}

// Search cache.
auto cache_key =
build_sdpa_cache_key(encoder, q, k, v, do_causal, output_logsumexp);
auto cache_key = build_sdpa_cache_key(
encoder, q, k, v, do_causal, mask_arr, output_logsumexp);
auto it = sdpa_cache().find(cache_key);
if (it == sdpa_cache().end()) {
auto graph = build_sdpa_graph(
handle, q, k, v, do_causal, output_logsumexp, o, stats);
handle, q, k, v, do_causal, mask_arr, output_logsumexp, o, stats);
it = sdpa_cache().emplace(cache_key, std::move(graph)).first;
}
auto& graph = it->second;
Expand All @@ -369,6 +381,9 @@ void sdpa_cudnn(
{V, const_cast<void*>(gpu_ptr<void>(v))},
{SCALE, &scale},
{O, gpu_ptr<void>(o)}};
if (mask_arr) {
variant_pack[BIAS] = const_cast<void*>(gpu_ptr<void>(*mask_arr));
}
if (output_logsumexp) {
variant_pack[STATS] = gpu_ptr<void>(stats);
}
Expand All @@ -384,6 +399,7 @@ void sdpa_backward_cudnn(
const array& o,
const array& stats,
bool do_causal,
const std::optional<array>& mask_arr,
const array& d_o,
array& d_q,
array& d_k,
Expand All @@ -406,13 +422,16 @@ void sdpa_backward_cudnn(
encoder.set_output_array(d_q);
encoder.set_output_array(d_k);
encoder.set_output_array(d_v);
if (mask_arr) {
encoder.set_input_array(*mask_arr);
}

// Search cache.
auto cache_key = build_sdpa_cache_key(encoder, q, k, v, do_causal);
auto cache_key = build_sdpa_cache_key(encoder, q, k, v, do_causal, mask_arr);
auto it = sdpa_backward_cache().find(cache_key);
if (it == sdpa_backward_cache().end()) {
auto graph = build_sdpa_backward_graph(
handle, q, k, v, do_causal, o, d_o, stats, d_q, d_k, d_v);
handle, q, k, v, do_causal, mask_arr, o, d_o, stats, d_q, d_k, d_v);
it = sdpa_backward_cache().emplace(cache_key, std::move(graph)).first;
}
auto& graph = it->second;
Expand All @@ -428,6 +447,9 @@ void sdpa_backward_cudnn(
{D_Q, gpu_ptr<void>(d_q)},
{D_K, gpu_ptr<void>(d_k)},
{D_V, gpu_ptr<void>(d_v)}};
if (mask_arr) {
variant_pack[BIAS] = const_cast<void*>(gpu_ptr<void>(*mask_arr));
}

execute_graph(encoder, handle, graph, variant_pack);
}
Expand Down Expand Up @@ -469,7 +491,11 @@ bool ScaledDotProductAttention::use_fallback(

return !supports_sdpa_vector(
q, k, v, has_mask, has_arr_mask, do_causal, output_logsumexp) &&
!supports_sdpa_cudnn(q, k, v, has_mask, do_causal, s);
!supports_sdpa_cudnn(q, k, v, s);
}

bool ScaledDotProductAttention::supports_bool_mask() {
return false;
}

void ScaledDotProductAttention::eval_gpu(
Expand All @@ -487,6 +513,11 @@ void ScaledDotProductAttention::eval_gpu(
bool has_mask = inputs.size() - has_sinks_ > 3;
bool has_arr_mask = has_mask && !do_causal_;

std::optional<array> mask_arr;
if (has_arr_mask) {
mask_arr = prepare_sdpa_input(inputs[3], s);
}

if (supports_sdpa_vector(
q, k, v, has_mask, has_arr_mask, do_causal_, output_logsumexp_)) {
if (has_sinks_) {
Expand All @@ -495,7 +526,17 @@ void ScaledDotProductAttention::eval_gpu(
sdpa_vector(q, k, v, scale_, out, do_causal_, std::nullopt, s);
}
} else {
sdpa_cudnn(q, k, v, scale_, out, stats, do_causal_, output_logsumexp_, s);
sdpa_cudnn(
q,
k,
v,
scale_,
out,
stats,
do_causal_,
mask_arr,
output_logsumexp_,
s);
}
}

Expand All @@ -515,21 +556,29 @@ void ScaledDotProductAttentionVJP::eval_gpu(

auto& s = stream();

assert(inputs.size() == 6);
assert(inputs.size() >= 6);
int primals_size = inputs.size() - 3;
bool has_arr_mask = primals_size > 3 + has_sinks_;

array q = prepare_sdpa_input(inputs[0], s);
array k = prepare_sdpa_input(inputs[1], s);
array v = prepare_sdpa_input(inputs[2], s);
array o = prepare_sdpa_input(inputs[3], s);
array stats = prepare_sdpa_input(inputs[4], s);
array d_o = prepare_sdpa_input(inputs[5], s);
array o = prepare_sdpa_input(inputs[primals_size], s);
array stats = prepare_sdpa_input(inputs[primals_size + 1], s);
array d_o = prepare_sdpa_input(inputs[primals_size + 2], s);

std::optional<array> mask_arr;
if (has_arr_mask) {
mask_arr = prepare_sdpa_input(inputs[3], s);
}

assert(outputs.size() == 3);
auto& d_q = outputs[0];
auto& d_k = outputs[1];
auto& d_v = outputs[2];

sdpa_backward_cudnn(
q, k, v, scale_, o, stats, do_causal_, d_o, d_q, d_k, d_v, s);
q, k, v, scale_, o, stats, do_causal_, mask_arr, d_o, d_q, d_k, d_v, s);
}

} // namespace fast
Expand Down
4 changes: 4 additions & 0 deletions mlx/backend/metal/scaled_dot_product_attention.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -569,6 +569,10 @@ bool ScaledDotProductAttention::use_fallback(
return !(supports_sdpa_full || supports_sdpa_vector);
}

bool ScaledDotProductAttention::supports_bool_mask() {
return true;
}

void ScaledDotProductAttention::eval_gpu(
const std::vector<array>& inputs,
std::vector<array>& outputs) {
Expand Down
4 changes: 4 additions & 0 deletions mlx/backend/no_gpu/primitives.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,10 @@ bool fast::ScaledDotProductAttention::use_fallback(
return true;
}

bool fast::ScaledDotProductAttention::supports_bool_mask() {
return false;
}

bool fast::ScaledDotProductAttentionVJP::use_fallback(
const array& q,
Stream s) {
Expand Down
11 changes: 10 additions & 1 deletion mlx/fast.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -800,6 +800,15 @@ array scaled_dot_product_attention(
is_training,
output_logsumexp,
stream)) {
if (has_bool_mask && !ScaledDotProductAttention::supports_bool_mask()) {
// Convert bool mask to additive mask.
float inf = std::numeric_limits<float>::infinity();
array& mask = inputs[3];
mask = where(
mask,
full_like(mask, 0, final_type, s),
full_like(mask, -inf, final_type, s));
}
Shape out_shape{q.shape(0), q.shape(1), q.shape(2), v.shape(-1)};
auto primitive = std::make_shared<ScaledDotProductAttention>(
stream, fallback, scale, do_causal, has_sinks, output_logsumexp);
Expand Down Expand Up @@ -839,7 +848,7 @@ std::vector<array> ScaledDotProductAttention::vjp(

std::vector<Shape> shapes;
std::vector<Dtype> dtypes;
for (int i = 0; i < primals.size(); ++i) {
for (int i = 0; i < /* outputs size */ 3; ++i) {
shapes.push_back(primals[i].shape());
dtypes.push_back(primals[i].dtype());
}
Expand Down
1 change: 1 addition & 0 deletions mlx/fast_primitives.h
Original file line number Diff line number Diff line change
Expand Up @@ -228,6 +228,7 @@ class ScaledDotProductAttention : public Custom {
bool is_training,
bool output_logsumexp,
Stream s);
static bool supports_bool_mask();

void eval_cpu(const std::vector<array>& inputs, std::vector<array>& outputs)
override {
Expand Down
Loading
Loading