Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
18 commits
Select commit Hold shift + click to select a range
7c1362e
metal: 1-bit uint32 wide-load qdot over upstream #3161; non-aligned-K…
jeethu Oct 11, 2026
11fde19
python/tests: make test_qmv_fast_tail_non_aligned_k honest and future…
jeethu Oct 11, 2026
db2bcd8
quantized: bias-free symmetric 1/2-bit affine (affine_sym) on upstrea…
jeethu Oct 11, 2026
276d40f
quantized: implied-bias affine quantized_matmul/dequantize on upstrea…
jeethu Oct 11, 2026
bc90a8d
quantized: qmv_fast for affine 4-bit K % 256 and 1-bit K % 512 on ups…
jeethu Oct 11, 2026
b4f0dd8
quantized: affine qmv_wide by bit width and batch, single-tile 6-7 ro…
jeethu Oct 11, 2026
5f515ea
metal: dispatch/commit/sync/wait counters (gemma4-qat#181, #282)
jeethu Oct 11, 2026
f0980d5
metal_kernel: build and hash each kernel variant once, bounded memo (…
jeethu Oct 11, 2026
0bf8503
metal: gate NAX off for gen-17 (M5-class) at every NAX dispatch site …
jeethu Oct 11, 2026
faa5ecc
metal: keep transcendentals precise unless MathMode::Fast (#282)
jeethu Oct 11, 2026
385e703
fast: spec_decode_verify op (CPU composition + fused Metal) (#282)
jeethu Oct 11, 2026
bcc9922
python/tests: clarify qmv tail/qmv_wide test comments (#282)
jeethu Oct 11, 2026
8b46f3f
fast: spec_decode_verify materializes non-row-contiguous drafts (#282)
jeethu Oct 11, 2026
52d3ef5
fast: spec_decode_verify CPU path handles K = 0 like Metal (#282)
jeethu Oct 11, 2026
2e4e816
fast: spec_decode_verify returns empty results for B = 0 without disp…
jeethu Oct 11, 2026
78bdb20
fast: spec_decode_verify takes the composition for K = 0 on every bac…
jeethu Oct 11, 2026
87d00a6
metal: condense the gen-17 NAX gate comment (#282)
jeethu Oct 11, 2026
2fc35a3
fast: spec_decode_verify rejects non-integer draft tokens (#282)
jeethu Oct 11, 2026
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
1 change: 1 addition & 0 deletions docs/src/python/fast.rst
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@ Fast
cross_entropy
rope
scaled_dot_product_attention
spec_decode_verify
metal_kernel
cuda_kernel
precompiled_cuda_kernel
179 changes: 132 additions & 47 deletions mlx/backend/common/metal_kernel.cpp
Original file line number Diff line number Diff line change
@@ -1,7 +1,9 @@
// Copyright © 2024 Apple Inc.

#include <iostream>
#include <mutex>
#include <sstream>
#include <unordered_map>

#include "mlx/backend/common/compiled.h"
#include "mlx/backend/metal/metal.h"
Expand Down Expand Up @@ -217,6 +219,59 @@ std::string make_template_hash(const std::string& template_def) {
return template_hash;
}

// The name, source and source hash of one kernel variant.
struct KernelVariant {
std::string name;
std::shared_ptr<const std::string> source;
size_t source_hash;
};

// Shared by all copies of one kernel function. Holds at most
// `metal_kernel_max_cached_variants` entries and is cleared when full. Callers
// keep an owning pointer to their entry, so clearing is safe at any time.
struct VariantCache {
std::mutex mutex;
std::unordered_map<std::string, std::shared_ptr<const KernelVariant>>
variants;
};

// Encodes each call argument that the name and source depend on. Template
// arguments include their kind and name: `int` 1 and `bool` true give the same
// name but a different signature.
std::string variant_key(
const std::vector<array>& inputs,
const std::vector<Dtype>& output_dtypes,
const std::vector<std::pair<std::string, TemplateArg>>& template_args) {
std::string key;
auto append_bytes = [&key](auto v) {
key.append(reinterpret_cast<const char*>(&v), sizeof(v));
};
for (const auto& arr : inputs) {
key += static_cast<char>(arr.dtype().val());
key += arr.ndim() == 0 ? 's'
: arr.size() < max_constant_array_size ? 'c'
: 'd';
}
for (const auto& dtype : output_dtypes) {
key += static_cast<char>(dtype.val());
}
for (const auto& [name, arg] : template_args) {
key += static_cast<char>(arg.index());
append_bytes(name.size());
key += name;
std::visit(
[&](auto v) {
if constexpr (std::is_same_v<decltype(v), Dtype>) {
key += static_cast<char>(v.val());
} else {
append_bytes(v);
}
},
arg);
}
return key;
}

} // namespace

CustomKernelFunction metal_kernel(
Expand Down Expand Up @@ -270,7 +325,10 @@ CustomKernelFunction metal_kernel(
}
}

auto cache = std::make_shared<VariantCache>();

return [=,
cache = std::move(cache),
shape_infos = std::move(shape_infos),
attributes = std::move(attributes)](
const std::vector<array>& inputs,
Expand Down Expand Up @@ -298,60 +356,86 @@ CustomKernelFunction metal_kernel(

auto s = resolve_metal_kernel_stream(s_);

std::string kernel_name = "custom_kernel_" + name;
std::string template_def = "";
if (!template_args.empty()) {
template_def = write_template(template_args);
auto template_hash = make_template_hash(template_def);
kernel_name += "_";
kernel_name += template_hash;
auto key = variant_key(inputs, output_dtypes, template_args);
std::shared_ptr<const KernelVariant> variant;
{
std::lock_guard lock(cache->mutex);
if (auto it = cache->variants.find(key); it != cache->variants.end()) {
variant = it->second;
}
}
if (!variant) {
std::string kernel_name = "custom_kernel_" + name;
std::string template_def = "";
if (!template_args.empty()) {
template_def = write_template(template_args);
auto template_hash = make_template_hash(template_def);
kernel_name += "_";
kernel_name += template_hash;
}

// The generated source depends on the dtypes of the inputs and outputs
// and on how each input is passed (see `write_signature`). Include them
// in the kernel name so that a given name always maps to the same source.
for (const auto& arr : inputs) {
kernel_name += "_";
kernel_name += get_type_string(arr.dtype());
if (arr.ndim() == 0) {
kernel_name += "s";
} else if (arr.size() < max_constant_array_size) {
kernel_name += "c";
// The generated source depends on the dtypes of the inputs and outputs
// and on how each input is passed (see `write_signature`). Include them
// in the kernel name so that a given name always maps to the same
// source.
for (const auto& arr : inputs) {
kernel_name += "_";
kernel_name += get_type_string(arr.dtype());
if (arr.ndim() == 0) {
kernel_name += "s";
} else if (arr.size() < max_constant_array_size) {
kernel_name += "c";
}
}
for (const auto& dtype : output_dtypes) {
kernel_name += "_";
kernel_name += get_type_string(dtype);
}
}
for (const auto& dtype : output_dtypes) {
kernel_name += "_";
kernel_name += get_type_string(dtype);
}

std::string kernel_source = write_signature(
kernel_name,
header,
source,
input_names,
inputs,
output_names,
output_dtypes,
template_args,
attributes,
shape_infos,
atomic_outputs);
std::string kernel_source = write_signature(
kernel_name,
header,
source,
input_names,
inputs,
output_names,
output_dtypes,
template_args,
attributes,
shape_infos,
atomic_outputs);

if (!template_args.empty()) {
template_def = kernel_name + template_def;
kernel_source += "\ntemplate [[host_name(\"";
kernel_source += kernel_name;
kernel_source += "\")]] [[kernel]] decltype(";
kernel_source += template_def;
kernel_source += ") ";
kernel_source += template_def;
kernel_source += ";\n";
if (!template_args.empty()) {
template_def = kernel_name + template_def;
kernel_source += "\ntemplate [[host_name(\"";
kernel_source += kernel_name;
kernel_source += "\")]] [[kernel]] decltype(";
kernel_source += template_def;
kernel_source += ") ";
kernel_source += template_def;
kernel_source += ";\n";
}

auto source_hash = std::hash<std::string>{}(kernel_source);
variant = std::make_shared<const KernelVariant>(KernelVariant{
std::move(kernel_name),
std::make_shared<const std::string>(std::move(kernel_source)),
source_hash});
std::lock_guard lock(cache->mutex);
if (auto it = cache->variants.find(key); it != cache->variants.end()) {
variant = it->second;
} else {
if (cache->variants.size() >= metal_kernel_max_cached_variants) {
cache->variants.clear();
}
cache->variants.emplace(std::move(key), variant);
}
}

if (verbose) {
std::cout << "Generated source code for `" << name << "`:" << std::endl
<< "```" << std::endl
<< kernel_source << std::endl
<< *variant->source << std::endl
<< "```" << std::endl;
}

Expand All @@ -360,8 +444,8 @@ CustomKernelFunction metal_kernel(
std::move(output_dtypes),
std::make_shared<CustomKernel>(
s,
std::move(kernel_name),
std::move(kernel_source),
variant->name,
variant->source,
grid,
threadgroup,
shape_infos,
Expand All @@ -370,7 +454,8 @@ CustomKernelFunction metal_kernel(
std::vector<ScalarArg>{},
false,
0,
compile_options.serialize()),
compile_options.serialize(),
variant->source_hash),
std::move(inputs));
};
}
Expand Down
16 changes: 16 additions & 0 deletions mlx/backend/common/quantized.h
Original file line number Diff line number Diff line change
@@ -1,5 +1,11 @@
// Copyright © 2026 Apple Inc.

#pragma once

#include <optional>

#include "mlx/array.h"

namespace mlx::core {

inline constexpr short get_pack_factor(int bits, int wsize = 8) {
Expand All @@ -11,4 +17,14 @@ inline constexpr short get_bytes_per_pack(int bits, int wsize = 8) {
return power_of_2_bits ? (wsize / 8) : (bits == 5 ? 5 : 3);
}

// Implied bias (affine mode): a 0-d `biases` array holds a single factor f
// standing for the per-group bias `scales * T(f)` in the scales' dtype T.
inline bool is_implied_bias(const array& biases) {
return biases.ndim() == 0;
}

inline bool is_implied_bias(const std::optional<array>& biases) {
return biases.has_value() && is_implied_bias(*biases);
}

} // namespace mlx::core
63 changes: 63 additions & 0 deletions mlx/backend/cpu/quantized.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -943,8 +943,71 @@ void QuantizedMatmul::eval_cpu(const std::vector<array>& inputs, array& out) {
encoder.set_input_array(scales);
encoder.set_output_array(out);
if (mode_ == QuantizationMode::Affine) {
if (inputs.size() < 4) {
throw std::runtime_error(
"[QuantizedMatmul::eval_cpu] Bias-free affine is Metal-only for "
"now; pass biases (e.g. derive_biases) on CPU.");
}
auto biases = ensure_row_contiguous(inputs[3], encoder, stream());
encoder.set_input_array(biases);
if (is_implied_bias(biases)) {
// Implied bias: expand the factor f into a temporary `scales * T(f)`
// for the per-group-bias kernels.
array full_biases(scales.shape(), scales.dtype(), nullptr, {});
full_biases.set_data(allocator::malloc(full_biases.nbytes()));
encoder.add_temporary(full_biases);
encoder.set_output_array(full_biases);
encoder.dispatch([out = array::unsafe_weak_copy(out),
x = array::unsafe_weak_copy(x),
w = array::unsafe_weak_copy(w),
scales = array::unsafe_weak_copy(scales),
factor = array::unsafe_weak_copy(biases),
full_biases = array::unsafe_weak_copy(full_biases),
group_size_ = group_size_,
bits_ = bits_,
transpose_ = transpose_]() mutable {
float f;
switch (factor.dtype()) {
case float32:
f = factor.data<float>()[0];
break;
case float16:
f = static_cast<float>(factor.data<float16_t>()[0]);
break;
case bfloat16:
f = static_cast<float>(factor.data<bfloat16_t>()[0]);
break;
default:
throw std::runtime_error(
"[QuantizedMatmul::eval_cpu] Implied-bias factor must be a "
"real floating type.");
}
auto fill = [&](auto* dst, const auto* src) {
using T = std::remove_pointer_t<decltype(dst)>;
for (size_t i = 0; i < scales.size(); ++i) {
dst[i] = static_cast<T>(f * static_cast<float>(src[i]));
}
};
switch (scales.dtype()) {
case float32:
fill(full_biases.data<float>(), scales.data<float>());
break;
case float16:
fill(full_biases.data<float16_t>(), scales.data<float16_t>());
break;
case bfloat16:
fill(full_biases.data<bfloat16_t>(), scales.data<bfloat16_t>());
break;
default:
throw std::runtime_error(
"[QuantizedMatmul::eval_cpu] Only real floating scales are "
"supported with an implied bias.");
}
_qmm_dispatch(
out, x, w, scales, full_biases, group_size_, bits_, transpose_);
});
return;
}
encoder.dispatch([out = array::unsafe_weak_copy(out),
x = array::unsafe_weak_copy(x),
w = array::unsafe_weak_copy(w),
Expand Down
5 changes: 2 additions & 3 deletions mlx/backend/cuda/custom_kernel.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -310,14 +310,13 @@ void CustomKernel::eval_gpu(
// Compile the custom kernel
std::string kernel_name =
(is_precompiled_) ? name_ : "mlx::core::cu::" + name_;
std::string module_name =
fmt::format("{}_{:x}", name_, std::hash<std::string>{}(source_));
std::string module_name = fmt::format("{}_{:x}", name_, source_hash_);
cu::JitModule& mod = cu::get_jit_module(
encoder.device(),
module_name,
[&]() {
return std::make_tuple(
is_precompiled_, source_, std::vector{kernel_name});
is_precompiled_, *source_, std::vector{kernel_name});
},
false);

Expand Down
2 changes: 1 addition & 1 deletion mlx/backend/cuda/eval.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -83,7 +83,7 @@ void finalize(Stream s) {
cu::get_command_encoder(s).commit();
}

void synchronize(Stream s) {
void synchronize(Stream s, bool) {
nvtx3::scoped_range r("gpu::synchronize");
cu::get_command_encoder(s).synchronize();
}
Expand Down
1 change: 1 addition & 0 deletions mlx/backend/cuda/primitives.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -49,6 +49,7 @@ NO_GPU_MULTI(SVD)
NO_GPU(Inverse)
NO_GPU_MULTI(Eig)
NO_GPU_MULTI(Eigh)
NO_GPU_USE_FALLBACK(fast::SpecDecodeVerify)

namespace fast {
NO_GPU_MULTI(GatedDeltaUpdate)
Expand Down
4 changes: 4 additions & 0 deletions mlx/backend/cuda/quantized/quantized.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,10 @@ void QuantizedMatmul::eval_gpu(const std::vector<array>& inputs, array& out) {
if (inputs.size() > 3) {
biases = inputs[3];
}
if (mode_ == QuantizationMode::Affine && !biases) {
throw std::runtime_error(
"[QuantizedMatmul::eval_gpu] Bias-free affine is Metal-only for now.");
}

auto supports = [&](auto&& f) {
return f(
Expand Down
3 changes: 2 additions & 1 deletion mlx/backend/gpu/eval.h
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,8 @@ void new_stream(Stream s);
void new_thread_unsafe_stream(Stream s);
void eval(array& arr);
void finalize(Stream s);
void synchronize(Stream s);
// explicit_sync is false for internal flushes (e.g. eval error recovery).
void synchronize(Stream s, bool explicit_sync = true);
void clear_streams();

} // namespace mlx::core::gpu
Loading