Layr-Labs/mlx

diff: ignored:
+8102
-2151
+471
-0

This is an overview of the changes in Layr-Labs/mlx, a fork of ml-explore/mlx.

The fork is the core of Layr-Labs’ Apple-silicon inference stack — mlx → mlx-c → mlx-swift → mlx-swift-lm — and carries the kernel, allocator and Metal-runtime work that serving production MoE models (Gemma 4, Qwen 3.5⁄3.6, GPT-OSS) needed before or beyond what upstream ships: sorted expert-tile quantized matmul routes, a masked-tail MXFP4 decode kernel, FP32-safe SDPA partials, coherent allocator accounting for admission control, declared-mutable custom-kernel inputs for paged KV, and a use-after-free fix in Metal eval. Every change is meant to be additive and rebasable; each is gated by shape, device or environment checks so the upstream path stays the default elsewhere.

One section is not fork work at all: the fork carries upstream’s own v0.32.2 changes (applied in PR #10 rather than merged), and those files are listed as such so the page stays honest about what Layr-Labs actually changed.

gpu::eval captured the MTL::CommandBuffer* before eval_gpu ran and attached the input-buffer liveness handler to it afterwards. Any primitive that calls CommandEncoder::synchronize() mid-eval — the expert-tile route’s sortedness-retract check does — commits and replaces the encoder’s buffer, so the handler landed on a dangling pointer (latent on Gemma E=128, a deterministic SIGSEGV on Qwen E=256). The buffer is now fetched from the encoder after eval_gpu, so inputs outlive exactly the tail of the primitive’s work. (PR #5)

diff --git ml-explore/mlx/mlx/backend/metal/eval.cpp Layr-Labs/mlx/mlx/backend/metal/eval.cpp index 43a6fb993288a5ad2e1422ccea6357e89a921d33..aac977187836ab316c1f0831f8f0f88215170744 100644 --- ml-explore/mlx/mlx/backend/metal/eval.cpp +++ Layr-Labs/mlx/mlx/backend/metal/eval.cpp @@ -30,7 +30,6 @@ void eval(array& arr) { auto pool = metal::new_scoped_memory_pool(); auto s = arr.primitive().stream(); auto& encoder = metal::get_command_encoder(s); - auto* command_buffer = encoder.get_command_buffer();   auto outputs = arr.outputs(); { @@ -41,7 +40,8 @@ if (arr.is_tracer()) { inputs = arr.inputs(); }   - debug_set_primitive_buffer_label(command_buffer, arr.primitive()); + debug_set_primitive_buffer_label( + encoder.get_command_buffer(), arr.primitive()); arr.primitive().eval_gpu(arr.inputs(), outputs); } std::unordered_set<std::shared_ptr<array::Data>> buffers; @@ -63,7 +63,13 @@ encoder.commit([s, buffers = std::move(buffers)]() { scheduler::notify_task_completion(s); }); } else { - command_buffer->addCompletedHandler( + // Fetch the command buffer AFTER eval_gpu: primitives that synchronize + // mid-eval (e.g. the expert-tile route's descriptor retract check) + // commit and REPLACE the encoder's buffer, so a pointer captured before + // eval_gpu would be stale here — attaching the buffer-liveness handler + // to it is a use-after-free. The current buffer holds the tail of this + // primitive's work, which is exactly what the inputs must outlive. + encoder.get_command_buffer()->addCompletedHandler( [buffers = std::move(buffers)](MTL::CommandBuffer* cbuf) {}); } }

An opt-in route (MLX_GATHER_QMM_EXPERT_SLICES=1|trust) for gathered, affine-quantized MoE expert matmuls: a descriptor kernel builds sorted expert tiles on device (one instantiation per expert count) and a single expert-count-agnostic tile kernel consumes them. classify_gemma4_expert_qmm is a pure route table admitting exactly the shapes the kernel is built for; availability is all-or-nothing on a source-matched metallib so an older library fails the route closed; trust skips the descriptor-retract readback when the caller guarantees sorted indices. Diagnostics counters (mlx_metal_gemma4_expert_qmm_diagnostics_*) are exposed through a C ABI so mlx-c / mlx-swift can arm, snapshot and reset them. Also in these kernels: every affine quantized-vector bias-sum operand is promoted to the accumulator type before summation (PR #15).

diff --git ml-explore/mlx/mlx/backend/common/gemma4_expert_qmm.h Layr-Labs/mlx/mlx/backend/common/gemma4_expert_qmm.h new file mode 100644 index 0000000000000000000000000000000000000000..1a1ce37542e0695eb006370d7901f3668e34126d --- /dev/null +++ Layr-Labs/mlx/mlx/backend/common/gemma4_expert_qmm.h @@ -0,0 +1,293 @@ +// Copyright © 2023-2024 Apple Inc. + +#pragma once + +#include <stdint.h> + +#include "mlx/api.h" + +#if defined(__APPLE__) +#ifdef __cplusplus +extern "C" { +#endif + +typedef struct mlx_metal_gemma4_expert_qmm_diagnostics { + uint8_t requested; + uint8_t aot_available; + uint8_t nax_available; + uint8_t armed; + uint64_t attempts; + uint64_t hits; + uint64_t fallback_nax; + uint64_t fallback_outer_route; + uint64_t fallback_quantization; + uint64_t fallback_topology; + uint64_t fallback_assignment_count; + uint64_t fallback_geometry; + uint64_t fallback_metallib_unavailable; + uint64_t fallback_sortedness_retracted; +} mlx_metal_gemma4_expert_qmm_diagnostics; + +MLX_API void mlx_metal_gemma4_expert_qmm_diagnostics_snapshot( + mlx_metal_gemma4_expert_qmm_diagnostics* diagnostics); +MLX_API void mlx_metal_gemma4_expert_qmm_diagnostics_reset(void); +MLX_API void mlx_metal_gemma4_expert_qmm_diagnostics_clear_and_arm(void); +MLX_API void mlx_metal_gemma4_expert_qmm_diagnostics_snapshot_and_disarm( + mlx_metal_gemma4_expert_qmm_diagnostics* diagnostics); + +#ifdef __cplusplus +} +#endif +#endif + +#ifdef __cplusplus + +#include <atomic> + +namespace mlx::core::metal { + +enum class Gemma4ExpertQMMRoute : uint8_t { + not_requested, + hit, + fallback_nax, + fallback_outer_route, + fallback_quantization, + fallback_topology, + fallback_assignment_count, + fallback_geometry, + fallback_metallib_unavailable, + fallback_sortedness_retracted, +}; + +struct Gemma4ExpertQMMRouteInput { + bool requested{false}; + bool aot_available{false}; + bool nax_available{false}; + bool outer_route{false}; + + bool affine{false}; + bool transpose{false}; + bool has_bias{false}; + bool indices_uint32{false}; + bool indices_contiguous{false}; + bool x_bfloat16{false}; + bool x_contiguous{false}; + bool w_uint32{false}; + bool w_contiguous{false}; + bool scales_bfloat16{false}; + bool scales_contiguous{false}; + bool biases_bfloat16{false}; + bool biases_contiguous{false}; + + int group_size{0}; + int bits{0}; + int expert_count{0}; + int assignments{0}; + int index_count{0}; + int k{0}; + int n{0}; + + int x_rank{0}; + int x_dim0{0}; + int x_dim1{0}; + int x_dim2{0}; + int w_rank{0}; + int w_dim0{0}; + int w_dim1{0}; + int w_dim2{0}; + int scales_rank{0}; + int scales_dim0{0}; + int scales_dim1{0}; + int scales_dim2{0}; + int biases_rank{0}; + int biases_dim0{0}; + int biases_dim1{0}; + int biases_dim2{0}; +}; + +inline Gemma4ExpertQMMRoute classify_gemma4_expert_qmm( + const Gemma4ExpertQMMRouteInput& input) { + if (!input.requested) { + return Gemma4ExpertQMMRoute::not_requested; + } + if (!input.outer_route) { + return Gemma4ExpertQMMRoute::fallback_outer_route; + } + // The existing NAX route owns every supported BF16/transposed RHS call and + // must win before the Gemma 4 specialization or its AOT capability matters. + if (input.nax_available) { + return Gemma4ExpertQMMRoute::fallback_nax; + } + if (!input.affine || !input.transpose || !input.has_bias || + !input.indices_uint32 || !input.indices_contiguous || !input.x_bfloat16 || + !input.x_contiguous || !input.w_uint32 || !input.w_contiguous || + !input.scales_bfloat16 || !input.scales_contiguous || + !input.biases_bfloat16 || !input.biases_contiguous || + input.group_size != 64 || input.bits != 4) { + return Gemma4ExpertQMMRoute::fallback_quantization; + } + const bool gemma4 = input.expert_count == 128; + const bool qwen36 = input.expert_count == 256; + if ((!gemma4 && !qwen36) || input.x_rank != 3 || + input.x_dim0 != input.assignments || input.x_dim1 != 1 || + input.x_dim2 != input.k || input.w_rank != 3 || + input.w_dim0 != input.expert_count || input.scales_rank != 3 || + input.scales_dim0 != input.expert_count || input.biases_rank != 3 || + input.biases_dim0 != input.expert_count || + input.index_count != input.assignments) { + return Gemma4ExpertQMMRoute::fallback_topology; + } + if (input.assignments != 4096 && input.assignments != 8192 && + input.assignments != 16384) { + return Gemma4ExpertQMMRoute::fallback_assignment_count; + } + + // Whole-projection geometry for one expert matrix [E, n, k] at W4/g64: + // packed weight columns k/8 (eight 4-bit values per uint32) and one + // scale/bias per 64-wide group, k/64 columns. The quantization gate above + // guarantees bits==4 and group_size==64, so the divisions are exact. + auto projection = [&input](int k, int n) { + return input.k == k && input.n == n && input.w_dim1 == n && + input.w_dim2 == k / 8 && input.scales_dim1 == n && + input.scales_dim2 == k / 64 && input.biases_dim1 == n && + input.biases_dim2 == k / 64; + }; + // Gemma 4 26B-A4B (E=128): gate/up [128,1408,2816] and down [128,2816,704]. + // Qwen 3.5/3.6 35B-A3B (E=256): fused gate_up [256,1024,2048], split + // gate/up [256,512,2048], and down [256,2048,512]. The tile kernel itself + // is expert-count agnostic; only the descriptor builder instantiation + // differs (one thread per expert). + const bool hit_geometry = gemma4 + ? (projection(2816, 1408) || projection(704, 2816)) + : (projection(2048, 1024) || projection(2048, 512) || + projection(512, 2048)); + if (!hit_geometry) { + return Gemma4ExpertQMMRoute::fallback_geometry; + } + if (!input.aot_available) { + return Gemma4ExpertQMMRoute::fallback_metallib_unavailable; + } + return Gemma4ExpertQMMRoute::hit; +} + +struct Gemma4ExpertQMMCounterSnapshot { + uint64_t hits{0}; + uint64_t fallback_nax{0}; + uint64_t fallback_outer_route{0}; + uint64_t fallback_quantization{0}; + uint64_t fallback_topology{0}; + uint64_t fallback_assignment_count{0}; + uint64_t fallback_geometry{0}; + uint64_t fallback_metallib_unavailable{0}; + uint64_t fallback_sortedness_retracted{0}; + bool armed{false}; + + uint64_t attempts() const { + return hits + fallback_nax + fallback_outer_route + fallback_quantization + + fallback_topology + fallback_assignment_count + fallback_geometry + + fallback_metallib_unavailable + fallback_sortedness_retracted; + } +}; + +class Gemma4ExpertQMMCounters { + public: + bool armed() const { + return armed_.load(std::memory_order_relaxed); + } + + // Recording is called only after the caller's relaxed-atomic armed branch. + // Keeping the branch at that boundary makes the unarmed inference path free + // of counter atomic operations while the engine-idle arm/disarm contract + // makes access to armed_ well-defined. + void record(Gemma4ExpertQMMRoute route) { + std::atomic<uint64_t>* counter = nullptr; + switch (route) { + case Gemma4ExpertQMMRoute::not_requested: + return; + case Gemma4ExpertQMMRoute::hit: + counter = &hits_; + break; + case Gemma4ExpertQMMRoute::fallback_nax: + counter = &fallback_nax_; + break; + case Gemma4ExpertQMMRoute::fallback_outer_route: + counter = &fallback_outer_route_; + break; + case Gemma4ExpertQMMRoute::fallback_quantization: + counter = &fallback_quantization_; + break; + case Gemma4ExpertQMMRoute::fallback_topology: + counter = &fallback_topology_; + break; + case Gemma4ExpertQMMRoute::fallback_assignment_count: + counter = &fallback_assignment_count_; + break; + case Gemma4ExpertQMMRoute::fallback_geometry: + counter = &fallback_geometry_; + break; + case Gemma4ExpertQMMRoute::fallback_metallib_unavailable: + counter = &fallback_metallib_unavailable_; + break; + case Gemma4ExpertQMMRoute::fallback_sortedness_retracted: + counter = &fallback_sortedness_retracted_; + break; + } + counter->fetch_add(1, std::memory_order_relaxed); + } + + Gemma4ExpertQMMCounterSnapshot snapshot() const { + return { + hits_.load(std::memory_order_relaxed), + fallback_nax_.load(std::memory_order_relaxed), + fallback_outer_route_.load(std::memory_order_relaxed), + fallback_quantization_.load(std::memory_order_relaxed), + fallback_topology_.load(std::memory_order_relaxed), + fallback_assignment_count_.load(std::memory_order_relaxed), + fallback_geometry_.load(std::memory_order_relaxed), + fallback_metallib_unavailable_.load(std::memory_order_relaxed), + fallback_sortedness_retracted_.load(std::memory_order_relaxed), + armed_.load(std::memory_order_relaxed), + }; + } + + Gemma4ExpertQMMCounterSnapshot snapshot_and_disarm() { + const bool was_armed = armed_.load(std::memory_order_relaxed); + armed_.store(false, std::memory_order_relaxed); + auto result = snapshot(); + result.armed = was_armed; + return result; + } + + void reset() { + hits_.store(0, std::memory_order_relaxed); + fallback_nax_.store(0, std::memory_order_relaxed); + fallback_outer_route_.store(0, std::memory_order_relaxed); + fallback_quantization_.store(0, std::memory_order_relaxed); + fallback_topology_.store(0, std::memory_order_relaxed); + fallback_assignment_count_.store(0, std::memory_order_relaxed); + fallback_geometry_.store(0, std::memory_order_relaxed); + fallback_metallib_unavailable_.store(0, std::memory_order_relaxed); + fallback_sortedness_retracted_.store(0, std::memory_order_relaxed); + } + + void clear_and_arm() { + reset(); + armed_.store(true, std::memory_order_relaxed); + } + + private: + std::atomic<bool> armed_{false}; + std::atomic<uint64_t> hits_{0}; + std::atomic<uint64_t> fallback_nax_{0}; + std::atomic<uint64_t> fallback_outer_route_{0}; + std::atomic<uint64_t> fallback_quantization_{0}; + std::atomic<uint64_t> fallback_topology_{0}; + std::atomic<uint64_t> fallback_assignment_count_{0}; + std::atomic<uint64_t> fallback_geometry_{0}; + std::atomic<uint64_t> fallback_metallib_unavailable_{0}; + std::atomic<uint64_t> fallback_sortedness_retracted_{0}; +}; + +} // namespace mlx::core::metal + +#endif
diff --git ml-explore/mlx/mlx/backend/metal/device.cpp Layr-Labs/mlx/mlx/backend/metal/device.cpp index 65df5c108cffddcb46cf3b822a18f3f5402db42f..c4d7de0586a79a07ae7aa9a1c9bee44f3263cfed 100644 --- ml-explore/mlx/mlx/backend/metal/device.cpp +++ Layr-Labs/mlx/mlx/backend/metal/device.cpp @@ -1,5 +1,7 @@ // Copyright © 2023-2024 Apple Inc.   +#include <algorithm> +#include <cctype> #include <cstdlib> #include <sstream>   @@ -496,19 +498,15 @@ concurrent_outputs_.clear(); all_inputs_.clear(); }   -void CommandEncoder::signal_event( - std::shared_ptr<EventImpl> event, - uint64_t value) { +void CommandEncoder::signal_event(Event event, uint64_t value) { end_encoding(); - buffer_->encodeSignalEvent(event->mtl_event(), value); + buffer_->encodeSignalEvent(event.cast<EventImpl>().mtl_event(), value); signal_events_.push_back({std::move(event), value}); }   -void CommandEncoder::wait_event( - std::shared_ptr<EventImpl> event, - uint64_t value) { +void CommandEncoder::wait_event(Event event, uint64_t value) { end_encoding(); - buffer_->encodeWait(event->mtl_event(), value); + buffer_->encodeWait(event.cast<EventImpl>().mtl_event(), value); wait_events_.push_back(std::move(event)); }   @@ -525,35 +523,37 @@ buffer_->addCompletedHandler( [&error_ = error_, wait_events = std::move(wait_events_), signal_events = std::move(signal_events_), - completion = std::move(completion)](MTL::CommandBuffer* cbuf) { + completion = std::move(completion)](MTL::CommandBuffer* cbuf) mutable { if (completion) { completion(); } // If any of the waited event has error in it, poison the encoder. for (auto& event : wait_events) { - if (event->error()) { - error_ = event->error(); + if (error_.store_if_valid(event.load_error())) { break; } } // Set error only when no error happended before, to preserve the // earliest error. - if (!error_ && cbuf->status() == MTL::CommandBufferStatusError) { - error_ = std::make_shared<std::string>(fmt::format( - "[METAL] Command buffer execution failed: {}.", - cbuf->error()->localizedDescription()->utf8String())); + bool has_error = error_.valid(); + if (!has_error && cbuf->status() == MTL::CommandBufferStatusError) { + error_.set_message( + std::make_shared<std::string>(fmt::format( + "[METAL] Command buffer execution failed: {}.", + cbuf->error()->localizedDescription()->utf8String()))); + has_error = true; } // Poison all the signaled events when error happened. - if (error_) { + if (has_error) { for (auto& [event, value] : signal_events) { - event->set_error(error_); + event.set_error(error_); } } // Metal won't signal the events for us on error, manually signal them // to avoid infinite waiting. if (cbuf->status() == MTL::CommandBufferStatusError) { for (auto& [event, value] : signal_events) { - event->signal(value); + event.cast<EventImpl>().signal(value); } } }); @@ -570,20 +570,17 @@ end_encoding(); commit(); cbuf->waitUntilCompleted();   - if (error_ && !exiting_) { - auto error = std::move(error_); - throw std::runtime_error(*error); + if (!exiting_) { + error_.check(); } }   MTL::ComputeCommandEncoder* CommandEncoder::get_command_encoder() { if (!encoder_) { + error_.check(); encoder_ = NS::RetainPtr( buffer_->computeCommandEncoder(MTL::DispatchTypeConcurrent)); fence_ = NS::TransferPtr(device_.mtl_device()->newFence()); - // Reset error when user starts to encode new commands, they are supposed to - // have handled the error in synchronize() or Event::wait(). - error_.reset(); } return encoder_.get(); } @@ -591,6 +588,53 @@ Device::Device() : device_(load_device()), residency_sets_(device_.get()) { auto pool = new_scoped_memory_pool(); default_library_ = NS::TransferPtr(load_default_library(device_.get())); + + std::string expert_qmm_env = env::get_var("MLX_GATHER_QMM_EXPERT_SLICES", ""); + std::transform( + expert_qmm_env.begin(), + expert_qmm_env.end(), + expert_qmm_env.begin(), + [](unsigned char c) { return static_cast<char>(std::tolower(c)); }); + gemma4_expert_qmm_requested_ = expert_qmm_env == "1" || + expert_qmm_env == "true" || expert_qmm_env == "on" || + expert_qmm_env == "yes" || expert_qmm_env == "trust"; + // "trust" additionally skips the descriptor-retract readback: the caller + // asserts its sorted-indices contract is machine-guaranteed (the Swift + // SwitchGLU path sorts on-device), so the host never drains the stream to + // observe a retracted build. Under trust, a genuinely mis-sorted input + // produces undefined tile output instead of the legacy fallback. + gemma4_expert_qmm_trust_sorted_ = expert_qmm_env == "trust"; + + constexpr const char* descriptor_kernel = + "build_gemma4_sorted_expert_tiles_bm32"; + constexpr const char* descriptor_kernel_e256 = + "build_sorted_expert_tiles_bm32_e256"; + constexpr const char* tile_kernel = + "affine_gather_qmm_gemma4_expert_tiles_bfloat16_t_gs_64_b_4_" + "alN_true_bm_32_bn_32_bk_32"; + auto has_default_function = [this](const char* name) { + auto ns_name = NS::String::string(name, NS::ASCIIStringEncoding); + auto function = NS::TransferPtr(default_library_->newFunction(ns_name)); + return function.get() != nullptr; + }; + // All expert-tile symbols ship from one source-matched metallib + // (scripts/fetch-metallib.sh completeness contract), so availability is + // all-or-nothing: a metallib missing any of them predates this revision + // and must fail the whole route closed. + gemma4_expert_qmm_aot_available_ = has_default_function(descriptor_kernel) && + has_default_function(descriptor_kernel_e256) && + has_default_function(tile_kernel); + if (gemma4_expert_qmm_requested_ && gemma4_expert_qmm_aot_available_) { + try { + // Resolve the pipelines once so missing or incompatible packaged AOT + // assets fail closed before an inference command encoder is touched. + get_kernel(descriptor_kernel); + get_kernel(descriptor_kernel_e256); + get_kernel(tile_kernel); + } catch (...) { + gemma4_expert_qmm_aot_available_ = false; + } + } arch_ = env::metal_gpu_arch(); if (arch_.empty()) { arch_ = std::string(device_->architecture()->name()->utf8String()); @@ -971,3 +1015,69 @@ #endif }   } // namespace mlx::core::metal + +#if defined(__APPLE__) +namespace { +void gemma4_expert_qmm_diagnostics_snapshot( + mlx_metal_gemma4_expert_qmm_diagnostics* diagnostics, + bool disarm) { + if (diagnostics == nullptr) { + return; + } + *diagnostics = {}; + try { + auto& d = mlx::core::metal::device(mlx::core::Device::gpu); + const auto counters = disarm + ? d.gemma4_expert_qmm_counter_snapshot_and_disarm() + : d.gemma4_expert_qmm_counter_snapshot(); + diagnostics->requested = d.gemma4_expert_qmm_requested(); + diagnostics->aot_available = d.gemma4_expert_qmm_aot_available(); + diagnostics->nax_available = mlx::core::metal::is_nax_available(); + diagnostics->armed = counters.armed; + diagnostics->attempts = counters.attempts(); + diagnostics->hits = counters.hits; + diagnostics->fallback_nax = counters.fallback_nax; + diagnostics->fallback_outer_route = counters.fallback_outer_route; + diagnostics->fallback_quantization = counters.fallback_quantization; + diagnostics->fallback_topology = counters.fallback_topology; + diagnostics->fallback_assignment_count = counters.fallback_assignment_count; + diagnostics->fallback_geometry = counters.fallback_geometry; + diagnostics->fallback_metallib_unavailable = + counters.fallback_metallib_unavailable; + diagnostics->fallback_sortedness_retracted = + counters.fallback_sortedness_retracted; + } catch (...) { + // Diagnostics are optional. A missing Metal device must remain observable + // as an all-zero snapshot rather than escaping an exception through C ABI. + } +} +} // namespace + +extern "C" void mlx_metal_gemma4_expert_qmm_diagnostics_snapshot( + mlx_metal_gemma4_expert_qmm_diagnostics* diagnostics) { + gemma4_expert_qmm_diagnostics_snapshot(diagnostics, false); +} + +extern "C" void mlx_metal_gemma4_expert_qmm_diagnostics_reset(void) { + try { + mlx::core::metal::device(mlx::core::Device::gpu) + .reset_gemma4_expert_qmm_counters(); + } catch (...) { + // Reset is best-effort on hosts without an accessible Metal device. + } +} + +extern "C" void mlx_metal_gemma4_expert_qmm_diagnostics_clear_and_arm(void) { + try { + mlx::core::metal::device(mlx::core::Device::gpu) + .clear_and_arm_gemma4_expert_qmm_counters(); + } catch (...) { + // Arming is best-effort on hosts without an accessible Metal device. + } +} + +extern "C" void mlx_metal_gemma4_expert_qmm_diagnostics_snapshot_and_disarm( + mlx_metal_gemma4_expert_qmm_diagnostics* diagnostics) { + gemma4_expert_qmm_diagnostics_snapshot(diagnostics, true); +} +#endif
diff --git ml-explore/mlx/mlx/backend/metal/device.h Layr-Labs/mlx/mlx/backend/metal/device.h index 3bb1e9e3b32bbfac95660e6d29e7c78f702a7af9..5599047bf9309ffadf634208e0f0d8b61850c333 100644 --- ml-explore/mlx/mlx/backend/metal/device.h +++ Layr-Labs/mlx/mlx/backend/metal/device.h @@ -6,11 +6,11 @@ #include <Metal/Metal.hpp> #include <functional> #include <mutex> #include <shared_mutex> -#include <string> #include <unordered_map> #include <unordered_set>   #include "mlx/array.h" +#include "mlx/backend/common/gemma4_expert_qmm.h" #include "mlx/backend/common/metal_kernel.h" #include "mlx/backend/metal/resident.h" #include "mlx/device.h" @@ -21,7 +21,6 @@ using MTLFCList = std::vector<std::tuple<const void*, MTL::DataType, NS::UInteger>>;   class Device; -class EventImpl;   class MLX_API CommandEncoder { public: @@ -92,8 +91,8 @@ }   void barrier(); void end_encoding(); - void wait_event(std::shared_ptr<EventImpl> event, uint64_t value); - void signal_event(std::shared_ptr<EventImpl> event, uint64_t value); + void wait_event(Event event, uint64_t value); + void signal_event(Event event, uint64_t value); bool needs_commit() const; void commit(std::function<void()> completion = nullptr); void synchronize(); @@ -119,11 +118,11 @@ ResidencySets& residency_sets_; uint64_t sets_attached_{0};   // The events hooked to current command buffer. - std::vector<std::shared_ptr<EventImpl>> wait_events_; - std::vector<std::tuple<std::shared_ptr<EventImpl>, uint64_t>> signal_events_; + std::vector<Event> wait_events_; + std::vector<std::tuple<Event, uint64_t>> signal_events_;   // Error from previous commited command buffer. - std::shared_ptr<std::string> error_; + Error error_;   // Encoder for issuing GPU commands. // The members are used within a single ComputeCommandEncoder and will be @@ -201,6 +200,46 @@ ResidencySets& residency_sets() { return residency_sets_; }   + bool gemma4_expert_qmm_requested() const { + return gemma4_expert_qmm_requested_; + } + + // MLX_GATHER_QMM_EXPERT_SLICES=trust: skip the descriptor-retract + // readback in the expert-tile route (no mid-eval stream drain). The + // caller asserts sorted indices are machine-guaranteed; a violation + // yields undefined tile output instead of the legacy fallback. + bool gemma4_expert_qmm_trust_sorted() const { + return gemma4_expert_qmm_trust_sorted_; + } + + bool gemma4_expert_qmm_aot_available() const { + return gemma4_expert_qmm_aot_available_; + } + bool gemma4_expert_qmm_diagnostics_armed() const { + return gemma4_expert_qmm_counters_.armed(); + } + + // Call only inside a route boundary guarded by + // gemma4_expert_qmm_diagnostics_armed(). + void record_armed_gemma4_expert_qmm(Gemma4ExpertQMMRoute route) { + gemma4_expert_qmm_counters_.record(route); + } + + Gemma4ExpertQMMCounterSnapshot gemma4_expert_qmm_counter_snapshot() const { + return gemma4_expert_qmm_counters_.snapshot(); + } + Gemma4ExpertQMMCounterSnapshot + gemma4_expert_qmm_counter_snapshot_and_disarm() { + return gemma4_expert_qmm_counters_.snapshot_and_disarm(); + } + + void reset_gemma4_expert_qmm_counters() { + gemma4_expert_qmm_counters_.reset(); + } + void clear_and_arm_gemma4_expert_qmm_counters() { + gemma4_expert_qmm_counters_.clear_and_arm(); + } + private: NS::SharedPtr<MTL::Library> build_library_( const std::string& source_string, @@ -240,6 +279,10 @@ std::shared_mutex kernel_mtx_; std::shared_mutex library_mtx_; std::unordered_map<std::string, NS::SharedPtr<MTL::Library>> library_map_; NS::SharedPtr<MTL::Library> default_library_; + bool gemma4_expert_qmm_requested_{false}; + bool gemma4_expert_qmm_trust_sorted_{false}; + bool gemma4_expert_qmm_aot_available_{false}; + Gemma4ExpertQMMCounters gemma4_expert_qmm_counters_; std::unordered_map< MTL::Library*, std::unordered_map<std::string, NS::SharedPtr<MTL::ComputePipelineState>>>
diff --git ml-explore/mlx/mlx/backend/metal/kernels/quantized.h Layr-Labs/mlx/mlx/backend/metal/kernels/quantized.h index 6d87dc770fd9dc7a7fce52b5a9c7439146ef3970..17e74d1db14b0328d04d54f3e3df55d4ddddcae4 100644 --- ml-explore/mlx/mlx/backend/metal/kernels/quantized.h +++ Layr-Labs/mlx/mlx/backend/metal/kernels/quantized.h @@ -36,7 +36,7 @@ U sum = 0;   if (bits == 2) { for (int i = 0; i < values_per_thread; i += 4) { - sum += x[i] + x[i + 1] + x[i + 2] + x[i + 3]; + sum += U(x[i]) + U(x[i + 1]) + U(x[i + 2]) + U(x[i + 3]); x_thread[i] = x[i]; x_thread[i + 1] = x[i + 1] / 4.0f; x_thread[i + 2] = x[i + 2] / 16.0f; @@ -46,8 +46,8 @@ }   else if (bits == 3) { for (int i = 0; i < values_per_thread; i += 8) { - sum += x[i] + x[i + 1] + x[i + 2] + x[i + 3] + x[i + 4] + x[i + 5] + - x[i + 6] + x[i + 7]; + sum += U(x[i]) + U(x[i + 1]) + U(x[i + 2]) + U(x[i + 3]) + U(x[i + 4]) + + U(x[i + 5]) + U(x[i + 6]) + U(x[i + 7]); x_thread[i] = x[i]; x_thread[i + 1] = x[i + 1] / 8.0f; x_thread[i + 2] = x[i + 2] / 64.0f; @@ -61,7 +61,7 @@ }   else if (bits == 4) { for (int i = 0; i < values_per_thread; i += 4) { - sum += x[i] + x[i + 1] + x[i + 2] + x[i + 3]; + sum += U(x[i]) + U(x[i + 1]) + U(x[i + 2]) + U(x[i + 3]); x_thread[i] = x[i]; x_thread[i + 1] = x[i + 1] / 16.0f; x_thread[i + 2] = x[i + 2] / 256.0f; @@ -71,8 +71,8 @@ }   else if (bits == 5) { for (int i = 0; i < values_per_thread; i += 8) { - sum += x[i] + x[i + 1] + x[i + 2] + x[i + 3] + x[i + 4] + x[i + 5] + - x[i + 6] + x[i + 7]; + sum += U(x[i]) + U(x[i + 1]) + U(x[i + 2]) + U(x[i + 3]) + U(x[i + 4]) + + U(x[i + 5]) + U(x[i + 6]) + U(x[i + 7]); x_thread[i] = x[i]; x_thread[i + 1] = x[i + 1] / 32.0f; x_thread[i + 2] = x[i + 2] / 4.0f; @@ -86,7 +86,7 @@ }   else if (bits == 6) { for (int i = 0; i < values_per_thread; i += 4) { - sum += x[i] + x[i + 1] + x[i + 2] + x[i + 3]; + sum += U(x[i]) + U(x[i + 1]) + U(x[i + 2]) + U(x[i + 3]); x_thread[i] = x[i]; x_thread[i + 1] = x[i + 1] / 64.0f; x_thread[i + 2] = x[i + 2] / 16.0f; @@ -115,7 +115,7 @@ U sum = 0;   if (bits == 2) { for (int i = 0; i < N; i += 4) { - sum += x[i] + x[i + 1] + x[i + 2] + x[i + 3]; + sum += U(x[i]) + U(x[i + 1]) + U(x[i + 2]) + U(x[i + 3]); x_thread[i] = x[i]; x_thread[i + 1] = x[i + 1] / 4.0f; x_thread[i + 2] = x[i + 2] / 16.0f; @@ -125,8 +125,8 @@ }   else if (bits == 3) { for (int i = 0; i < N; i += 8) { - sum += x[i] + x[i + 1] + x[i + 2] + x[i + 3] + x[i + 4] + x[i + 5] + - x[i + 6] + x[i + 7]; + sum += U(x[i]) + U(x[i + 1]) + U(x[i + 2]) + U(x[i + 3]) + U(x[i + 4]) + + U(x[i + 5]) + U(x[i + 6]) + U(x[i + 7]);   x_thread[i] = x[i]; x_thread[i + 1] = x[i + 1] / 8.0f; @@ -141,7 +141,7 @@ }   else if (bits == 4) { for (int i = 0; i < N; i += 4) { - sum += x[i] + x[i + 1] + x[i + 2] + x[i + 3]; + sum += U(x[i]) + U(x[i + 1]) + U(x[i + 2]) + U(x[i + 3]); x_thread[i] = x[i]; x_thread[i + 1] = x[i + 1] / 16.0f; x_thread[i + 2] = x[i + 2] / 256.0f; @@ -151,8 +151,8 @@ }   else if (bits == 5) { for (int i = 0; i < N; i += 8) { - sum += x[i] + x[i + 1] + x[i + 2] + x[i + 3] + x[i + 4] + x[i + 5] + - x[i + 6] + x[i + 7]; + sum += U(x[i]) + U(x[i + 1]) + U(x[i + 2]) + U(x[i + 3]) + U(x[i + 4]) + + U(x[i + 5]) + U(x[i + 6]) + U(x[i + 7]); x_thread[i] = x[i]; x_thread[i + 1] = x[i + 1] / 32.0f; x_thread[i + 2] = x[i + 2] / 4.0f; @@ -166,7 +166,7 @@ }   else if (bits == 6) { for (int i = 0; i < N; i += 4) { - sum += x[i] + x[i + 1] + x[i + 2] + x[i + 3]; + sum += U(x[i]) + U(x[i + 1]) + U(x[i + 2]) + U(x[i + 3]); x_thread[i] = x[i]; x_thread[i + 1] = x[i + 1] / 64.0f; x_thread[i + 2] = x[i + 2] / 16.0f; @@ -489,17 +489,16 @@ bits == 2 || bits == 3 || bits == 4 || bits == 5 || bits == 6 || bits == 8, "Template undefined for bits not in {2, 3, 4, 5, 6, 8}");   + const float s = float(scale); + const float b = float(bias); + if (bits == 2) { - U s[4] = { - scale, - scale / static_cast<U>(4.0f), - scale / static_cast<U>(16.0f), - scale / static_cast<U>(64.0f)}; + float sc[4] = {s, s / 4.0f, s / 16.0f, s / 64.0f}; for (int i = 0; i < (N / 4); i++) { - w_local[4 * i] = s[0] * (w[i] & 0x03) + bias; - w_local[4 * i + 1] = s[1] * (w[i] & 0x0c) + bias; - w_local[4 * i + 2] = s[2] * (w[i] & 0x30) + bias; - w_local[4 * i + 3] = s[3] * (w[i] & 0xc0) + bias; + w_local[4 * i] = static_cast<U>(sc[0] * (w[i] & 0x03) + b); + w_local[4 * i + 1] = static_cast<U>(sc[1] * (w[i] & 0x0c) + b); + w_local[4 * i + 2] = static_cast<U>(sc[2] * (w[i] & 0x30) + b); + w_local[4 * i + 3] = static_cast<U>(sc[3] * (w[i] & 0xc0) + b); } }   @@ -508,22 +507,24 @@ for (int i = 0; i < (N / 8); i++) { w_local += 8 * i; w += 3 * i;   - w_local[0] = (w[0] & 0x7) * scale + bias; - w_local[1] = ((w[0] & 0x38) >> 3) * scale + bias; - w_local[2] = (((w[0] & 0xc0) >> 6) + ((w[1] & 0x1) << 2)) * scale + bias; - w_local[3] = ((w[1] & 0xe) >> 1) * scale + bias; - w_local[4] = ((w[1] & 0x70) >> 4) * scale + bias; - w_local[5] = (((w[1] & 0x80) >> 7) + ((w[2] & 0x3) << 1)) * scale + bias; - w_local[6] = ((w[2] & 0x1c) >> 2) * scale + bias; - w_local[7] = ((w[2] & 0xe0) >> 5) * scale + bias; + w_local[0] = static_cast<U>((w[0] & 0x7) * s + b); + w_local[1] = static_cast<U>(((w[0] & 0x38) >> 3) * s + b); + w_local[2] = + static_cast<U>((((w[0] & 0xc0) >> 6) + ((w[1] & 0x1) << 2)) * s + b); + w_local[3] = static_cast<U>(((w[1] & 0xe) >> 1) * s + b); + w_local[4] = static_cast<U>(((w[1] & 0x70) >> 4) * s + b); + w_local[5] = + static_cast<U>((((w[1] & 0x80) >> 7) + ((w[2] & 0x3) << 1)) * s + b); + w_local[6] = static_cast<U>(((w[2] & 0x1c) >> 2) * s + b); + w_local[7] = static_cast<U>(((w[2] & 0xe0) >> 5) * s + b); } }   else if (bits == 4) { - U s[2] = {scale, scale / static_cast<U>(16.0f)}; + float sc[2] = {s, s / 16.0f}; for (int i = 0; i < (N / 2); i++) { - w_local[2 * i] = s[0] * (w[i] & 0x0f) + bias; - w_local[2 * i + 1] = s[1] * (w[i] & 0xf0) + bias; + w_local[2 * i] = static_cast<U>(sc[0] * (w[i] & 0x0f) + b); + w_local[2 * i + 1] = static_cast<U>(sc[1] * (w[i] & 0xf0) + b); } }   @@ -532,14 +533,18 @@ for (int i = 0; i < (N / 8); i++) { w_local += 8 * i; w += 5 * i;   - w_local[0] = (w[0] & 0x1f) * scale + bias; - w_local[1] = (((w[0] & 0xe0) >> 5) + ((w[1] & 0x3) << 3)) * scale + bias; - w_local[2] = ((w[1] & 0x7c) >> 2) * scale + bias; - w_local[3] = (((w[1] & 0x80) >> 7) + ((w[2] & 0xf) << 1)) * scale + bias; - w_local[4] = (((w[2] & 0xf0) >> 4) + ((w[3] & 0x1) << 4)) * scale + bias; - w_local[5] = ((w[3] & 0x3e) >> 1) * scale + bias; - w_local[6] = (((w[3] & 0xc0) >> 6) + ((w[4] & 0x7) << 2)) * scale + bias; - w_local[7] = ((w[4] & 0xf8) >> 3) * scale + bias; + w_local[0] = static_cast<U>((w[0] & 0x1f) * s + b); + w_local[1] = + static_cast<U>((((w[0] & 0xe0) >> 5) + ((w[1] & 0x3) << 3)) * s + b); + w_local[2] = static_cast<U>(((w[1] & 0x7c) >> 2) * s + b); + w_local[3] = + static_cast<U>((((w[1] & 0x80) >> 7) + ((w[2] & 0xf) << 1)) * s + b); + w_local[4] = + static_cast<U>((((w[2] & 0xf0) >> 4) + ((w[3] & 0x1) << 4)) * s + b); + w_local[5] = static_cast<U>(((w[3] & 0x3e) >> 1) * s + b); + w_local[6] = + static_cast<U>((((w[3] & 0xc0) >> 6) + ((w[4] & 0x7) << 2)) * s + b); + w_local[7] = static_cast<U>(((w[4] & 0xf8) >> 3) * s + b); } }   @@ -547,16 +552,18 @@ else if (bits == 6) { for (int i = 0; i < (N / 4); i++) { w_local += 4 * i; w += 3 * i; - w_local[0] = (w[0] & 0x3f) * scale + bias; - w_local[1] = (((w[0] >> 6) & 0x03) + ((w[1] & 0x0f) << 2)) * scale + bias; - w_local[2] = (((w[1] >> 4) & 0x0f) + ((w[2] & 0x03) << 4)) * scale + bias; - w_local[3] = ((w[2] >> 2) & 0x3f) * scale + bias; + w_local[0] = static_cast<U>((w[0] & 0x3f) * s + b); + w_local[1] = + static_cast<U>((((w[0] >> 6) & 0x03) + ((w[1] & 0x0f) << 2)) * s + b); + w_local[2] = + static_cast<U>((((w[1] >> 4) & 0x0f) + ((w[2] & 0x03) << 4)) * s + b); + w_local[3] = static_cast<U>(((w[2] >> 2) & 0x3f) * s + b); } }   else if (bits == 8) { for (int i = 0; i < N; i++) { - w_local[i] = scale * w[i] + bias; + w_local[i] = static_cast<U>(s * w[i] + b); } } } @@ -1190,6 +1197,134 @@ const bool aligned_N, const int BM = 32, const int BK = 32, const int BN = 32> +METAL_FUNC void qmm_t_expert_impl( + const device uint32_t* w, + const device T* scales, + const device T* biases, + const device T* x, + device T* y, + threadgroup T* Xs, + threadgroup T* Ws, + const constant int& K, + const constant int& N, + const int M, + const constant int& K_eff, + uint3 tid [[threadgroup_position_in_grid]], + uint lid [[thread_index_in_threadgroup]], + uint simd_gid [[simdgroup_index_in_threadgroup]], + uint simd_lid [[thread_index_in_simdgroup]]) { + static_assert(BK >= SIMD_SIZE, "BK should be larger than SIMD_SIZE"); + static_assert(BK % SIMD_SIZE == 0, "BK should be divisible by SIMD_SIZE"); + + (void)lid; + + constexpr int WM = 2; + constexpr int WN = 2; + constexpr int pack_factor = get_pack_factor<bits, 8>(); + constexpr int bytes_per_pack = get_bytes_per_pack<bits>(); + + constexpr int BK_padded = (BK + 16 / sizeof(T)); + + // Instantiate the appropriate BlockMMA and Loader + using mma_t = mlx::steel:: + BlockMMA<T, T, BM, BN, BK, WM, WN, false, true, BK_padded, BK_padded>; + using loader_x_t = + mlx::steel::BlockLoader<T, BM, BK, BK_padded, 1, WM * WN * SIMD_SIZE>; + using loader_w_t = QuantizedBlockLoader< + T, + BN, + BK, + BK_padded, + 1, + WM * WN * SIMD_SIZE, + group_size, + bits>; + + // Set the block + const int K_w = K * bytes_per_pack / pack_factor; + const int K_g = K / group_size; + const int y_row = tid.y * BM; + const int y_col = tid.x * BN; + + auto wl = (const device uint8_t*)w; + + x += y_row * static_cast<int64_t>(K); + wl += y_col * K_w; + scales += y_col * K_g; + biases += y_col * K_g; + y += y_row * static_cast<int64_t>(N) + y_col; + + // Make the x loader and mma operation + const short num_els = min(BM, M - y_row); + const short num_outs = min(BN, N - y_col); + loader_x_t loader_x(x, K, Xs, simd_gid, simd_lid); + loader_w_t loader_w(wl, scales, biases, K, Ws, simd_gid, simd_lid); + mma_t mma_op(simd_gid, simd_lid); + + if (num_els < BM) { + if (!aligned_N && num_outs < BN) { + for (int k = 0; k < K_eff; k += BK) { + threadgroup_barrier(mem_flags::mem_threadgroup); + loader_x.load_safe(short2(BK, num_els)); + loader_w.load_safe(short2(BK, num_outs)); + threadgroup_barrier(mem_flags::mem_threadgroup); + mma_op.mma(Xs, Ws); + loader_x.next(); + loader_w.next(); + } + } else { + for (int k = 0; k < K_eff; k += BK) { + threadgroup_barrier(mem_flags::mem_threadgroup); + loader_x.load_safe(short2(BK, num_els)); + loader_w.load_unsafe(); + threadgroup_barrier(mem_flags::mem_threadgroup); + mma_op.mma(Xs, Ws); + loader_x.next(); + loader_w.next(); + } + } + } else { + if (!aligned_N && num_outs < BN) { + for (int k = 0; k < K_eff; k += BK) { + threadgroup_barrier(mem_flags::mem_threadgroup); + loader_x.load_unsafe(); + loader_w.load_safe(short2(BK, num_outs)); + threadgroup_barrier(mem_flags::mem_threadgroup); + mma_op.mma(Xs, Ws); + loader_x.next(); + loader_w.next(); + } + } else { + for (int k = 0; k < K_eff; k += BK) { + threadgroup_barrier(mem_flags::mem_threadgroup); + loader_x.load_unsafe(); + loader_w.load_unsafe(); + threadgroup_barrier(mem_flags::mem_threadgroup); + + mma_op.mma(Xs, Ws); + loader_x.next(); + loader_w.next(); + } + } + } + + // Store results to device memory + threadgroup_barrier(mem_flags::mem_threadgroup); + if (num_els < BM || num_outs < BN) { + mma_op.store_result_safe(y, N, short2(num_outs, num_els)); + } else { + mma_op.store_result(y, N); + } +} + +template < + typename T, + const int group_size, + const int bits, + const bool aligned_N, + const int BM = 32, + const int BK = 32, + const int BN = 32> METAL_FUNC void qmm_t_impl( const device uint32_t* w, const device T* scales, @@ -1596,7 +1731,8 @@ typename T, int group_size, int bits, bool batched, - bool has_global_scale = false> + bool has_global_scale = false, + int results_per_simdgroup = 4> [[kernel]] void affine_qmv_fast( const device uint32_t* w [[buffer(0)]], const device T* scales [[buffer(1)]], @@ -1653,7 +1789,8 @@ typename T, int group_size, const int bits, bool batched, - bool has_global_scale = false> + bool has_global_scale = false, + int results_per_simdgroup = 4> [[kernel]] void affine_qmv( const device uint32_t* w [[buffer(0)]], const device T* scales [[buffer(1)]], @@ -2411,6 +2548,230 @@ b_strides, tid); qmm_n_impl<T, group_size, bits, BM, BK, BN>( w, scales, biases, x, y, Xs, Ws, K, N, M, tid, lid, simd_gid, simd_lid); +} + +// Descriptor builder for the sorted expert-tile route. One thread per expert +// (threadgroup size == NE); instantiated for NE=128 (Gemma 4) and NE=256 +// (Qwen 3.5/3.6 MoE) in quantized.metal. +template <int NE> +[[kernel]] void build_sorted_expert_tiles_bm32( + const device uint32_t* indices [[buffer(0)]], + device uint4* descriptors [[buffer(1)]], + device uint* count [[buffer(2)]], + const constant int& M [[buffer(3)]], + uint lid [[thread_index_in_threadgroup]], + uint simd_gid [[simdgroup_index_in_threadgroup]], + uint simd_lid [[thread_index_in_simdgroup]]) { + constexpr uint expert_count = uint(NE); + constexpr uint BM = 32; + constexpr uint simdgroup_count = expert_count / 32; + threadgroup uint segment_starts[expert_count + 1]; + threadgroup uint inclusive_tile_offsets[expert_count]; + threadgroup uint violation_votes[simdgroup_count]; + + // One thread finds each expert's first sorted row. The last thread also + // supplies the sentinel, so all NE+1 boundaries are ready after one barrier. + int lower = 0; + int upper = M; + while (lower < upper) { + const int midpoint = lower + (upper - lower) / 2; + if (indices[midpoint] < lid) { + lower = midpoint + 1; + } else { + upper = midpoint; + } + } + segment_starts[lid] = uint(lower); + if (lid == expert_count - 1) { + segment_starts[expert_count] = uint(M); + } + + // The binary search is only sound when `indices` is non-decreasing across + // the whole array; a mis-sorted input would silently mis-attribute rows to + // experts. The check is twofold. Each thread verifies its own segment + // boundary against the generalized invariant + // `indices[start - 1] < lid <= indices[start]` (edge threads have only one + // neighbor to check). Independently of that search, a strided adjacent-pair + // scan validates `indices[i - 1] <= indices[i]` for every i in [1, M): + // thread `lid` covers i = lid + 1, lid + NE + 1, ..., so the NE threads + // between them inspect every adjacent pair exactly once (for the reachable + // M in {4096, 8192, 16384} that is a bounded number of iterations each). + // Adjacent-pair monotonicity is transitive, so a clean scan is a sound and + // complete proof that the array is globally non-decreasing; no + // intra-segment inversion can escape it. The simdgroups vote with simd_or + // over the conjunction, and the votes fold threadgroup-wide through shared + // memory; the barrier below orders both loops' results before the fold. On + // any violation the kernel retracts the descriptor count below so the tile + // kernel early-returns and the host re-routes to the order-agnostic legacy + // path. + bool boundary_ok = true; + if (lower > 0) { + boundary_ok = boundary_ok && indices[lower - 1] < lid; + } + if (lower < M) { + boundary_ok = boundary_ok && indices[lower] >= lid; + } + bool adjacent_ok = true; + for (int i = int(lid) + 1; i < M; i += int(expert_count)) { + adjacent_ok = adjacent_ok && indices[i - 1] <= indices[i]; + } + const uint violation_vote = simd_or((boundary_ok && adjacent_ok) ? 0u : 1u); + if (simd_lid == 0) { + violation_votes[simd_gid] = violation_vote; + } + threadgroup_barrier(mem_flags::mem_threadgroup); + + // Every thread folds the per-simdgroup votes uniformly. + bool sorted_violation = false; + for (uint group = 0; group < simdgroup_count; ++group) { + sorted_violation = sorted_violation || violation_votes[group] != 0u; + } + // count[1] is the violation observability slot; count[0] is the descriptor + // count. + if (lid == 0) { + count[1] = sorted_violation ? 1u : 0u; + } + + const uint segment_rows = segment_starts[lid + 1] - segment_starts[lid]; + inclusive_tile_offsets[lid] = (segment_rows + BM - 1) / BM; + threadgroup_barrier(mem_flags::mem_threadgroup); + + // log2(NE) uniform Hillis-Steele strides form an inclusive scan over the + // experts. The read barrier precedes each in-place update and the write + // barrier makes that stride visible to the next one. + for (uint stride = 1; stride < expert_count; stride <<= 1) { + const uint addend = + lid >= stride ? inclusive_tile_offsets[lid - stride] : 0; + threadgroup_barrier(mem_flags::mem_threadgroup); + if (lid >= stride) { + inclusive_tile_offsets[lid] += addend; + } + threadgroup_barrier(mem_flags::mem_threadgroup); + } + + const uint descriptor_count = inclusive_tile_offsets[expert_count - 1]; + if (lid == expert_count - 1) { + // A retracted count keeps the tile kernel's capacity check memory-safe + // (every threadgroup early-returns) and unambiguously signals the host: + // the assignment-count route gate guarantees M is 4096/8192/16384, so a + // valid build always emits at least one tile. + count[0] = sorted_violation ? 0u : descriptor_count; + } + + // Every thread emits a strided share of the bounded descriptor array. + // upper_bound over the inclusive offsets maps each slot back to its expert. + for (uint slot = lid; slot < descriptor_count; slot += expert_count) { + uint expert_lower = 0; + uint expert_upper = expert_count; + while (expert_lower < expert_upper) { + const uint midpoint = expert_lower + (expert_upper - expert_lower) / 2; + if (inclusive_tile_offsets[midpoint] <= slot) { + expert_lower = midpoint + 1; + } else { + expert_upper = midpoint; + } + } + const uint expert = expert_lower; + const uint expert_tile_begin = + expert == 0 ? 0 : inclusive_tile_offsets[expert - 1]; + const uint row = segment_starts[expert] + (slot - expert_tile_begin) * BM; + const uint row_count = min(BM, segment_starts[expert + 1] - row); + descriptors[slot] = uint4(row, row_count, expert, 0); + } +} + +template < + typename T, + const int group_size, + const int bits, + const bool aligned_N, + const int BM = 32, + const int BK = 32, + const int BN = 32> +[[kernel]] void affine_gather_qmm_gemma4_expert_tiles( + const device T* x [[buffer(0)]], + const device uint32_t* w [[buffer(1)]], + const device T* scales [[buffer(2)]], + const device T* biases [[buffer(3)]], + const device uint4* descriptors [[buffer(4)]], + const device uint* count [[buffer(5)]], + device T* y [[buffer(6)]], + const constant int& K [[buffer(7)]], + const constant int& N [[buffer(8)]], + uint3 tid [[threadgroup_position_in_grid]], + uint lid [[thread_index_in_threadgroup]], + uint simd_gid [[simdgroup_index_in_threadgroup]], + uint simd_lid [[thread_index_in_simdgroup]]) { + static_assert(BM == 32, "Gemma 4 expert tiles require BM=32"); + static_assert(BK == 32, "Gemma 4 expert tiles require BK=32"); + static_assert(BN == 32, "Gemma 4 expert tiles require BN=32"); + + // tid.y is uniform across the threadgroup. Empty capacity slots return + // before any threadgroup allocation is consumed by a microkernel barrier. + const uint descriptor_count = count[0]; + if (tid.y >= descriptor_count) { + return; + } + + constexpr int pack_factor = get_pack_factor<bits, 8>(); + constexpr int bytes_per_pack = get_bytes_per_pack<bits>(); + constexpr int BK_padded = BK + 16 / sizeof(T); + threadgroup T Xs[BM * BK_padded]; + threadgroup T Ws[BN * BK_padded]; + + const uint4 descriptor = descriptors[tid.y]; + const size_t row_start = size_t(descriptor.x); + const int row_count = int(descriptor.y); + const size_t expert = size_t(descriptor.z); + + const int K_w = K * bytes_per_pack / pack_factor; + const int K_g = K / group_size; + const size_t expert_w_stride = size_t(N) * size_t(K_w); + const size_t expert_sb_stride = size_t(N) * size_t(K_g); + + x += row_start * size_t(K); + y += row_start * size_t(N); + const device uint8_t* expert_w = + reinterpret_cast<const device uint8_t*>(w) + expert * expert_w_stride; + scales += expert * expert_sb_stride; + biases += expert * expert_sb_stride; + + const uint3 local_tid = uint3(tid.x, 0, 0); + if (row_count <= 16) { + qmm_t_expert_impl<T, group_size, bits, aligned_N, 16, BK, BN>( + reinterpret_cast<const device uint32_t*>(expert_w), + scales, + biases, + x, + y, + Xs, + Ws, + K, + N, + row_count, + K, + local_tid, + lid, + simd_gid, + simd_lid); + } else { + qmm_t_expert_impl<T, group_size, bits, aligned_N, BM, BK, BN>( + reinterpret_cast<const device uint32_t*>(expert_w), + scales, + biases, + x, + y, + Xs, + Ws, + K, + N, + row_count, + K, + local_tid, + lid, + simd_gid, + simd_lid); + } }   template <
diff --git ml-explore/mlx/mlx/backend/metal/kernels/quantized.metal Layr-Labs/mlx/mlx/backend/metal/kernels/quantized.metal index 069482cbaf0b67c24411afd3f80d7498a9a4d301..75e788ea2f6e39b5c28d8ef482212c2ca9137282 100644 --- ml-explore/mlx/mlx/backend/metal/kernels/quantized.metal +++ Layr-Labs/mlx/mlx/backend/metal/kernels/quantized.metal @@ -179,4 +179,31 @@ instantiate_quantized_groups(5) \ instantiate_quantized_groups(6) \ instantiate_quantized_groups(8)   -instantiate_quantized_all() // clang-format on +instantiate_quantized_all() + +instantiate_kernel( + "affine_gather_qmm_gemma4_expert_tiles_bfloat16_t_gs_64_b_4_alN_true_bm_32_bn_32_bk_32", + affine_gather_qmm_gemma4_expert_tiles, + bfloat16_t, + 64, + 4, + true, + 32, + 32, + 32) + +// Sorted expert-tile descriptor builders. The E=128 instantiation keeps the +// historical Gemma 4 host name; E=256 serves Qwen 3.5/3.6 MoE. The tile +// kernel instantiation above is expert-count agnostic (K/N are runtime +// arguments) and is shared by both routes. +instantiate_kernel( + "build_gemma4_sorted_expert_tiles_bm32", + build_sorted_expert_tiles_bm32, + 128) + +instantiate_kernel( + "build_sorted_expert_tiles_bm32_e256", + build_sorted_expert_tiles_bm32, + 256) + + // clang-format on
diff --git ml-explore/mlx/tests/gpu_tests.cpp Layr-Labs/mlx/tests/gpu_tests.cpp index 8bef07a616f6b72a6d4abbe2b71238763c690984..18018a463273b0b0106a28f0dd2f174ab8c44710 100644 --- ml-explore/mlx/tests/gpu_tests.cpp +++ Layr-Labs/mlx/tests/gpu_tests.cpp @@ -6,6 +6,7 @@ #include <cmath> #include <future>   #include "doctest/doctest.h" +#include "mlx/backend/common/gemma4_expert_qmm.h" #include "mlx/mlx.h"   using namespace mlx::core; @@ -712,3 +713,396 @@ } } CHECK(worst <= 1e-5); } + +TEST_CASE("test Gemma 4 expert QMM pure route table") { + using metal::classify_gemma4_expert_qmm; + using metal::Gemma4ExpertQMMRoute; + using metal::Gemma4ExpertQMMRouteInput; + + auto gate_up = [](int assignments) { + Gemma4ExpertQMMRouteInput input; + input.requested = true; + input.aot_available = true; + input.outer_route = true; + input.affine = true; + input.transpose = true; + input.has_bias = true; + input.indices_uint32 = true; + input.indices_contiguous = true; + input.x_bfloat16 = true; + input.x_contiguous = true; + input.w_uint32 = true; + input.w_contiguous = true; + input.scales_bfloat16 = true; + input.scales_contiguous = true; + input.biases_bfloat16 = true; + input.biases_contiguous = true; + input.group_size = 64; + input.bits = 4; + input.expert_count = 128; + input.assignments = assignments; + input.index_count = assignments; + input.k = 2816; + input.n = 1408; + input.x_rank = 3; + input.x_dim0 = assignments; + input.x_dim1 = 1; + input.x_dim2 = 2816; + input.w_rank = 3; + input.w_dim0 = 128; + input.w_dim1 = 1408; + input.w_dim2 = 352; + input.scales_rank = 3; + input.scales_dim0 = 128; + input.scales_dim1 = 1408; + input.scales_dim2 = 44; + input.biases_rank = 3; + input.biases_dim0 = 128; + input.biases_dim1 = 1408; + input.biases_dim2 = 44; + return input; + }; + auto down = [&gate_up](int assignments) { + auto input = gate_up(assignments); + input.k = 704; + input.n = 2816; + input.x_dim2 = 704; + input.w_dim1 = 2816; + input.w_dim2 = 88; + input.scales_dim1 = 2816; + input.scales_dim2 = 11; + input.biases_dim1 = 2816; + input.biases_dim2 = 11; + return input; + }; + + for (int assignments : {4096, 8192, 16384}) { + CHECK( + classify_gemma4_expert_qmm(gate_up(assignments)) == + Gemma4ExpertQMMRoute::hit); + CHECK( + classify_gemma4_expert_qmm(down(assignments)) == + Gemma4ExpertQMMRoute::hit); + } + + auto exact = gate_up(4096); + auto check_miss = [&exact](auto mutate, Gemma4ExpertQMMRoute expected) { + auto input = exact; + mutate(input); + CHECK(classify_gemma4_expert_qmm(input) == expected); + }; + check_miss( + [](auto& x) { x.requested = false; }, + Gemma4ExpertQMMRoute::not_requested); + check_miss( + [](auto& x) { x.nax_available = true; }, + Gemma4ExpertQMMRoute::fallback_nax); + check_miss( + [](auto& x) { x.outer_route = false; }, + Gemma4ExpertQMMRoute::fallback_outer_route); + check_miss( + [](auto& x) { x.affine = false; }, + Gemma4ExpertQMMRoute::fallback_quantization); + check_miss( + [](auto& x) { x.transpose = false; }, + Gemma4ExpertQMMRoute::fallback_quantization); + check_miss( + [](auto& x) { x.has_bias = false; }, + Gemma4ExpertQMMRoute::fallback_quantization); + check_miss( + [](auto& x) { x.group_size = 32; }, + Gemma4ExpertQMMRoute::fallback_quantization); + check_miss( + [](auto& x) { x.bits = 8; }, Gemma4ExpertQMMRoute::fallback_quantization); + check_miss( + [](auto& x) { x.indices_uint32 = false; }, + Gemma4ExpertQMMRoute::fallback_quantization); + check_miss( + [](auto& x) { x.indices_contiguous = false; }, + Gemma4ExpertQMMRoute::fallback_quantization); + check_miss( + [](auto& x) { x.x_bfloat16 = false; }, + Gemma4ExpertQMMRoute::fallback_quantization); + check_miss( + [](auto& x) { x.x_contiguous = false; }, + Gemma4ExpertQMMRoute::fallback_quantization); + check_miss( + [](auto& x) { x.w_uint32 = false; }, + Gemma4ExpertQMMRoute::fallback_quantization); + check_miss( + [](auto& x) { x.w_contiguous = false; }, + Gemma4ExpertQMMRoute::fallback_quantization); + check_miss( + [](auto& x) { x.scales_bfloat16 = false; }, + Gemma4ExpertQMMRoute::fallback_quantization); + check_miss( + [](auto& x) { x.scales_contiguous = false; }, + Gemma4ExpertQMMRoute::fallback_quantization); + check_miss( + [](auto& x) { x.biases_bfloat16 = false; }, + Gemma4ExpertQMMRoute::fallback_quantization); + check_miss( + [](auto& x) { x.biases_contiguous = false; }, + Gemma4ExpertQMMRoute::fallback_quantization); + check_miss( + [](auto& x) { x.expert_count = 127; }, + Gemma4ExpertQMMRoute::fallback_topology); + check_miss( + [](auto& x) { x.x_rank = 4; }, Gemma4ExpertQMMRoute::fallback_topology); + check_miss( + [](auto& x) { x.w_rank = 2; }, Gemma4ExpertQMMRoute::fallback_topology); + check_miss( + [](auto& x) { x.scales_rank = 2; }, + Gemma4ExpertQMMRoute::fallback_topology); + check_miss( + [](auto& x) { x.biases_rank = 2; }, + Gemma4ExpertQMMRoute::fallback_topology); + check_miss( + [](auto& x) { x.index_count -= 1; }, + Gemma4ExpertQMMRoute::fallback_topology); + for (int assignments : {8, 16, 32, 4095, 4097}) { + check_miss( + [assignments](auto& x) { + x.assignments = assignments; + x.index_count = assignments; + x.x_dim0 = assignments; + }, + Gemma4ExpertQMMRoute::fallback_assignment_count); + } + check_miss( + [](auto& x) { x.w_dim2 = 176; }, Gemma4ExpertQMMRoute::fallback_geometry); + check_miss( + [](auto& x) { x.w_dim1 += 1; }, Gemma4ExpertQMMRoute::fallback_geometry); + check_miss( + [](auto& x) { + x.k += 32; + x.x_dim2 = x.k; + }, + Gemma4ExpertQMMRoute::fallback_geometry); + check_miss( + [](auto& x) { x.n -= 32; }, Gemma4ExpertQMMRoute::fallback_geometry); + check_miss( + [](auto& x) { x.aot_available = false; }, + Gemma4ExpertQMMRoute::fallback_metallib_unavailable); + + auto nax_without_aot = exact; + nax_without_aot.nax_available = true; + nax_without_aot.aot_available = false; + CHECK( + classify_gemma4_expert_qmm(nax_without_aot) == + Gemma4ExpertQMMRoute::fallback_nax); +} + +TEST_CASE("test Qwen 3.6 expert QMM pure route table") { + using metal::classify_gemma4_expert_qmm; + using metal::Gemma4ExpertQMMRoute; + using metal::Gemma4ExpertQMMRouteInput; + + // Base input: Qwen 3.5/3.6 35B-A3B expert projection at W4/g64, + // parametrized by whole-projection [E=256, n, k]. + auto qwen = [](int assignments, int k, int n) { + Gemma4ExpertQMMRouteInput input; + input.requested = true; + input.aot_available = true; + input.outer_route = true; + input.affine = true; + input.transpose = true; + input.has_bias = true; + input.indices_uint32 = true; + input.indices_contiguous = true; + input.x_bfloat16 = true; + input.x_contiguous = true; + input.w_uint32 = true; + input.w_contiguous = true; + input.scales_bfloat16 = true; + input.scales_contiguous = true; + input.biases_bfloat16 = true; + input.biases_contiguous = true; + input.group_size = 64; + input.bits = 4; + input.expert_count = 256; + input.assignments = assignments; + input.index_count = assignments; + input.k = k; + input.n = n; + input.x_rank = 3; + input.x_dim0 = assignments; + input.x_dim1 = 1; + input.x_dim2 = k; + input.w_rank = 3; + input.w_dim0 = 256; + input.w_dim1 = n; + input.w_dim2 = k / 8; + input.scales_rank = 3; + input.scales_dim0 = 256; + input.scales_dim1 = n; + input.scales_dim2 = k / 64; + input.biases_rank = 3; + input.biases_dim0 = 256; + input.biases_dim1 = n; + input.biases_dim2 = k / 64; + return input; + }; + + // Fused gate_up, split gate/up, and down projections hit at the chunked + // prefill assignment counts (T x top-8 for T in {512, 1024, 2048}). + for (int assignments : {4096, 8192, 16384}) { + CHECK( + classify_gemma4_expert_qmm(qwen(assignments, 2048, 1024)) == + Gemma4ExpertQMMRoute::hit); + CHECK( + classify_gemma4_expert_qmm(qwen(assignments, 2048, 512)) == + Gemma4ExpertQMMRoute::hit); + CHECK( + classify_gemma4_expert_qmm(qwen(assignments, 512, 2048)) == + Gemma4ExpertQMMRoute::hit); + } + + auto exact = qwen(4096, 2048, 1024); + auto check_miss = [&exact](auto mutate, Gemma4ExpertQMMRoute expected) { + auto input = exact; + mutate(input); + CHECK(classify_gemma4_expert_qmm(input) == expected); + }; + // Expert counts other than the two instantiated builders miss on topology. + check_miss( + [](auto& x) { + x.expert_count = 255; + x.w_dim0 = 255; + x.scales_dim0 = 255; + x.biases_dim0 = 255; + }, + Gemma4ExpertQMMRoute::fallback_topology); + // E=256 with Gemma geometry (and vice versa) must miss on geometry: the + // shape table is tied to the expert count, never mixed. + check_miss( + [](auto& x) { + x.k = 2816; + x.n = 1408; + x.x_dim2 = 2816; + x.w_dim1 = 1408; + x.w_dim2 = 352; + x.scales_dim1 = 1408; + x.scales_dim2 = 44; + x.biases_dim1 = 1408; + x.biases_dim2 = 44; + }, + Gemma4ExpertQMMRoute::fallback_geometry); + check_miss( + [](auto& x) { x.w_dim2 = 128; }, Gemma4ExpertQMMRoute::fallback_geometry); + check_miss( + [](auto& x) { x.n -= 32; }, Gemma4ExpertQMMRoute::fallback_geometry); + // T=128 chunks (1024 assignments) intentionally stay on the legacy path. + for (int assignments : {8, 1024, 4095, 4097}) { + check_miss( + [assignments](auto& x) { + x.assignments = assignments; + x.index_count = assignments; + x.x_dim0 = assignments; + }, + Gemma4ExpertQMMRoute::fallback_assignment_count); + } + check_miss( + [](auto& x) { x.bits = 8; }, Gemma4ExpertQMMRoute::fallback_quantization); + check_miss( + [](auto& x) { x.aot_available = false; }, + Gemma4ExpertQMMRoute::fallback_metallib_unavailable); + + // The Gemma table must also reject Qwen geometry under E=128. + auto gemma_with_qwen_geometry = exact; + gemma_with_qwen_geometry.expert_count = 128; + gemma_with_qwen_geometry.w_dim0 = 128; + gemma_with_qwen_geometry.scales_dim0 = 128; + gemma_with_qwen_geometry.biases_dim0 = 128; + CHECK( + classify_gemma4_expert_qmm(gemma_with_qwen_geometry) == + Gemma4ExpertQMMRoute::fallback_geometry); +} + +TEST_CASE("test Gemma 4 expert QMM counter invariant") { + metal::Gemma4ExpertQMMCounters counters; + using Route = metal::Gemma4ExpertQMMRoute; + counters.record(Route::not_requested); + counters.record(Route::hit); + counters.record(Route::fallback_nax); + counters.record(Route::fallback_outer_route); + counters.record(Route::fallback_quantization); + counters.record(Route::fallback_topology); + counters.record(Route::fallback_assignment_count); + counters.record(Route::fallback_geometry); + counters.record(Route::fallback_metallib_unavailable); + counters.record(Route::fallback_sortedness_retracted); + + auto snapshot = counters.snapshot(); + CHECK(snapshot.hits == 1); + CHECK(snapshot.fallback_nax == 1); + CHECK(snapshot.fallback_outer_route == 1); + CHECK(snapshot.fallback_quantization == 1); + CHECK(snapshot.fallback_topology == 1); + CHECK(snapshot.fallback_assignment_count == 1); + CHECK(snapshot.fallback_geometry == 1); + CHECK(snapshot.fallback_metallib_unavailable == 1); + CHECK(snapshot.fallback_sortedness_retracted == 1); + CHECK(snapshot.attempts() == 9); + + counters.reset(); + snapshot = counters.snapshot(); + CHECK(snapshot.attempts() == 0); + CHECK( + snapshot.attempts() == + snapshot.hits + snapshot.fallback_nax + snapshot.fallback_outer_route + + snapshot.fallback_quantization + snapshot.fallback_topology + + snapshot.fallback_assignment_count + snapshot.fallback_geometry + + snapshot.fallback_metallib_unavailable + + snapshot.fallback_sortedness_retracted); +} + +TEST_CASE("test Gemma 4 expert QMM arm disarm cycle") { + metal::Gemma4ExpertQMMCounters counters; + using Route = metal::Gemma4ExpertQMMRoute; + + // Counters start disarmed with an empty interval. + CHECK(!counters.armed()); + + // Arm: the interval opens with zeroed counters. + counters.clear_and_arm(); + CHECK(counters.armed()); + CHECK(counters.snapshot().attempts() == 0); + + // Record across the measured interval, including the retract class the + // sortedness fail-safe attributes mis-sorted indices to. + counters.record(Route::hit); + counters.record(Route::fallback_sortedness_retracted); + counters.record(Route::fallback_metallib_unavailable); + + // Disarm snapshots the interval and reports the previous armed state. + auto interval = counters.snapshot_and_disarm(); + CHECK(interval.armed); + CHECK(!counters.armed()); + + // The attempts == hits + sum(fallback classes) invariant holds across the + // cycle, with the sortedness-retract class included in the sum. + CHECK(interval.attempts() == 3); + CHECK(interval.hits == 1); + CHECK(interval.fallback_sortedness_retracted == 1); + CHECK(interval.fallback_metallib_unavailable == 1); + CHECK( + interval.attempts() == + interval.hits + interval.fallback_nax + interval.fallback_outer_route + + interval.fallback_quantization + interval.fallback_topology + + interval.fallback_assignment_count + interval.fallback_geometry + + interval.fallback_metallib_unavailable + + interval.fallback_sortedness_retracted); + + // The snapshot stays readable while disarmed. + CHECK(!counters.snapshot().armed); + CHECK(counters.snapshot().attempts() == 3); + + // Re-arming clears the interval again, and disarming it reports armed. + counters.clear_and_arm(); + CHECK(counters.armed()); + auto reopened = counters.snapshot_and_disarm(); + CHECK(reopened.armed); + CHECK(reopened.attempts() == 0); + CHECK(!counters.armed()); +} \ No newline at end of file

GPT-OSS 20B gathered MXFP4 decode has K=2880, which misses the fast vector kernel’s 512-element alignment and fell to the general path. fp_gather_qmv_fast_tail runs five full blocks plus a 320-value tail read only by the participating lanes, preserving FP32 accumulation and the gathered indices/strides. It is enabled automatically only on the physical M4 Max architecture (applegpu_g16s) for the exact MXFP4/group-32⁄4-bit, E=32, K=2880, N=2880|5760, FP32/BF16 shape; MLX_GPTOSS_MXFP4_DECODE_FAST_TAIL=0 restores the original route and MLX_GPTOSS_MXFP4_PREFILL_TILE=m32n32k32 opts into a separately measured prefill tile (gptoss_mxfp4_policy.h). (PR #14) fp_quantized.metal also builds the two m32n32k32 prefill-tile kernels (float and bfloat16) ahead of time, so CMake builds with the default MLX_METAL_JIT=OFF can load them. (PR #24)

diff --git ml-explore/mlx/mlx/backend/metal/gptoss_mxfp4_policy.h Layr-Labs/mlx/mlx/backend/metal/gptoss_mxfp4_policy.h new file mode 100644 index 0000000000000000000000000000000000000000..f0eaf66349c6952584059b25c4e622c436b1413f --- /dev/null +++ Layr-Labs/mlx/mlx/backend/metal/gptoss_mxfp4_policy.h @@ -0,0 +1,24 @@ +// Copyright © 2026 Eigen Labs. +#pragma once + +#include <cstdlib> +#include <string_view> + +namespace mlx::core::metal { + +struct GPTOSSMXFP4PrefillTile { + int bm = 16; + int bn = 32; + int bk = 32; + int wm = 1; + int wn = 2; +}; + +inline GPTOSSMXFP4PrefillTile gptoss_mxfp4_prefill_tile(const char* option) { + const std::string_view value = option ? option : ""; + if (value == "m32n32k32") + return {32, 32, 32, 2, 2}; + return {}; +} + +} // namespace mlx::core::metal
diff --git ml-explore/mlx/mlx/backend/metal/kernels/fp_quantized.h Layr-Labs/mlx/mlx/backend/metal/kernels/fp_quantized.h index 6e77569f56ba9676dba950a5b4a41df5747fe7cc..5b4b04f3e1fa82c2bbfcefc3c4ed179ffce57e19 100644 --- ml-explore/mlx/mlx/backend/metal/kernels/fp_quantized.h +++ Layr-Labs/mlx/mlx/backend/metal/kernels/fp_quantized.h @@ -137,11 +137,12 @@ }   template <typename U, int bits> inline void dequantize(uint8_t w, U scale, threadgroup U* w_local) { + const float s = float(scale); if constexpr (bits == 4) { - w_local[0] = scale * Dequantize<4, U>{}(w); - w_local[1] = scale * Dequantize<4, U>{}(w >> 4); + w_local[0] = static_cast<U>(s * Dequantize<4, float>{}(w)); + w_local[1] = static_cast<U>(s * Dequantize<4, float>{}(w >> 4)); } else { - w_local[0] = scale * Dequantize<8, U>{}(w); + w_local[0] = static_cast<U>(s * Dequantize<8, float>{}(w)); } }   @@ -326,7 +327,8 @@ typename T, int group_size, int bits, bool has_global_scale = false, - int results_per_simdgroup = 4> + int results_per_simdgroup = 4, + bool allow_aligned_tail = false> METAL_FUNC void fp_qmv_fast_impl( const device uint32_t* w, const device uint8_t* scales, @@ -363,7 +365,10 @@ scales += out_row * in_vec_size_g + simd_lid / scale_step_per_thread; x += tid.x * in_vec_size + simd_lid * values_per_thread; y += tid.x * out_vec_size + out_row;   - for (int k = 0; k < in_vec_size; k += block_size) { + const int full_size = allow_aligned_tail + ? (in_vec_size / block_size) * block_size + : in_vec_size; + for (int k = 0; k < full_size; k += block_size) { load_vector<T, U, values_per_thread>(x, x_thread);   for (int row = 0; row < results_per_simdgroup; row++) { @@ -377,6 +382,20 @@ ws += block_size * bytes_per_pack / pack_factor; scales += block_size / group_size; x += block_size; + } + + if constexpr (allow_aligned_tail) { + // K is group-aligned: every active lane owns a complete packed vector. + const int tail_values = in_vec_size - full_size; + if (int(simd_lid) * values_per_thread < tail_values) { + load_vector<T, U, values_per_thread>(x, x_thread); + for (int row = 0; row < results_per_simdgroup; row++) { + const device uint8_t* wl = ws + row * in_vec_size_w; + const device uint8_t* sl = scales + row * in_vec_size_g; + U scale = dequantize_scale<U, group_size>(sl[0]); + result[row] += qdot<U, values_per_thread, bits>(wl, x_thread, scale); + } + } }   float inv_scale_enc = 1.0f; @@ -1231,7 +1250,8 @@ typename T, int group_size, int bits, bool batched, - bool has_global_scale = false> + bool has_global_scale = false, + int results_per_simdgroup = 4> [[kernel]] void fp_qmv( const device uint32_t* w, const device uint8_t* scales, @@ -1608,6 +1628,65 @@ simd_lid); }   template <typename T, int group_size, int bits, bool has_global_scale = false> +[[kernel]] void fp_gather_qmv_fast_tail( + const device uint32_t* w, + const device uint8_t* scales, + const device float* global_scale, + const device T* x, + const device uint32_t* lhs_indices, + const device uint32_t* rhs_indices, + device T* y, + const constant int& in_vec_size, + const constant int& out_vec_size, + const constant int& x_batch_ndims, + const constant int* x_shape, + const constant int64_t* x_strides, + const constant int& w_batch_ndims, + const constant int* w_shape, + const constant int64_t* w_strides, + const constant int64_t* s_strides, + const constant int& batch_ndims, + const constant int* batch_shape, + const constant int64_t* lhs_strides, + const constant int64_t* rhs_strides, + uint3 tid [[threadgroup_position_in_grid]], + uint simd_gid [[simdgroup_index_in_threadgroup]], + uint simd_lid [[thread_index_in_simdgroup]]) { + int M = x_shape[x_batch_ndims]; + adjust_matrix_offsets( + x, + w, + scales, + lhs_indices, + rhs_indices, + y, + out_vec_size * M, + batch_ndims, + batch_shape, + lhs_strides, + rhs_strides, + x_batch_ndims, + x_shape, + x_strides, + w_batch_ndims, + w_shape, + w_strides, + s_strides, + tid); + fp_qmv_fast_impl<T, group_size, bits, has_global_scale, 4, true>( + w, + scales, + global_scale, + x, + y, + in_vec_size, + out_vec_size, + tid, + simd_gid, + simd_lid); +} + +template <typename T, int group_size, int bits, bool has_global_scale = false> [[kernel]] void fp_gather_qmv( const device uint32_t* w, const device uint8_t* scales, @@ -2121,7 +2200,7 @@ }   float scale_dec_b; float w_thread = w[index]; - if (use_mx_scale) { + if constexpr (use_mx_scale) { scale_dec_b = simd_max(abs(w_thread)); } else { float w_max_l = simd_max(simd_lid < 16 ? abs(w_thread) : 0.0); @@ -2131,6 +2210,8 @@ } scale_dec_b /= bits == 4 ? F4E2M1_MAX : F8E4M3_MAX; if constexpr (has_global_scale) { scale_dec_b *= scale_enc; + } else if constexpr (use_mx_scale) { + scale_dec_b = mx_scale_round_up(scale_dec_b); }   using ScaleType = metal::conditional_t<use_mx_scale, fp8_e8m0, fp8_e4m3>; @@ -2213,7 +2294,7 @@ }   float scale_dec_b; float w_thread = w[index]; - if (use_mx_scale) { + if constexpr (use_mx_scale) { scale_dec_b = simd_max(abs(w_thread)); } else { float w_max_l = simd_max(simd_lid < 16 ? abs(w_thread) : 0.0); @@ -2223,6 +2304,8 @@ } scale_dec_b /= bits == 4 ? F4E2M1_MAX : F8E4M3_MAX; if constexpr (has_global_scale) { scale_dec_b *= scale_enc; + } else if constexpr (use_mx_scale) { + scale_dec_b = mx_scale_round_up(scale_dec_b); }   using ScaleType = metal::conditional_t<use_mx_scale, fp8_e8m0, fp8_e4m3>;
diff --git ml-explore/mlx/mlx/backend/metal/kernels/fp_quantized.metal Layr-Labs/mlx/mlx/backend/metal/kernels/fp_quantized.metal index d8d462288adf65b09230bb3e4584a90fad4e8d1a..ea12ae9afb573ecb3d3a9f71af285bcd92a949b9 100644 --- ml-explore/mlx/mlx/backend/metal/kernels/fp_quantized.metal +++ Layr-Labs/mlx/mlx/backend/metal/kernels/fp_quantized.metal @@ -222,3 +222,11 @@ instantiate_quantized_types(float) instantiate_quantized_types(bfloat16_t) instantiate_quantized_types(float16_t) // clang-format on + +// Exact GPT-OSS MXFP4 gathered-vector tail experiment. +instantiate_quantized(mxfp4, gather_qmv_fast_tail, float, 32, 4) +instantiate_quantized(mxfp4, gather_qmv_fast_tail, bfloat16_t, 32, 4) + +// GPT-OSS MXFP4 prefill tile m32n32k32, selected by MLX_GPTOSS_MXFP4_PREFILL_TILE. +instantiate_gather_qmm_rhs(fp_gather_qmm_rhs, gather_qmm_rhs_nt, float, 32, 32, 32, 2, 2, true, mxfp4, 32, 4) +instantiate_gather_qmm_rhs(fp_gather_qmm_rhs, gather_qmm_rhs_nt, bfloat16_t, 32, 32, 32, 2, 2, true, mxfp4, 32, 4)
diff --git ml-explore/mlx/mlx/backend/metal/quantized.cpp Layr-Labs/mlx/mlx/backend/metal/quantized.cpp index f659e16c9356a066154cc5511596bba2c28e444a..888389a2dd158ed095220d959f0d845819e300c3 100644 --- ml-explore/mlx/mlx/backend/metal/quantized.cpp +++ Layr-Labs/mlx/mlx/backend/metal/quantized.cpp @@ -5,6 +5,7 @@ #include "mlx/backend/common/broadcasting.h" #include "mlx/backend/common/compiled.h" #include "mlx/backend/gpu/copy.h" #include "mlx/backend/metal/device.h" +#include "mlx/backend/metal/gptoss_mxfp4_policy.h" #include "mlx/backend/metal/kernels.h" #include "mlx/backend/metal/reduce.h" #include "mlx/backend/metal/unary.h" @@ -14,6 +15,10 @@ #include "mlx/primitives.h" #include "mlx/utils.h"   namespace mlx::core { + +using metal::classify_gemma4_expert_qmm; +using metal::Gemma4ExpertQMMRoute; +using metal::Gemma4ExpertQMMRouteInput;   namespace {   @@ -249,21 +254,18 @@ auto w = ensure_row_contiguous(w_pre, d, s); if (dequantize) { auto scales = ensure_row_contiguous(inputs[1], d, s); - compute_encoder.set_input_array(w, 0); - compute_encoder.set_input_array(scales, 1); if (has_biases) { auto biases = ensure_row_contiguous(inputs[2], d, s); compute_encoder.set_input_array(biases, 2); } else if (has_global_scale) { compute_encoder.set_input_array(inputs[2], 2); } + compute_encoder.set_input_array(w, 0); + compute_encoder.set_input_array(scales, 1); compute_encoder.set_output_array(out, 3); } else { auto& scales = outputs[1]; scales.set_data(allocator::malloc(scales.nbytes())); - compute_encoder.set_input_array(w, 0); - compute_encoder.set_output_array(out, 1); - compute_encoder.set_output_array(scales, 2); if (has_biases) { auto& biases = outputs[2]; biases.set_data(allocator::malloc(biases.nbytes())); @@ -271,6 +273,9 @@ compute_encoder.set_output_array(biases, 3); } else if (has_global_scale) { compute_encoder.set_input_array(inputs[1], 3); } + compute_encoder.set_input_array(w, 0); + compute_encoder.set_output_array(out, 1); + compute_encoder.set_output_array(scales, 2); }   auto type_string = dequantize ? get_type_string(out.dtype()) @@ -649,7 +654,7 @@ constexpr int num_simdgroups = 2; constexpr int bk = 32; int bn = std::min(group_size, 32) * num_simdgroups; MTL::Size group_dims = MTL::Size(bk, num_simdgroups, 1); - MTL::Size grid_dims = MTL::Size(M, N / bn, B); + MTL::Size grid_dims = MTL::Size(M, (N + bn - 1) / bn, B);   auto x_shape = x.shape(); auto x_strides = x.strides(); @@ -830,7 +835,8 @@ int B = out.size() / M / N;   int wm = 2; int wn = 2; - int bm = 64; + // Use smaller bm when one block covers all of M. + int bm = (M <= 32) ? 32 : 64; int bn = 64; int bk = 64; MTL::Size group_dims(32, wn, wm); @@ -971,7 +977,7 @@ if (transpose) { kernel = get_qmm_nax_kernel_wrapped( d, kname, - "gather_qmm_t_nax_", + "gather_qmm_t_nax", mode, type_string, group_size, @@ -986,7 +992,7 @@ } else { kernel = get_qmm_nax_kernel_wrapped( d, kname, - "gather_qmm_n_nax_", + "gather_qmm_n_nax", mode, type_string, group_size, @@ -1334,9 +1340,22 @@ std::string kname; kname.reserve(64); std::string type_string = get_type_string(x.dtype()); bool fast = N % bn == 0 && K % qmv_fast_k_alignment(bits) == 0; + bool fast_tail = false; + if (mode == "mxfp4" && group_size == 32 && bits == 4 && !global_scale && + w.ndim() == 3 && w.shape(0) == 32 && K == 2880 && + (N == 2880 || N == 5760) && + (x.dtype() == float32 || x.dtype() == bfloat16)) { + const char* option = std::getenv("MLX_GPTOSS_MXFP4_DECODE_FAST_TAIL"); + const std::string_view physical_arch = + d.mtl_device()->architecture()->name()->utf8String(); + fast_tail = option ? std::string_view(option) == "1" + : physical_arch == "applegpu_g16s"; + } concatenate( kname, - mode + (fast ? "_gather_qmv_fast_" : "_gather_qmv_"), + mode + + (fast_tail ? "_gather_qmv_fast_tail_" + : (fast ? "_gather_qmv_fast_" : "_gather_qmv_")), type_string, "_gs_", group_size, @@ -1347,7 +1366,8 @@ auto kernel = get_quantized_kernel_wrapped( d, kname, - (fast ? "gather_qmv_fast" : "gather_qmv"), + (fast_tail ? "gather_qmv_fast_tail" + : (fast ? "gather_qmv_fast" : "gather_qmv")), mode, type_string, group_size, @@ -1582,6 +1602,114 @@ compute_encoder.dispatch_threadgroups(grid_dims, group_dims); }   +Gemma4ExpertQMMRoute try_gemma4_expert_qmm( + const array& x, + const array& w, + const array& scales, + const array& biases, + const array& indices, + array& out, + int M, + int N, + int K, + metal::Device& d, + const Stream& s) { + // The classifier admits exactly E=128 (Gemma 4) and E=256 (Qwen 3.5/3.6 + // MoE); w is rank-3 [E, N_w, K_w] by the same gate. The tile kernel is + // expert-count agnostic — only the descriptor builder (one thread per + // expert) is instantiated per expert count. + const int expert_count = w.shape(0); + const char* descriptor_kernel_name = expert_count == 256 + ? "build_sorted_expert_tiles_bm32_e256" + : "build_gemma4_sorted_expert_tiles_bm32"; + constexpr const char* tile_kernel_name = + "affine_gather_qmm_gemma4_expert_tiles_bfloat16_t_gs_64_b_4_" + "alN_true_bm_32_bn_32_bk_32"; + + MTL::ComputePipelineState* descriptor_kernel = nullptr; + MTL::ComputePipelineState* tile_kernel = nullptr; + try { + descriptor_kernel = d.get_kernel(descriptor_kernel_name); + tile_kernel = d.get_kernel(tile_kernel_name); + } catch (...) { + return Gemma4ExpertQMMRoute::fallback_metallib_unavailable; + } + + constexpr int bm = 32; + constexpr int bn = 32; + constexpr int wm = 2; + constexpr int wn = 2; + // Upper bound on descriptors: sum over experts of ceil(rows_e / bm) with + // sum(rows_e) == M. With k experts carrying a nonzero remainder the total + // is at most (M - k) / bm + k <= M / bm + E - 1 for the reachable M/E. + const int max_tile_count = (M + bm - 1) / bm + expert_count - 1; + + array descriptors({max_tile_count, 4}, uint32, nullptr, {}); + descriptors.set_data(allocator::malloc(descriptors.nbytes())); + // count[0] is the descriptor count; count[1] is the builder's sortedness + // violation observability slot. + array tile_count({2}, uint32, nullptr, {}); + tile_count.set_data(allocator::malloc(tile_count.nbytes())); + + auto& compute_encoder = metal::get_command_encoder(s); + compute_encoder.add_temporary(descriptors); + compute_encoder.add_temporary(tile_count); + compute_encoder.set_compute_pipeline_state(descriptor_kernel); + compute_encoder.set_input_array(indices, 0); + compute_encoder.set_output_array(descriptors, 1); + compute_encoder.set_output_array(tile_count, 2); + compute_encoder.set_bytes(M, 3); + compute_encoder.dispatch_threads( + MTL::Size(expert_count, 1, 1), MTL::Size(expert_count, 1, 1)); + + // The descriptor builder fail-safes to count[0] == 0 when the purportedly + // sorted indices violate the non-decreasing invariant its binary search + // relies on. A zero count is unambiguous here: the selector's assignment + // gate guarantees M is one of 4096/8192/16384, so a valid build always + // emits at least one tile. Drain the stream and re-route retracted calls + // to the order-agnostic legacy path rather than running the tile kernel. + // + // MLX_GATHER_QMM_EXPERT_SLICES=trust skips this drain entirely: the tile + // grid below is already over-dispatched (max_tile_count threadgroups; + // the kernel early-returns slots >= count[0]), so the readback exists + // ONLY to observe a retracted build. Under trust the caller asserts the + // sorted contract is machine-guaranteed (the Swift SwitchGLU prefill + // path sorts on-device just before this call); a genuine violation then + // yields undefined output for this matmul instead of the legacy result. + // Measured cost of the drain: ~120 stream drains per 512-token prefill + // chunk (3 gathers x 40 MoE layers) — it cancels the tile kernel's + // 12-23%/unit win end-to-end. A device-side legacy fallback (dispatch- + // diet item 1.3) would make trust the only behavior. + if (!d.gemma4_expert_qmm_trust_sorted()) { + compute_encoder.synchronize(); + const uint32_t* counts = tile_count.data<uint32_t>(); + if (counts[0] == 0) { + // count[1] flags a detected sortedness violation; attribute the + // retract to its own bucket. Any other unusable build keeps the + // metallib bucket. + return counts[1] == 1u + ? Gemma4ExpertQMMRoute::fallback_sortedness_retracted + : Gemma4ExpertQMMRoute::fallback_metallib_unavailable; + } + } + + compute_encoder.set_compute_pipeline_state(tile_kernel); + int c = 0; + compute_encoder.set_input_array(x, c++); + compute_encoder.set_input_array(w, c++); + compute_encoder.set_input_array(scales, c++); + compute_encoder.set_input_array(biases, c++); + compute_encoder.set_input_array(descriptors, c++); + compute_encoder.set_input_array(tile_count, c++); + compute_encoder.set_output_array(out, c++); + compute_encoder.set_bytes(K, c++); + compute_encoder.set_bytes(N, c++); + + compute_encoder.dispatch_threadgroups( + MTL::Size((N + bn - 1) / bn, max_tile_count, 1), MTL::Size(32, wn, wm)); + return Gemma4ExpertQMMRoute::hit; +} + void gather_qmm_rhs( const array& x_, const array& w_, @@ -1598,8 +1726,17 @@ int K, metal::Device& d, const Stream& s, const std::string mode) { - if (metal::is_nax_available() && transpose && - (env::enable_tf32() || x_.dtype() != float32)) { + const bool nax_takes_call = metal::is_nax_available() && transpose && + (env::enable_tf32() || x_.dtype() != float32); + if (nax_takes_call) { + if (d.gemma4_expert_qmm_diagnostics_armed() && + d.gemma4_expert_qmm_requested()) { + Gemma4ExpertQMMRouteInput route_input; + route_input.requested = true; + route_input.outer_route = true; + route_input.nax_available = true; + d.record_armed_gemma4_expert_qmm(classify_gemma4_expert_qmm(route_input)); + } return gather_qmm_rhs_nax( /* const array& x_ = */ x_, /* const array& w_ = */ w_, @@ -1647,9 +1784,98 @@ if (biases_) { biases = ensure_row_contiguous(*biases_, d, s); }   + if (d.gemma4_expert_qmm_requested()) { + auto shape_dim = [](const array& value, int axis) { + return value.ndim() > axis ? value.shape(axis) : 0; + }; + Gemma4ExpertQMMRouteInput route_input; + route_input.requested = true; + route_input.aot_available = d.gemma4_expert_qmm_aot_available(); + route_input.nax_available = false; + route_input.outer_route = true; + route_input.affine = mode == "affine"; + route_input.transpose = transpose; + route_input.has_bias = biases_.has_value(); + route_input.indices_uint32 = indices.dtype() == uint32; + route_input.indices_contiguous = indices.flags().row_contiguous; + route_input.x_bfloat16 = x.dtype() == bfloat16; + route_input.x_contiguous = x.flags().row_contiguous; + route_input.w_uint32 = w.dtype() == uint32; + route_input.w_contiguous = w.flags().row_contiguous; + route_input.scales_bfloat16 = scales.dtype() == bfloat16; + route_input.scales_contiguous = scales.flags().row_contiguous; + // Classification reads the raw bias tensor; normalization is spent only + // in the winning-route branch below. The legacy block at the end of this + // function retains its original normalization point and ordering. + route_input.biases_bfloat16 = biases_ && biases_->dtype() == bfloat16; + route_input.biases_contiguous = biases_ && biases_->flags().row_contiguous; + route_input.group_size = group_size; + route_input.bits = bits; + route_input.expert_count = w.size() / w.shape(-1) / w.shape(-2); + route_input.assignments = M; + route_input.index_count = indices.size(); + route_input.k = K; + route_input.n = N; + route_input.x_rank = x.ndim(); + route_input.x_dim0 = shape_dim(x, 0); + route_input.x_dim1 = shape_dim(x, 1); + route_input.x_dim2 = shape_dim(x, 2); + route_input.w_rank = w.ndim(); + route_input.w_dim0 = shape_dim(w, 0); + route_input.w_dim1 = shape_dim(w, 1); + route_input.w_dim2 = shape_dim(w, 2); + route_input.scales_rank = scales.ndim(); + route_input.scales_dim0 = shape_dim(scales, 0); + route_input.scales_dim1 = shape_dim(scales, 1); + route_input.scales_dim2 = shape_dim(scales, 2); + if (biases_) { + route_input.biases_rank = biases_->ndim(); + route_input.biases_dim0 = shape_dim(*biases_, 0); + route_input.biases_dim1 = shape_dim(*biases_, 1); + route_input.biases_dim2 = shape_dim(*biases_, 2); + } + + auto route = classify_gemma4_expert_qmm(route_input); + if (route == Gemma4ExpertQMMRoute::hit) { + // A hit requires has_bias, so dereferencing biases_ is safe. + array expert_biases = ensure_row_contiguous(*biases_, d, s); + route = try_gemma4_expert_qmm( + x, w, scales, expert_biases, indices, out, M, N, K, d, s); + if (route == Gemma4ExpertQMMRoute::hit) { + if (d.gemma4_expert_qmm_diagnostics_armed()) { + d.record_armed_gemma4_expert_qmm(route); + } + return; + } + // A retracted build attributes to its own counter bucket + // (fallback_sortedness_retracted); missing AOT kernels keep the + // metallib bucket. In both cases the legacy route below produces the + // correct result. + } + if (d.gemma4_expert_qmm_diagnostics_armed()) { + d.record_armed_gemma4_expert_qmm(route); + } + } + + // Legacy gather path. Its normalization and dispatch order intentionally + // remain the global behavior for every non-exact call. + // TODO: Tune the block sizes int bm = 16, bn = 32, bk = 32; int wm = 1, wn = 2; + + if (mode == "mxfp4" && transpose && group_size == 32 && bits == 4 && + w.ndim() == 3 && w.shape(0) == 32 && K == 2880 && + (N == 2880 || N == 5760) && M >= 64 && + (x.dtype() == float32 || x.dtype() == bfloat16)) { + const auto tile = metal::gptoss_mxfp4_prefill_tile( + std::getenv("MLX_GPTOSS_MXFP4_PREFILL_TILE")); + bm = tile.bm; + bn = tile.bn; + bk = tile.bk; + wm = tile.wm; + wn = tile.wn; + }   const bool align_M = (M % bm) == 0; const bool align_N = (N % bn) == 0; @@ -1873,6 +2099,15 @@ x, w, scales, biases, out, group_size_, bits_, M, N, K, d, s, mode); return; }   +// Single source for the sorted-RHS expert route gate. Both the diagnostics +// record below and the dispatch decision evaluate this one predicate so a +// future tuning change cannot desynchronize them. +// TODO: Tune 16 and 4 here a bit better. +static constexpr bool +takes_sorted_rhs_route(int M, int B, int E, bool right_sorted) { + return M == 1 && B >= 16 && right_sorted && B / E >= 4; +} + void GatherQMM::eval_gpu(const std::vector<array>& inputs, array& out) { auto& s = stream(); auto& d = metal::device(s.device); @@ -1897,11 +2132,18 @@ int E = w.size() / w.shape(-1) / w.shape(-2); int vector_limit = transpose_ ? get_qmv_batch_limit(K, N, d) : 4; auto mode = quantization_mode_to_string(mode_);   + if (d.gemma4_expert_qmm_diagnostics_armed() && + d.gemma4_expert_qmm_requested() && + !takes_sorted_rhs_route(M, B, E, right_sorted_)) { + Gemma4ExpertQMMRouteInput route_input; + route_input.requested = true; + route_input.outer_route = false; + d.record_armed_gemma4_expert_qmm(classify_gemma4_expert_qmm(route_input)); + } + // We are walking x in order and w is also in order so we can batch up the // matmuls and reuse reading x and w. - // - // TODO: Tune 16 and 4 here a bit better. - if (M == 1 && B >= 16 && right_sorted_ == true && B / E >= 4) { + if (takes_sorted_rhs_route(M, B, E, right_sorted_)) { gather_qmm_rhs( x, w,
diff --git ml-explore/mlx/tests/CMakeLists.txt Layr-Labs/mlx/tests/CMakeLists.txt index 859215acb1e912ab43e5cbeaaa9f7d34bd6d69f4..c80d20c1e3ab110bfb13bc02c76f5ead37725100 100644 --- ml-explore/mlx/tests/CMakeLists.txt +++ Layr-Labs/mlx/tests/CMakeLists.txt @@ -11,7 +11,7 @@ set(METAL_TEST_SOURCES gpu_tests.cpp) endif()   if(MLX_BUILD_METAL) - list(APPEND METAL_TEST_SOURCES residency_tests.cpp) + list(APPEND METAL_TEST_SOURCES residency_tests.cpp gptoss_mxfp4_tests.cpp) endif()   include(${doctest_SOURCE_DIR}/scripts/cmake/doctest.cmake)
diff --git ml-explore/mlx/tests/gptoss_mxfp4_tests.cpp Layr-Labs/mlx/tests/gptoss_mxfp4_tests.cpp new file mode 100644 index 0000000000000000000000000000000000000000..5a3f82e6d0c24ac6a6251f2e6a3171168e214f3d --- /dev/null +++ Layr-Labs/mlx/tests/gptoss_mxfp4_tests.cpp @@ -0,0 +1,214 @@ +// Copyright © 2026 Eigen Labs. +#include <algorithm> +#include <cmath> +#include <cstdint> +#include <cstdlib> +#include <optional> +#include <string> +#include <vector> + +#include "doctest/doctest.h" +#include "mlx/mlx.h" + +using namespace mlx::core; + +namespace { +struct ScopedPrefillTile { + std::optional<std::string> previous; + explicit ScopedPrefillTile(const char* value) { + if (const auto* old = std::getenv("MLX_GPTOSS_MXFP4_PREFILL_TILE")) + previous = old; + setenv("MLX_GPTOSS_MXFP4_PREFILL_TILE", value, 1); + } + ~ScopedPrefillTile() { + if (previous) + setenv("MLX_GPTOSS_MXFP4_PREFILL_TILE", previous->c_str(), 1); + else + unsetenv("MLX_GPTOSS_MXFP4_PREFILL_TILE"); + } +}; +} // namespace + +TEST_CASE( + "gptoss mxfp4 prefill tiles preserve expert boundaries and partial rows") { + constexpr int E = 32, N = 2880, K = 2880; + auto packed = reshape( + multiply( + arange(E * N * K / 8, uint32, Device::gpu), + array(uint32_t{2654435761}), + Device::gpu), + {E, N, K / 8}); + auto scales = full({E, N, K / 32}, array(uint8_t{126}), Device::gpu); + eval(packed, scales); + + for (const auto dtype : {float32, bfloat16}) { + for (int rows : {128, 131, 257}) { + std::vector<uint32_t> ids(rows); + for (int i = 0; i < rows; ++i) + ids[i] = (i * E) / rows; + const array indices(ids.data(), {rows}); + auto x = astype( + reshape( + sin(arange(rows * K, float32, Device::gpu) * array(0.017f), + Device::gpu), + {rows, 1, K}), + dtype, + Device::gpu); + array reference(0.0f); + { + ScopedPrefillTile tile("legacy"); + reference = gather_qmm( + x, + packed, + scales, + std::nullopt, + std::nullopt, + indices, + true, + 32, + 4, + "mxfp4", + true, + Device::gpu); + eval(reference); + } + for (const auto* name : {"m32n32k32"}) { + ScopedPrefillTile tile(name); + const auto actual = gather_qmm( + x, + packed, + scales, + std::nullopt, + std::nullopt, + indices, + true, + 32, + 4, + "mxfp4", + true, + Device::gpu); + eval(actual); + INFO("tile=", name, " rows=", rows, " dtype=", dtype); + const auto delta = + abs(astype(actual, float32) - astype(reference, float32)); + const float worst = max(delta).item<float>(); + const float magnitude = + max(abs(astype(reference, float32))).item<float>(); + CHECK(std::isfinite(worst)); + CHECK( + worst <= + (dtype == float32 ? 1e-4f : 0.016f) * std::max(1.0f, magnitude)); + } + } + } +} + +namespace { +struct ScopedDecodeFastTail { + std::optional<std::string> previous; + explicit ScopedDecodeFastTail(const char* value) { + if (const auto* old = std::getenv("MLX_GPTOSS_MXFP4_DECODE_FAST_TAIL")) + previous = old; + setenv("MLX_GPTOSS_MXFP4_DECODE_FAST_TAIL", value, 1); + } + ~ScopedDecodeFastTail() { + if (previous) + setenv("MLX_GPTOSS_MXFP4_DECODE_FAST_TAIL", previous->c_str(), 1); + else + unsetenv("MLX_GPTOSS_MXFP4_DECODE_FAST_TAIL"); + } +}; +} // namespace + +TEST_CASE("gptoss mxfp4 fast gather covers the final 320 input values") { + constexpr int E = 32, N = 2880, K = 2880; + auto packed = reshape( + multiply( + arange(E * N * K / 8, uint32, Device::gpu), + array(uint32_t{2654435761}), + Device::gpu), + {E, N, K / 8}); + auto scales = full({E, N, K / 32}, array(uint8_t{126}), Device::gpu); + eval(packed, scales); + for (const auto dtype : {float32, bfloat16}) { + for (int batch : {1, 2, 4, 8}) { + std::vector<uint32_t> lhs_ids(batch * 4), rhs_ids(batch * 4); + for (int row = 0; row < batch; ++row) { + for (int expert = 0; expert < 4; ++expert) { + lhs_ids[row * 4 + expert] = row; + rhs_ids[row * 4 + expert] = (row * 7 + expert * 5) % E; + } + } + const array lhs(lhs_ids.data(), {batch, 4}); + const array rhs(rhs_ids.data(), {batch, 4}); + for (bool tail_edges_only : {true, false}) { + std::vector<float> values(batch * K, 0.0f); + for (int row = 0; row < batch; ++row) { + if (tail_edges_only) { + values[row * K + 2559] = 0.75f; + values[row * K + 2560] = 0.25f; + values[row * K + 2879] = -0.5f; + } else { + for (int k = 0; k < K; ++k) + values[row * K + k] = std::sin((row * K + k) * 0.017f); + } + } + auto x = + astype(array(values.data(), {batch, 1, K}), dtype, Device::gpu); + array reference(0.0f); + { + ScopedDecodeFastTail flag("0"); + reference = gather_qmm( + x, + packed, + scales, + std::nullopt, + lhs, + rhs, + true, + 32, + 4, + "mxfp4", + false, + Device::gpu); + eval(reference); + } + ScopedDecodeFastTail flag("1"); + const auto actual = gather_qmm( + x, + packed, + scales, + std::nullopt, + lhs, + rhs, + true, + 32, + 4, + "mxfp4", + false, + Device::gpu); + eval(actual); + INFO( + "batch=", + batch, + " dtype=", + dtype, + " tail_edges_only=", + tail_edges_only); + if (tail_edges_only) { + CHECK(array_equal(actual, reference).item<bool>()); + } else { + const auto delta = + abs(astype(actual, float32) - astype(reference, float32)); + const float worst = max(delta).item<float>(); + const float magnitude = + max(abs(astype(reference, float32))).item<float>(); + CHECK(std::isfinite(worst)); + CHECK( + worst <= + (dtype == float32 ? 1e-4f : 0.016f) * std::max(1.0f, magnitude)); + } + } + } + } +}

Admission accounting used to combine independent reads of the allocator counters with logical buffer sizes, so concurrent allocator activity could produce an incoherent snapshot and alignment / cache reuse could make the real backing larger than the request. get_memory_snapshot() reads active / cache / peak under one lock; get_allocation_size_upper_bound() and the detached, lock-free AllocationFootprintPolicy bound a buffer including alignment and the (now inclusive, overflow-checked) cache-reuse limit shared with BufferCache. Metal, CPU and CUDA allocators implement the same interface; the allocator’s own behaviour stays the authority. These describe allocator accounting, not process RSS. (PR #15)

diff --git ml-explore/mlx/mlx/backend/common/allocation_footprint.h Layr-Labs/mlx/mlx/backend/common/allocation_footprint.h new file mode 100644 index 0000000000000000000000000000000000000000..c7ec23de6a33ba32426df48c458116a773d17f4b --- /dev/null +++ Layr-Labs/mlx/mlx/backend/common/allocation_footprint.h @@ -0,0 +1,39 @@ +// Copyright © 2026 Apple Inc. + +#pragma once + +#include <algorithm> +#include <cstddef> +#include <limits> +#include <stdexcept> + +namespace mlx::core::allocator { + +inline size_t checked_allocation_add(size_t size, size_t extra) { + if (extra > std::numeric_limits<size_t>::max() - size) { + throw std::overflow_error("allocation footprint overflow"); + } + return size + extra; +} + +inline size_t round_allocation_size(size_t size, size_t alignment) { + if (alignment == 0) { + throw std::invalid_argument("allocation alignment is zero"); + } + auto remainder = size % alignment; + return checked_allocation_add(size, remainder ? alignment - remainder : 0); +} + +// Inclusive cache-reuse bound. A fresh buffer never exceeds this bound. +inline size_t maximum_reuse_size(size_t size, size_t page_size) { + if (size == 0) { + return 0; + } + auto two_pages = checked_allocation_add(page_size, page_size); + if (two_pages == 0) { + throw std::invalid_argument("cache page size is zero"); + } + return checked_allocation_add(size, std::min(size - 1, two_pages - 1)); +} + +} // namespace mlx::core::allocator
diff --git ml-explore/mlx/mlx/backend/common/buffer_cache.h Layr-Labs/mlx/mlx/backend/common/buffer_cache.h index 27648b2fbc0736331fbe3d71d5b82e212848f606..177dbc11ac7dcc09f1472871e4b99ee8b58b1f75 100644 --- ml-explore/mlx/mlx/backend/common/buffer_cache.h +++ Layr-Labs/mlx/mlx/backend/common/buffer_cache.h @@ -7,6 +7,8 @@ #include <cassert> #include <functional> #include <map>   +#include "mlx/backend/common/allocation_footprint.h" + namespace mlx::core {   template <typename T> @@ -28,10 +30,13 @@ BufferCache(const BufferCache&) = delete; BufferCache& operator=(const BufferCache&) = delete;   T* reuse_from_cache(size_t size) { + if (size == 0) { + return nullptr; + } // Find the closest buffer in pool. auto it = buffer_pool_.lower_bound(size); if (it == buffer_pool_.end() || - it->first >= std::min(2 * size, size + 2 * page_size_)) { + it->first > allocator::maximum_reuse_size(size, page_size_)) { return nullptr; }
diff --git ml-explore/mlx/mlx/backend/cuda/allocator.cpp Layr-Labs/mlx/mlx/backend/cuda/allocator.cpp index b04d72da0ee140ca94af1d811baa0454809a07bd..e4833178452e6bf7ce44a88f761affb3649cba21 100644 --- ml-explore/mlx/mlx/backend/cuda/allocator.cpp +++ Layr-Labs/mlx/mlx/backend/cuda/allocator.cpp @@ -24,6 +24,19 @@ // Any allocations smaller than this will try to use the small pool constexpr int small_block_size = 8;   +size_t normalized_allocation_size(size_t size) { + if (size == 0) { + return 0; + } + if (size <= small_block_size) { + return 8; + } + if (size < page_size) { + return next_power_of_2(size); + } + return allocator::round_allocation_size(size, page_size); +} + // The small pool size in bytes. This should be a multiple of the host page // size and small_block_size. constexpr int small_pool_size = 4 * page_size; @@ -168,13 +181,7 @@ if (size == 0) { return Buffer{new CudaBuffer{nullptr, 0, -1}}; }   - if (size <= small_block_size) { - size = 8; - } else if (size < page_size) { - size = next_power_of_2(size); - } else { - size = page_size * ((size + page_size - 1) / page_size); - } + size = normalized_allocation_size(size);   if (size <= small_block_size || stream == nullptr) { device = -1; @@ -339,6 +346,11 @@ std::lock_guard lock(mutex_); peak_memory_ = 0; }   +MemorySnapshot CudaAllocator::get_memory_snapshot() { + std::lock_guard lock(mutex_); + return {active_memory_, buffer_cache_.cache_size(), peak_memory_}; +} + size_t CudaAllocator::get_memory_limit() { return memory_limit_; } @@ -404,6 +416,21 @@ }   } // namespace allocator   +AllocationFootprintPolicy get_allocation_footprint_policy() noexcept { + return {cu::page_size, 0, cu::small_block_size, cu::page_size, cu::page_size}; +} + +size_t get_allocation_size_upper_bound(size_t size) { + size_t result; + if (!get_allocation_footprint_policy().upper_bound(size, result)) { + throw std::overflow_error("allocation footprint overflow"); + } + return result; +} + +MemorySnapshot get_memory_snapshot() { + return cu::allocator().get_memory_snapshot(); +} size_t get_active_memory() { return cu::allocator().get_active_memory(); } @@ -421,6 +448,13 @@ return cu::allocator().get_memory_limit(); } size_t get_cache_memory() { return cu::allocator().get_cache_memory(); +} +size_t get_num_resources() { + // CUDA allocator does not track a Metal-style resource count; report 0. + return 0; +} +size_t get_resource_limit() { + return 0; } size_t set_cache_limit(size_t limit) { return cu::allocator().set_cache_limit(limit);
diff --git ml-explore/mlx/mlx/backend/cuda/allocator.h Layr-Labs/mlx/mlx/backend/cuda/allocator.h index af76ad9059bf5e7f9e959ca59a76769e9e98b172..df703b8ff3ac3ae4508e0c5266e110b95016acf9 100644 --- ml-explore/mlx/mlx/backend/cuda/allocator.h +++ Layr-Labs/mlx/mlx/backend/cuda/allocator.h @@ -5,6 +5,7 @@ #include "mlx/allocator.h" #include "mlx/backend/common/buffer_cache.h" #include "mlx/backend/cuda/cuda_utils.h" +#include "mlx/memory.h"   #include <cuda_runtime.h> #include <mutex> @@ -58,6 +59,7 @@ // Replace the memory of |buf| with unified memory (managed memory or pinned // host memory), and copy the data over. Pass |stream| to copy asynchronously. void move_to_unified_memory(CudaBuffer& buf, cudaStream_t stream = nullptr);   + MemorySnapshot get_memory_snapshot(); size_t get_active_memory() const; size_t get_peak_memory() const; void reset_peak_memory();
diff --git ml-explore/mlx/mlx/backend/metal/allocator.cpp Layr-Labs/mlx/mlx/backend/metal/allocator.cpp index 903cce21879a94de36929c0e5568e9f3d6e09ca7..9b80840a18901710c9790b2f46f7ce9410a3f934 100644 --- ml-explore/mlx/mlx/backend/metal/allocator.cpp +++ Layr-Labs/mlx/mlx/backend/metal/allocator.cpp @@ -7,8 +7,11 @@ #include "mlx/memory.h"   #include <mach/vm_page_size.h> #include <unistd.h> +#include <algorithm> #include <cassert> +#include <cerrno> #include <cstdlib> +#include <limits>   namespace mlx::core {   @@ -60,6 +63,30 @@ auto memsize = std::get<size_t>(info.at("memory_size")); auto max_rec_size = std::get<size_t>(info.at("max_recommended_working_set_size")); resource_limit_ = std::get<size_t>(info.at("resource_limit")); + // Optional override (Darkbloom): MLX_RESOURCE_LIMIT lets an operator/test pin + // the Metal resource-COUNT ceiling below the OS default (~499000). Used to + // deterministically exercise the count-aware high-water trim, and as a safety + // valve to force earlier cache reclamation on a box seeing the resource-limit + // crash. The value may only LOWER the ceiling (it is clamped to the OS limit) + // — raising it above what the hardware/OS reports would invite the very crash + // this guards against. Strictly validated: a plain unsigned decimal that + // consumes the whole string, is non-zero, and does not overflow; anything + // else (empty, sign, junk, range error) is ignored and the OS limit stands. + if (const char* rl = std::getenv("MLX_RESOURCE_LIMIT")) { + while (*rl == ' ' || *rl == '\t') { + ++rl; + } + if (*rl >= '0' && *rl <= '9') { // unsigned decimal only (reject sign/junk) + errno = 0; + char* end = nullptr; + unsigned long long v = std::strtoull(rl, &end, 10); + bool consumed_all = end != rl && *end == '\0'; + if (consumed_all && errno != ERANGE && v > 0 && + v <= std::numeric_limits<size_t>::max()) { + resource_limit_ = std::min(static_cast<size_t>(v), resource_limit_); + } + } + } block_limit_ = std::min(1.5 * max_rec_size, 0.95 * memsize); gc_limit_ = std::min(static_cast<size_t>(0.95 * max_rec_size), block_limit_); max_pool_size_ = block_limit_; @@ -93,6 +120,11 @@ static_cast<size_t>(0.95 * device_->recommendedMaxWorkingSetSize())); return limit; };   +MemorySnapshot MetalAllocator::get_memory_snapshot() { + std::lock_guard lock(mutex_); + return {active_memory_, buffer_cache_.cache_size(), peak_memory_}; +} + size_t MetalAllocator::get_memory_limit() { return block_limit_; } @@ -122,7 +154,7 @@ }   // Align up memory if (size > vm_page_size) { - size = vm_page_size * ((size + vm_page_size - 1) / vm_page_size); + size = allocator::round_allocation_size(size, vm_page_size); }   // Try the cache @@ -131,12 +163,34 @@ MTL::Buffer* buf = buffer_cache_.reuse_from_cache(size); if (!buf) { size_t mem_required = get_active_memory() + get_cache_memory() + size;   - // If we have a lot of memory pressure try to reclaim memory from the cache + // If we have a lot of memory pressure try to reclaim memory from the cache. + // NOTE: release_cached_buffers takes a BYTES-to-free target; when the + // buffers are tiny this frees only a few entries even though the COUNT is + // the binding constraint, so the byte path alone cannot bound + // num_resources_ (see the count-aware reclaim below). if (mem_required >= gc_limit_ || num_resources_ >= resource_limit_) { num_resources_ -= buffer_cache_.release_cached_buffers(mem_required - gc_limit_); }   + // Count-aware reclaim (Darkbloom): the Metal resource COUNT limit + // (resource_limit_, ~iogpu.rsrc_limit/499000) is independent of byte usage. + // Under churn with many distinct buffer shapes (varied prompt lengths, + // growing KV caches, multiple co-resident models) freed buffers are + // recycled into the size-keyed cache and never reused at that exact size, + // so the cache ENTRY COUNT creeps toward the limit while byte usage stays + // modest — the byte-driven trim above never fires (its threshold is + // ~physical RAM). Once the count crosses a high-water mark, proactively + // clear the cache (pure reuse pool — clearing only costs re-allocation, + // never correctness) so the count drops back to the live working set. This + // makes the count limit unreachable by any request mix / batching method, + // while the existing byte limits keep total memory below physical RAM. + if (resource_limit_ > 0 && + num_resources_ >= (resource_limit_ * resource_high_water_num_) / + resource_high_water_den_) { + num_resources_ -= buffer_cache_.clear(); + } + // Allocate new buffer if needed if (num_resources_ >= resource_limit_) { std::ostringstream msg; @@ -242,6 +296,18 @@ }   } // namespace metal   +AllocationFootprintPolicy get_allocation_footprint_policy() noexcept { + return {vm_page_size, vm_page_size, 0, 0, vm_page_size}; +} + +size_t get_allocation_size_upper_bound(size_t size) { + size_t result; + if (!get_allocation_footprint_policy().upper_bound(size, result)) { + throw std::overflow_error("allocation footprint overflow"); + } + return result; +} + size_t set_cache_limit(size_t limit) { return metal::allocator().set_cache_limit(limit); } @@ -260,6 +326,9 @@ "the maximum working set size is not allowed."); } return metal::allocator().set_wired_limit(limit); } +MemorySnapshot get_memory_snapshot() { + return metal::allocator().get_memory_snapshot(); +} size_t get_active_memory() { return metal::allocator().get_active_memory(); } @@ -271,6 +340,12 @@ metal::allocator().reset_peak_memory(); } size_t get_cache_memory() { return metal::allocator().get_cache_memory(); +} +size_t get_num_resources() { + return metal::allocator().get_num_resources(); +} +size_t get_resource_limit() { + return metal::allocator().get_resource_limit(); } void clear_cache() { return metal::allocator().clear_cache();
diff --git ml-explore/mlx/mlx/backend/metal/allocator.h Layr-Labs/mlx/mlx/backend/metal/allocator.h index 885951364e52cd93a82ab05b59839cc38875a5bb..623055e92c016b6bb4a1af1f1372e4400a6b327d 100644 --- ml-explore/mlx/mlx/backend/metal/allocator.h +++ Layr-Labs/mlx/mlx/backend/metal/allocator.h @@ -9,6 +9,7 @@ #include "mlx/allocator.h" #include "mlx/backend/common/buffer_cache.h" #include "mlx/backend/metal/device.h" +#include "mlx/memory.h"   namespace mlx::core::metal {   @@ -23,6 +24,7 @@ virtual size_t size(Buffer buffer) const override; virtual Buffer make_buffer(void* ptr, size_t size) override; virtual void release(Buffer buffer) override;   + MemorySnapshot get_memory_snapshot(); size_t get_active_memory() { return active_memory_; }; @@ -36,6 +38,17 @@ }; size_t get_cache_memory() { return buffer_cache_.cache_size(); }; + // Live Metal resource (buffer) COUNT and its hard ceiling. The count limit + // (default iogpu.rsrc_limit, ~499000) is independent of the byte limits and + // is what malloc() throws on when reached; exposed so callers can observe and + // bound it (it can be far higher than byte usage implies when many tiny + // buffers accumulate in the cache). + size_t get_num_resources() { + return num_resources_; + }; + size_t get_resource_limit() { + return resource_limit_; + }; size_t set_cache_limit(size_t limit); size_t set_memory_limit(size_t limit); size_t get_memory_limit(); @@ -71,6 +84,14 @@ size_t max_pool_size_; size_t wired_limit_{0}; size_t num_resources_{0}; size_t resource_limit_{0}; + + // Count-aware cache-reclaim high-water mark (Darkbloom): when num_resources_ + // reaches resource_high_water_num_/den_ of resource_limit_ (90%), malloc() + // proactively clears the (pure-reuse) buffer cache so the resource COUNT can + // never reach resource_limit_ and throw, regardless of buffer-byte sizes. + // Integer fraction to avoid float work in the allocation hot path. + static constexpr size_t resource_high_water_num_ = 9; + static constexpr size_t resource_high_water_den_ = 10;   std::mutex mutex_; };
diff --git ml-explore/mlx/mlx/backend/no_gpu/allocator.cpp Layr-Labs/mlx/mlx/backend/no_gpu/allocator.cpp index a800e381a722754289b35c5a41b59676d6463bc4..f766707bce4e06a2e78db549fe18568d99d42979 100644 --- ml-explore/mlx/mlx/backend/no_gpu/allocator.cpp +++ Layr-Labs/mlx/mlx/backend/no_gpu/allocator.cpp @@ -41,6 +41,11 @@ virtual Buffer malloc(size_t size) override; virtual void free(Buffer buffer) override; virtual size_t size(Buffer buffer) const override;   + MemorySnapshot get_memory_snapshot() const { + std::lock_guard lock(mutex_); + return {active_memory_, buffer_cache_.cache_size(), peak_memory_}; + } + size_t get_active_memory() const { return active_memory_; }; @@ -184,6 +189,21 @@ }   } // namespace allocator   +AllocationFootprintPolicy get_allocation_footprint_policy() noexcept { + return {1, 0, 0, 0, 4096}; +} + +size_t get_allocation_size_upper_bound(size_t size) { + size_t result; + if (!get_allocation_footprint_policy().upper_bound(size, result)) { + throw std::overflow_error("allocation footprint overflow"); + } + return result; +} + +MemorySnapshot get_memory_snapshot() { + return allocator::common_allocator().get_memory_snapshot(); +} size_t get_active_memory() { return allocator::common_allocator().get_active_memory(); } @@ -202,6 +222,12 @@ }   size_t get_cache_memory() { return allocator::common_allocator().get_cache_memory(); +} +size_t get_num_resources() { + return 0; +} +size_t get_resource_limit() { + return 0; } size_t set_cache_limit(size_t limit) { return allocator::common_allocator().set_cache_limit(limit);
diff --git ml-explore/mlx/mlx/memory.h Layr-Labs/mlx/mlx/memory.h index f4eabc99760afdce5ea2fa313c8cd1fc71be9afd..f2e7bf225faee895ed93ee2fc2dce25eb6acd06a 100644 --- ml-explore/mlx/mlx/memory.h +++ Layr-Labs/mlx/mlx/memory.h @@ -2,12 +2,91 @@ // Copyright © 2025 Apple Inc.   #pragma once   +#include <algorithm> #include <cstdlib> +#include <limits>   #include "mlx/api.h"   namespace mlx::core {   +struct MemorySnapshot { + size_t active_memory; + size_t cache_memory; + size_t peak_memory; +}; + +/* Read allocator accounting under one lock. This does not synchronize streams + * or include allocations that have not entered allocator accounting yet. */ +MLX_API MemorySnapshot get_memory_snapshot(); + +/* Bound one allocator buffer, including alignment and larger cache reuse. + * This does not allocate, synchronize, or inspect the live buffer cache. */ +MLX_API size_t get_allocation_size_upper_bound(size_t size); + +// A detached value: evaluating bounds requires no allocator, error callback, +// exception, lock, or allocation. Capture once before entering admission locks. +struct AllocationFootprintPolicy { + size_t alignment; + size_t rounding_threshold; + size_t minimum_allocation; + size_t power_of_two_below; + size_t cache_page_size; + + bool upper_bound(size_t size, size_t& result) const noexcept { + if (size == 0) { + result = 0; + return true; + } + if (alignment == 0 || cache_page_size == 0) { + return false; + } + size = std::max(size, minimum_allocation); + if (power_of_two_below && size < power_of_two_below) { + size_t rounded = 1; + while (rounded < size) { + if (!add(rounded, rounded, rounded)) { + return false; + } + } + size = rounded; + } else if (size > rounding_threshold) { + auto remainder = size % alignment; + if (remainder && !add(size, alignment - remainder, size)) { + return false; + } + } + size_t two_pages; + return add(cache_page_size, cache_page_size, two_pages) && + add(size, std::min(size - 1, two_pages - 1), result); + } + + // For any positive n with a valid bound, B(n) <= n + this overhead. + bool maximum_extra_bytes(size_t& result) const noexcept { + if (alignment == 0 || cache_page_size == 0) { + return false; + } + auto normalization = std::max( + {alignment - 1, + minimum_allocation ? minimum_allocation - 1 : 0, + power_of_two_below ? power_of_two_below - 1 : 0}); + size_t two_pages; + return add(cache_page_size, cache_page_size, two_pages) && + add(normalization, two_pages - 1, result); + } + + private: + static bool add(size_t a, size_t b, size_t& result) noexcept { + if (b > std::numeric_limits<size_t>::max() - a) { + return false; + } + result = a + b; + return true; + } +}; + +MLX_API AllocationFootprintPolicy get_allocation_footprint_policy() noexcept; + /* Get the actively used memory in bytes. * * Note, this will not always match memory use reported by the system because @@ -32,6 +111,21 @@ * The cache includes memory not currently used that has not been returned * to the system allocator. * */ MLX_API size_t get_cache_memory(); + +/* Get the number of live Metal resources (buffers). + * + * This is a COUNT, independent of byte usage. The Metal backend throws when it + * reaches the resource limit (see get_resource_limit). Many small cached + * buffers can push this count high while byte usage stays low. + * */ +MLX_API size_t get_num_resources(); + +/* Get the hard ceiling on the number of live Metal resources (buffers). + * + * Defaults to the iogpu.rsrc_limit sysctl (~499000 when unset). Allocation + * throws once get_num_resources() reaches this value. + * */ +MLX_API size_t get_resource_limit();   /* Set the memory limit. * The memory limit is a guideline for the maximum amount of memory to use
diff --git ml-explore/mlx/tests/allocator_tests.cpp Layr-Labs/mlx/tests/allocator_tests.cpp index 9658e29f008667659088412ea0040a5af56e1cfc..d5ff60525b3d8b24a8857604eaf66662b2f6884c 100644 --- ml-explore/mlx/tests/allocator_tests.cpp +++ Layr-Labs/mlx/tests/allocator_tests.cpp @@ -1,7 +1,9 @@ // Copyright © 2023-2026 Apple Inc.   +#include <atomic> #include <chrono> #include <future> +#include <limits> #include <memory> #include <stdexcept> #include <thread> @@ -118,3 +120,131 @@ clear_thread.join();   set_cache_limit(old_limit); } + +TEST_CASE("test coherent memory snapshot during cache transfers") { + auto old_limit = set_cache_limit(1 << 20); + clear_cache(); + auto live = allocator::malloc(32768); + auto moving = allocator::malloc(65536); + auto moving_bytes = allocator::allocator().size(moving); + auto active = get_memory_snapshot(); + allocator::free(moving); + auto cached = get_memory_snapshot(); + CHECK_EQ(active.active_memory, cached.active_memory + moving_bytes); + CHECK_EQ(active.cache_memory + moving_bytes, cached.cache_memory); + CHECK_EQ(active.peak_memory, cached.peak_memory); + auto total = cached.active_memory + cached.cache_memory; + + std::atomic<bool> stop{false}; + std::promise<void> started; + auto started_future = started.get_future(); + auto churn = std::async(std::launch::async, [&] { + started.set_value(); + size_t transfers = 0; + do { + auto buffer = allocator::malloc(65536); + allocator::free(buffer); + ++transfers; + } while (!stop.load(std::memory_order_relaxed)); + return transfers; + }); + started_future.wait(); + + size_t inconsistent = 0; + for (int i = 0; i < 20000; ++i) { + auto snapshot = get_memory_snapshot(); + inconsistent += snapshot.active_memory + snapshot.cache_memory != total; + inconsistent += snapshot.active_memory > snapshot.peak_memory; + } + stop.store(true, std::memory_order_relaxed); + CHECK_GT(churn.get(), 0); + CHECK_EQ(inconsistent, 0); + + allocator::free(live); + clear_cache(); + set_cache_limit(old_limit); +} + +TEST_CASE("test memory snapshot does not wait for stream work") { + std::promise<void> started; + auto started_future = started.get_future(); + std::promise<void> release; + auto release_future = release.get_future().share(); + auto stream = new_stream(Device{Device::cpu}); + scheduler::enqueue(stream, [&] { + started.set_value(); + release_future.wait(); + }); + started_future.wait(); + + auto snapshot = + std::async(std::launch::async, [] { return get_memory_snapshot(); }); + auto status = snapshot.wait_for(std::chrono::seconds(1)); + release.set_value(); + synchronize(stream); + snapshot.get(); + CHECK_EQ(status, std::future_status::ready); +} + +TEST_CASE("allocation footprint bounds fresh and cached buffer owners") { + auto old_limit = set_cache_limit(1 << 20); + clear_cache(); + for (size_t size : {size_t(4), size_t(6000), size_t(24576), size_t(65536)}) { + auto bound = get_allocation_size_upper_bound(size); + CHECK_GE(bound, size); + auto buffer = allocator::malloc(size); + CHECK_LE(allocator::allocator().size(buffer), bound); + allocator::free(buffer); + } + clear_cache(); + auto large = allocator::malloc(49152); + auto larger_bytes = allocator::allocator().size(large); + allocator::free(large); + auto reused = allocator::malloc(32768); + CHECK_LE( + allocator::allocator().size(reused), + get_allocation_size_upper_bound(32768)); + CHECK_GE(larger_bytes, 49152); + allocator::free(reused); + clear_cache(); + set_cache_limit(old_limit); +} + +TEST_CASE("allocation prediction does not change allocator counters") { + auto before = get_memory_snapshot(); + CHECK_GE(get_allocation_size_upper_bound(24576), 24576); + CHECK_EQ(get_allocation_size_upper_bound(0), 0); + CHECK_THROWS( + get_allocation_size_upper_bound(std::numeric_limits<size_t>::max())); + auto after = get_memory_snapshot(); + CHECK_EQ(after.active_memory, before.active_memory); + CHECK_EQ(after.cache_memory, before.cache_memory); + CHECK_EQ(after.peak_memory, before.peak_memory); +} + +TEST_CASE( + "detached allocation policies preserve backend size classes without exceptions") { + const AllocationFootprintPolicy cpu{1, 0, 0, 0, 4096}; + const AllocationFootprintPolicy metal{16384, 16384, 0, 0, 16384}; + const AllocationFootprintPolicy cuda{16384, 0, 8, 16384, 16384}; + size_t result = 0; + CHECK(cpu.upper_bound(24576, result)); + CHECK_EQ(result, 32767); + CHECK(metal.upper_bound(24576, result)); + CHECK_EQ(result, 65535); + CHECK(cuda.upper_bound(1, result)); + CHECK_EQ(result, 15); + CHECK(cuda.upper_bound(8193, result)); + CHECK_EQ(result, 32767); + for (auto p : {cpu, metal, cuda}) { + size_t extra = 0; + REQUIRE(p.maximum_extra_bytes(extra)); + for (size_t n : + {1, 4, 8191, 8192, 8193, 16383, 16384, 16385, 24576, 65537}) { + REQUIRE(p.upper_bound(n, result)); + CHECK_GE(result, n); + CHECK_LE(result - n, extra); + } + CHECK_FALSE(p.upper_bound(std::numeric_limits<size_t>::max(), result)); + } +}

The two-pass vector SDPA cast unnormalized partial sums to the input dtype before combining them: BF16 lost cancellation residuals and FP16 could overflow even when the final output was representable. Partials now stay FP32 through the second pass and only the final result is cast; kernel entry names change so a stale shader cannot silently satisfy the new buffer contract. Python regressions cover cancellation and overflow across D64/D128/D256, 32⁄128 partitions, masked and unmasked. (PR #16)

diff --git ml-explore/mlx/mlx/backend/metal/kernels/scaled_dot_product_attention.metal Layr-Labs/mlx/mlx/backend/metal/kernels/scaled_dot_product_attention.metal index 44ccb834e96d335e66c6ff48b6ce401b6661851c..cfe5aec6c35b9899c6eb76933145d8ff0c095b35 100644 --- ml-explore/mlx/mlx/backend/metal/kernels/scaled_dot_product_attention.metal +++ Layr-Labs/mlx/mlx/backend/metal/kernels/scaled_dot_product_attention.metal @@ -9,7 +9,7 @@ // SDPA vector instantiations #define instantiate_sdpa_vector_aggregation(type, value_dim) \ instantiate_kernel( \ - "sdpa_vector_2pass_2_" #type "_" #value_dim, \ + "sdpa_vector_2pass_fp32partials_2_" #type "_" #value_dim, \ sdpa_vector_2pass_2, \ type, \ value_dim) @@ -22,21 +22,35 @@ type, \ qk_dim, \ value_dim) \ instantiate_kernel( \ - "sdpa_vector_2pass_1_" #type "_" #qk_dim "_" #value_dim, \ + "sdpa_vector_2pass_fp32partials_1_" #type "_" #qk_dim "_" #value_dim, \ sdpa_vector_2pass_1, \ type, \ qk_dim, \ value_dim)   +#define instantiate_sdpa_vector_gqa(type, qk_dim, value_dim, hpt) \ + instantiate_kernel( \ + "sdpa_vector_2pass_fp32partials_1_gqa_" #type "_" #qk_dim "_" #value_dim, \ + sdpa_vector_2pass_1_gqa, \ + type, \ + qk_dim, \ + value_dim, \ + 8, \ + hpt) + #define instantiate_sdpa_vector_heads(type) \ instantiate_sdpa_vector(type, 64, 64) \ instantiate_sdpa_vector(type, 96, 96) \ instantiate_sdpa_vector(type, 128, 128) \ instantiate_sdpa_vector(type, 192, 128) \ + instantiate_sdpa_vector(type, 192, 192) \ instantiate_sdpa_vector(type, 256, 256) \ + instantiate_sdpa_vector_gqa(type, 64, 64, 8) \ + instantiate_sdpa_vector_gqa(type, 128, 128, 4) \ instantiate_sdpa_vector_aggregation(type, 64) \ instantiate_sdpa_vector_aggregation(type, 96) \ instantiate_sdpa_vector_aggregation(type, 128) \ + instantiate_sdpa_vector_aggregation(type, 192) \ instantiate_sdpa_vector_aggregation(type, 256)   instantiate_sdpa_vector_heads(float)
diff --git ml-explore/mlx/mlx/backend/metal/kernels/sdpa_vector.h Layr-Labs/mlx/mlx/backend/metal/kernels/sdpa_vector.h index 1eec72be31da3b5d1ef2fc0cec40668423e7cdbc..3631e49daadbac2ca5c15d26c889b1d649840690 100644 --- ml-explore/mlx/mlx/backend/metal/kernels/sdpa_vector.h +++ Layr-Labs/mlx/mlx/backend/metal/kernels/sdpa_vector.h @@ -181,7 +181,7 @@ [[kernel]] void sdpa_vector_2pass_1( const device T* queries [[buffer(0)]], const device T* keys [[buffer(1)]], const device T* values [[buffer(2)]], - device T* out [[buffer(3)]], + device float* out [[buffer(3)]], device float* sums [[buffer(4)]], device float* maxs [[buffer(5)]], const constant int& N [[buffer(7)]], @@ -313,13 +313,156 @@ maxs[0] = max_score; }   for (int i = 0; i < v_per_thread; i++) { - out[i] = static_cast<T>(o[i]); + out[i] = o[i]; + } +} + +// Duplication-free variant for high gqa_factor decode: each simdgroup owns a +// contiguous token sub-chunk and computes HPT of its group's query heads, so +// each K/V byte is read G / HPT times instead of G times. Single-token +// queries without mask or sinks only; the partials layout matches +// sdpa_vector_2pass_2. +template <typename T, int D, int V, int G, int HPT> +[[kernel]] void sdpa_vector_2pass_1_gqa( + const device T* queries [[buffer(0)]], + const device T* keys [[buffer(1)]], + const device T* values [[buffer(2)]], + device float* out [[buffer(3)]], + device float* sums [[buffer(4)]], + device float* maxs [[buffer(5)]], + const constant int& N [[buffer(7)]], + const constant size_t& k_head_stride [[buffer(8)]], + const constant size_t& k_seq_stride [[buffer(9)]], + const constant size_t& v_head_stride [[buffer(10)]], + const constant size_t& v_seq_stride [[buffer(11)]], + const constant float& scale [[buffer(12)]], + uint3 tid [[threadgroup_position_in_grid]], + uint3 tpg [[threadgroups_per_grid]], + uint3 tidtg [[thread_position_in_threadgroup]], + uint simd_lid [[thread_index_in_simdgroup]]) { + constexpr int BD = 32; + constexpr int qk_per_thread = D / BD; + constexpr int v_per_thread = V / BD; + constexpr int NT = G / HPT; + + typedef float U; + + const int kv_head_idx = tid.x; + const int batch_idx = tid.y; + const int block_idx = tid.z; + const int blocks = tpg.z; + const int g = tidtg.y; + const int cchunk = g / NT; + const int h0 = (g % NT) * HPT; + const int num_kv_heads = tpg.x; + const int num_q_heads = num_kv_heads * G; + const int base_head = batch_idx * num_q_heads + kv_head_idx * G; + + const int chunk = (N + blocks - 1) / blocks; + const int kstart = block_idx * chunk; + const int kend = min(N, kstart + chunk); + const int sub = (chunk + HPT - 1) / HPT; + const int s0 = kstart + cchunk * sub; + const int s1 = min(kend, s0 + sub); + + const device T* kp = keys + kv_head_idx * k_head_stride + s0 * k_seq_stride + + simd_lid * qk_per_thread; + const device T* vp = values + kv_head_idx * v_head_stride + + s0 * v_seq_stride + simd_lid * v_per_thread; + + U q[HPT][qk_per_thread]; + for (int j = 0; j < HPT; j++) { + const device T* qp = + queries + (base_head + h0 + j) * D + simd_lid * qk_per_thread; + for (int i = 0; i < qk_per_thread; i++) { + q[j][i] = static_cast<U>(scale) * qp[i]; + } + } + + U max_score[HPT]; + U sum_exp_score[HPT]; + U o[HPT][v_per_thread]; + for (int j = 0; j < HPT; j++) { + max_score[j] = Limits<U>::finite_min; + sum_exp_score[j] = 0; + for (int i = 0; i < v_per_thread; i++) { + o[j][i] = 0; + } + } + + for (int t = s0; t < s1; t++) { + U kr[qk_per_thread]; + U vr[v_per_thread]; + for (int i = 0; i < qk_per_thread; i++) { + kr[i] = kp[i]; + } + for (int i = 0; i < v_per_thread; i++) { + vr[i] = vp[i]; + } + kp += k_seq_stride; + vp += v_seq_stride; + for (int j = 0; j < HPT; j++) { + U score = 0; + for (int i = 0; i < qk_per_thread; i++) { + score += q[j][i] * kr[i]; + } + score = simd_sum(score); + U new_max = max(max_score[j], score); + U factor = fast::exp(max_score[j] - new_max); + U exp_score = fast::exp(score - new_max); + max_score[j] = new_max; + sum_exp_score[j] = sum_exp_score[j] * factor + exp_score; + for (int i = 0; i < v_per_thread; i++) { + o[j][i] = o[j][i] * factor + exp_score * vr[i]; + } + } + } + + threadgroup U o_sh[G * HPT * V]; + threadgroup U se_sh[G * HPT]; + threadgroup U mx_sh[G * HPT]; + for (int j = 0; j < HPT; j++) { + int slot = (h0 + j) * HPT + cchunk; + U inv = sum_exp_score[j] > 0 ? 1 / sum_exp_score[j] : 0; + for (int i = 0; i < v_per_thread; i++) { + o_sh[slot * V + simd_lid * v_per_thread + i] = o[j][i] * inv; + } + if (simd_lid == 0) { + se_sh[slot] = sum_exp_score[j]; + mx_sh[slot] = max_score[j]; + } + } + threadgroup_barrier(mem_flags::mem_threadgroup); + + U gmax = Limits<U>::finite_min; + for (int s = 0; s < HPT; s++) { + gmax = max(gmax, mx_sh[g * HPT + s]); + } + U denom = 0; + U acc[v_per_thread] = {0}; + for (int s = 0; s < HPT; s++) { + U w = se_sh[g * HPT + s] * fast::exp(mx_sh[g * HPT + s] - gmax); + denom += w; + for (int i = 0; i < v_per_thread; i++) { + acc[i] += w * o_sh[(g * HPT + s) * V + simd_lid * v_per_thread + i]; + } + } + + const int o_offset = base_head + g; + device float* op = + out + o_offset * blocks * V + block_idx * V + simd_lid * v_per_thread; + for (int i = 0; i < v_per_thread; i++) { + op[i] = acc[i]; + } + if (simd_lid == 0) { + sums[o_offset * blocks + block_idx] = denom; + maxs[o_offset * blocks + block_idx] = gmax; } }   template <typename T, int D> [[kernel]] void sdpa_vector_2pass_2( - const device T* partials [[buffer(0)]], + const device float* partials [[buffer(0)]], const device float* sums [[buffer(1)]], const device float* maxs [[buffer(2)]], device T* out [[buffer(3)]],
diff --git ml-explore/mlx/mlx/backend/metal/scaled_dot_product_attention.cpp Layr-Labs/mlx/mlx/backend/metal/scaled_dot_product_attention.cpp index cb4af8523e44e6ca814dcc7a3200feab9ad12332..c2ee37f29eea63bef2893459f14b9d66d37ff758 100644 --- ml-explore/mlx/mlx/backend/metal/scaled_dot_product_attention.cpp +++ Layr-Labs/mlx/mlx/backend/metal/scaled_dot_product_attention.cpp @@ -28,12 +28,13 @@ const std::optional<array>& mask, const std::optional<array>& sinks) { using namespace mlx::steel;   - int wm = 4; - int wn = 1; - int bd = q.shape(-1); int bq = 64; int bk = 32; + + bool split_d = bd == 256; + int wm = 4; + int wn = split_d ? 2 : 1;   int B = q.shape(0); int H = q.shape(1); @@ -43,6 +44,36 @@ int qL = q.shape(2); int kL = k.shape(2);   + // The causal offset describes the true diagonal even if kL is widened + // below. + int qL_off = kL - qL; + + // Check if K/V are from chunked KV cache, and assume aligned K/V if so. + auto has_backing_rows = [](const array& kv, int rows) { + auto& st = kv.strides(); + if ((st[0] < 0) || (st[1] <= 0) || (st[2] <= 0) || (st[1] % st[2] != 0)) { + return false; + } + int64_t itemsize = kv.itemsize(); + // The rows must stay inside the head's row pitch (so they belong to the + // cache the slice was taken from) ... + int64_t pitch = st[1] / st[2]; + int64_t row0 = ((kv.offset() / itemsize) % st[1]) / st[2]; + if (row0 + rows > pitch) { + return false; + } + // ... and inside the buffer. + int64_t end = (kv.shape(0) - 1) * st[0] + (kv.shape(1) - 1) * st[1] + + (rows - 1) * st[2] + kv.shape(3); + return kv.offset() + end * itemsize <= int64_t(kv.buffer_size()); + }; + if (split_d && do_causal_ && !mask.has_value() && (kL % bk)) { + int kLp = bk * ((kL + bk - 1) / bk); + if (has_backing_rows(k, kLp) && has_backing_rows(v, kLp)) { + kL = kLp; + } + } + const bool align_Q = (qL % bq) == 0; const bool align_K = (kL % bk) == 0; const bool has_mask = mask.has_value(); @@ -59,7 +90,7 @@ std::string base_name; concatenate( base_name, - "steel_attention_", + split_d ? "steel_attention_dsplit_" : "steel_attention_", type_to_name(q), "_bq", bq, @@ -102,7 +133,8 @@ bk, bd, wm, wn, - (has_mask ? *mask : q)); + (has_mask ? *mask : q), + split_d);   compute_encoder.set_compute_pipeline_state(kernel);   @@ -131,7 +163,7 @@ /* int NK_aligned = */ NK_aligned,   /* int qL_rem = */ (qL - NQ_aligned * bq), /* int kL_rem = */ (kL - NK_aligned * bk), - /* int qL_off = */ (kL - qL), + /* int qL_off = */ qL_off,   /* int64_t Q_strides[3] = */ {q.strides(0), q.strides(1), q.strides(2)}, /* int64_t K_strides[3] = */ {k.strides(0), k.strides(1), k.strides(2)}, @@ -160,6 +192,7 @@ MTL::Size grid_dims = MTL::Size(NQ, H, B); MTL::Size group_dims = MTL::Size(32, wm, wn);   + check_kernel_threadgroup_size(kernel, group_dims, hash_name); compute_encoder.dispatch_threadgroups(grid_dims, group_dims); }   @@ -174,7 +207,16 @@ array& o, bool do_causal_, const std::optional<array>& mask, const std::optional<array>& sinks) { - if (metal::is_nax_available() && q.shape(3) != 80 && + int B = q.shape(0); + int H = q.shape(1); + int D = q.shape(3); + int gqa_factor = q.shape(1) / k.shape(1); + + int qL = q.shape(2); + int kL = k.shape(2); + + if (metal::is_nax_available() && + (D == 64 || D == 96 || D == 128 || D == 256) && (env::enable_tf32() || q.dtype() != float32)) { return sdpa_full_self_attention_nax( /* const Stream& s = */ s, @@ -198,14 +240,6 @@ int bd = q.shape(-1); int bq = 32; int bk = bd < 128 ? 32 : 16;   - int B = q.shape(0); - int H = q.shape(1); - int D = q.shape(3); - int gqa_factor = q.shape(1) / k.shape(1); - - int qL = q.shape(2); - int kL = k.shape(2); - const bool align_Q = (qL % bq) == 0; const bool align_K = (kL % bk) == 0; const bool has_mask = mask.has_value(); @@ -323,6 +357,7 @@ MTL::Size grid_dims = MTL::Size(NQ, H, B); MTL::Size group_dims = MTL::Size(32, wm, wn);   + check_kernel_threadgroup_size(kernel, group_dims, hash_name); compute_encoder.dispatch_threadgroups(grid_dims, group_dims); }   @@ -430,7 +465,13 @@ const std::optional<array>& sinks) { // Set the kernel name std::string kname; kname.reserve(64); - kname += "sdpa_vector_2pass_1_"; + kname += "sdpa_vector_2pass_fp32partials_1"; + if (!mask && !sinks && q.shape(2) == 1 && q.shape(1) == 8 * k.shape(1) && + q.shape(-1) == v.shape(-1) && (q.shape(-1) == 64 || q.shape(-1) == 128) && + k.shape(2) >= 8192) { + kname += "_gqa"; + } + kname += "_"; kname += get_type_string(q.dtype()); kname += "_"; kname += std::to_string(q.shape(-1)); @@ -494,7 +535,8 @@ intermediate_shape.insert( intermediate_shape.end(), out.shape().begin(), out.shape().end() - 1); intermediate_shape.push_back(blocks); intermediate_shape.push_back(out.shape().back()); - array intermediate(intermediate_shape, q.dtype(), nullptr, {}); + // Keep unnormalized partials in float until the final reduction. + array intermediate(intermediate_shape, float32, nullptr, {}); intermediate_shape.pop_back(); array sums(intermediate_shape, float32, nullptr, {}); array maxs(std::move(intermediate_shape), float32, nullptr, {}); @@ -566,7 +608,7 @@ compute_encoder.dispatch_threadgroups(grid_dims, group_dims);   // Final pass kname.clear(); - kname = "sdpa_vector_2pass_2_"; + kname = "sdpa_vector_2pass_fp32partials_2_"; kname += get_type_string(q.dtype()); kname += "_"; kname += std::to_string(v.shape(-1)); @@ -589,28 +631,23 @@ check_kernel_threadgroup_size(kernel, group_dims, kname); compute_encoder.dispatch_threadgroups(grid_dims, group_dims); }   -} // namespace - -bool ScaledDotProductAttention::use_fallback( +std::tuple<bool, std::string> has_fused_kernel( const array& q, const array& k, const array& v, bool has_mask, bool has_arr_mask, bool do_causal, - bool is_training, bool output_logsumexp, Stream s) { - if (is_training) { - // It's faster for training on Metal to use the unfused SDPA for both - // forward and backward. - return true; + if (s.device != Device::gpu) { + return {false, "the fused kernels require a GPU (Metal) stream."}; } if (output_logsumexp) { - return true; - } - if (s.device == Device::cpu) { - return true; + return { + false, + "the fused forward does not produce the logsumexp required for " + "the fused VJP; use default routing when training."}; }   const int value_head_dim = v.shape(-1); @@ -621,27 +658,113 @@ const int num_query_heads = q.shape(1); const int num_kv_heads = k.shape(1); const int gqa_factor = num_query_heads / num_kv_heads;   - const bool sdpa_vector_supported_head_dim = - (query_head_dim == value_head_dim && - (query_head_dim == 64 || query_head_dim == 96 || query_head_dim == 128 || - query_head_dim == 256)) || - (query_head_dim == 192 && value_head_dim == 128); - const bool sdpa_full_supported_head_dim = query_head_dim == value_head_dim && - (query_head_dim == 64 || query_head_dim == 80 || query_head_dim == 96 || - query_head_dim == 128); + std::ostringstream msg; + if (query_sequence_length > 8) { + const bool supported_head_dim = query_head_dim == value_head_dim && + (query_head_dim == 64 || query_head_dim == 72 || query_head_dim == 80 || + query_head_dim == 96 || query_head_dim == 128 || + query_head_dim == 192 || query_head_dim == 256); + if (!supported_head_dim) { + msg << "the full attention kernel supports head dims " + << "{64, 72, 80, 96, 128, 192, 256} with matching query/value head " + << "dims; got query head dim " << query_head_dim + << " and value head dim " << value_head_dim << "."; + return {false, msg.str()}; + } + if (has_mask && !has_arr_mask && + !(query_sequence_length <= key_sequence_length && do_causal)) { + msg << "the full attention kernel with a causal mask requires the " + << "query sequence to be no longer than the key sequence; got " + << "query length " << query_sequence_length << " and key length " + << key_sequence_length << "."; + return {false, msg.str()}; + } + } else { + const bool supported_head_dim = + (query_head_dim == value_head_dim && + (query_head_dim == 64 || query_head_dim == 96 || + query_head_dim == 128 || query_head_dim == 192 || + query_head_dim == 256)) || + (query_head_dim == 192 && value_head_dim == 128); + if (!supported_head_dim) { + msg << "the vector attention kernel supports head dims " + << "{64, 96, 128, 192, 256} with matching query/value head dims, " + << "or query head dim 192 with value head dim 128; got query head " + << "dim " << query_head_dim << " and value head dim " + << value_head_dim << "."; + return {false, msg.str()}; + } + if (query_sequence_length > key_sequence_length) { + msg << "the vector attention kernel requires the query sequence to be " + << "no longer than the key sequence; got query length " + << query_sequence_length << " and key length " << key_sequence_length + << "."; + return {false, msg.str()}; + } + if (query_sequence_length * gqa_factor > 32) { + msg << "the vector attention kernel requires the query length times " + << "the GQA factor to be at most 32; got query length " + << query_sequence_length << " and GQA factor " << gqa_factor << "."; + return {false, msg.str()}; + } + } + return {true, ""}; +}   - const bool sdpa_full_supported_mask = !has_mask || has_arr_mask || - (query_sequence_length <= key_sequence_length && do_causal); +} // namespace   - const bool supports_sdpa_full = query_sequence_length > 8 && - sdpa_full_supported_mask && sdpa_full_supported_head_dim; +bool ScaledDotProductAttention::use_fallback( + const array& q, + const array& k, + const array& v, + bool has_mask, + bool has_arr_mask, + bool do_causal, + bool is_training, + bool output_logsumexp, + bool force_fused, + Stream s) { + auto [has_fused, reason] = has_fused_kernel( + q, k, v, has_mask, has_arr_mask, do_causal, output_logsumexp, s); + if (force_fused) { + if (!has_fused) { + std::ostringstream msg; + msg << "[scaled_dot_product_attention] force_fused=True but no fused " + "kernel is available: " + << reason; + throw std::invalid_argument(msg.str()); + } + return false; + }   - const bool supports_sdpa_vector = (query_sequence_length <= 8) && - (query_sequence_length <= key_sequence_length) && - sdpa_vector_supported_head_dim && - (query_sequence_length * gqa_factor) <= 32; + if (is_training) { + // It's faster for training on Metal to use the unfused SDPA for both + // forward and backward. + return true; + } + if (!has_fused) { + return true; + }   - return !(supports_sdpa_full || supports_sdpa_vector); + const int query_sequence_length = q.shape(2); + const int query_head_dim = q.shape(-1); + const int value_head_dim = v.shape(-1); + + // Use headdim-split kernel when NAX is enabled and there are enough query + // blocks to fill the machine. + if (metal::is_nax_available() && + (env::enable_tf32() || q.dtype() != float32) && + query_sequence_length >= 1024 && query_head_dim == 256 && do_causal && + !has_arr_mask) { + return false; + } + + // Unfused path is faster for following shapes. + if (query_sequence_length > 8) { + return query_head_dim == 192 || query_head_dim == 256; + } else { + return query_head_dim == value_head_dim && query_head_dim == 192; + } }   bool ScaledDotProductAttention::supports_bool_mask() {
diff --git ml-explore/mlx/python/tests/test_fast_sdpa.py Layr-Labs/mlx/python/tests/test_fast_sdpa.py index 2418997bc892a66b592b8cb4c964a796dcaf1c05..3a9c04b075372b9f7f40d3be9835428df6b9d09a 100644 --- ml-explore/mlx/python/tests/test_fast_sdpa.py +++ Layr-Labs/mlx/python/tests/test_fast_sdpa.py @@ -2,6 +2,7 @@ import math import os import unittest from itertools import product +from unittest.mock import patch   import mlx.core as mx import mlx_tests @@ -118,6 +119,33 @@   class TestFastSDPA(mlx_tests.MLXTestCase): @unittest.skipIf(not mx.is_available(mx.gpu), "GPU kernel path only") + def test_sdpa_head_dim_72(self): + B, D, qH, kH = (1, 72, 8, 2) + for qL, kL, dtype, mask_str in product( + (64, 65), + (128, 127), + (mx.float16, mx.bfloat16, mx.float32), + (None, "additive", "bool", "causal"), + ): + with self.subTest(qL=qL, kL=kL, dtype=dtype, mask=mask_str): + q, k, v, scale, mask = prepare_inputs( + B, qL, kL, D, qH, kH, mask_str, False, dtype + ) + ref = mlx_ref_attn(q, k, v, scale, mask) + out = mx.fast.scaled_dot_product_attention( + q, k, v, scale=scale, mask=mask + ) + + if dtype == mx.float32: + atol = 1e-5 + elif dtype == mx.bfloat16: + atol = 5e-3 + else: + atol = 3e-4 + diff = mx.abs(out - ref) - atol * mx.abs(ref) + self.assertLessEqual(mx.max(diff).item(), atol) + + @unittest.skipIf(not mx.is_available(mx.gpu), "GPU kernel path only") def test_sdpa_head_dim_96(self): B, D, qH, kH = (1, 96, 8, 2) for qL, kL, dtype, mask_str in product( @@ -144,6 +172,73 @@ atol = 3e-4 diff = mx.abs(out - ref) - atol * mx.abs(ref) self.assertLessEqual(mx.max(diff).item(), atol)   + @unittest.skipIf(not mx.is_available(mx.gpu), "GPU kernel path only") + def test_sdpa_full_head_dim_256(self): + # On NAX devices, large nearly-square causal blocks take the fused + # path; everything else takes the unfused fallback. Ragged lengths + # exercise the kernel's unaligned pipelines, and K/V sliced out of a + # longer preallocated cache (the way mlx-lm hands them over) exercise + # the dispatch reading the slice past its end. All of it must be + # correct. + D = 256 + Nq, Nkv = 8, 2 + scale = D**-0.5 + mx.random.seed(0) + cases = [ + # fused on NAX: aligned square, ragged square (unaligned Q and + # K/V), aligned rectangle at the routing boundary, ragged + # rectangle + (2048, 2048, "causal", None), + (2049, 2049, "causal", None), + (2048, 2560, "causal", None), + (2049, 2560, "causal", None), + # fused on NAX with ragged K/V sliced out of a longer cache: the + # rows behind the slice must not leak into the output + (2049, 2049, "causal", 2304), + (2048, 2500, "causal", 2560), + (1031, 2049, "causal", 2304), + ] + for dtype in (mx.float32, mx.bfloat16): + for qL, kL, mask, cache_len in cases: + with self.subTest( + dtype=dtype, qL=qL, kL=kL, mask=mask, cache_len=cache_len + ): + q = (5e-1 * mx.random.normal(shape=(1, Nq, qL, D))).astype(dtype) + if cache_len is None: + k = (5e-1 * mx.random.normal(shape=(1, Nkv, kL, D))).astype( + dtype + ) + v = (5e-1 * mx.random.normal(shape=(1, Nkv, kL, D))).astype( + dtype + ) + else: + # Large, finite stale rows behind the slice: any of + # them reaching the output is loud. + k_cache = 1e2 * mx.random.normal(shape=(1, Nkv, cache_len, D)) + v_cache = 1e3 * mx.random.normal(shape=(1, Nkv, cache_len, D)) + k_cache[..., :kL, :] = 5e-1 * mx.random.normal( + shape=(1, Nkv, kL, D) + ) + v_cache[..., :kL, :] = 5e-1 * mx.random.normal( + shape=(1, Nkv, kL, D) + ) + k = k_cache.astype(dtype)[..., :kL, :] + v = v_cache.astype(dtype)[..., :kL, :] + k_rep = mx.repeat(k, Nq // Nkv, axis=1) + v_rep = mx.repeat(v, Nq // Nkv, axis=1) + ref = mlx_primitives_sdpa(q, k_rep, v_rep, scale, mask=mask) + out = mx.fast.scaled_dot_product_attention( + q, k, v, scale=scale, mask=mask + ) + self.assertEqual(out.shape, ref.shape) + if dtype == mx.float32: + # The fused shapes run through tf32 tensor ops when + # MLX_ENABLE_TF32 is on (the default). + tol = 1e-3 if qL >= 2048 else 1e-4 + else: + tol = 5e-3 + self.assertTrue(mx.allclose(ref, out, atol=tol, rtol=tol)) + def test_sdpa_vector_kv_transposed_head_seq(self): D = 64 Nq = 4 @@ -248,6 +343,68 @@ mask=m, ) self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4))   + def test_sdpa_vector_gqa_long(self): + scale = 1.0 + mx.random.seed(0) + for Nq, Nkv, D in [(32, 4, 128), (64, 8, 64)]: + for L in [8192, 8201]: + q = 5e-1 * mx.random.normal(shape=(1, Nq, 1, D)) + k = 5e-1 * mx.random.normal(shape=(1, Nkv, L + 32, D))[:, :, :L] + v = 5e-1 * mx.random.normal(shape=(1, Nkv, L + 32, D))[:, :, :L] + kr = mx.repeat(k, Nq // Nkv, axis=1) + vr = mx.repeat(v, Nq // Nkv, axis=1) + ref = mlx_primitives_sdpa(q, kr, vr, scale) + out = mx.fast.scaled_dot_product_attention(q, k, v, scale=scale) + self.assertTrue(mx.allclose(ref, out, atol=1e-4, rtol=1e-4)) + + @unittest.skipIf(not mx.is_available(mx.gpu), "GPU kernel path only") + def test_sdpa_two_pass_partial_cancellation(self): + # Uniform attention has the exact output 1 / L. Casting an + # unnormalized block sum to the input dtype loses the residual 1. + L = 8192 + for dtype, D, blocks, masked in product( + (mx.bfloat16, mx.float16, mx.float32), + (64, 128, 256), + (32, 128), + (False, True), + ): + with self.subTest(dtype=dtype, D=D, blocks=blocks, masked=masked): + amplitude = 2048 if dtype == mx.float16 else 256 + values = np.zeros((1, 1, L, D), dtype=np.float32) + values[:, :, 0, :] = amplitude + # D64/D128 without a mask uses contiguous GQA chunks. + gqa = D in (64, 128) and not masked + values[:, :, 1 if gqa else blocks, :] = 1 + values[:, :, L - 1 if gqa else 1, :] = -amplitude + q = mx.zeros((1, 8, 1, D), dtype=dtype) + k = mx.zeros((1, 1, L, D), dtype=dtype) + v = mx.array(values, dtype=dtype) + mask = mx.ones((L,), dtype=mx.bool_) if masked else None + with patch.dict(os.environ, MLX_SDPA_BLOCKS=str(blocks)): + out = mx.fast.scaled_dot_product_attention( + q, k, v, scale=1, mask=mask + ) + mx.eval(out) + self.assertEqual(out.dtype, dtype) + self.assertTrue(mx.array_equal(out, mx.full(out.shape, 1 / L, dtype))) + + @unittest.skipIf(not mx.is_available(mx.gpu), "GPU kernel path only") + def test_sdpa_two_pass_partial_overflow(self): + # The final mean is representable in FP16, but its partial sums are not. + L = 8192 + for D, blocks in product((64, 128, 256), (32, 128)): + with self.subTest(D=D, blocks=blocks): + q = mx.zeros((1, 8, 1, D), dtype=mx.float16) + k = mx.zeros((1, 1, L, D), dtype=mx.float16) + v = mx.full((1, 1, L, D), 32768, dtype=mx.float16) + with patch.dict(os.environ, MLX_SDPA_BLOCKS=str(blocks)): + out = mx.fast.scaled_dot_product_attention(q, k, v, scale=1) + mx.eval(out) + self.assertTrue(mx.all(mx.isfinite(out))) + self.assertTrue( + mx.array_equal(out, mx.full(out.shape, 32768, mx.float16)) + ) + def test_sdpa_fully_masked(self): Lkv = 8 mask = mx.array(False) @@ -694,6 +851,100 @@ loss_fast = lambda q, k, v: mx.fast.scaled_dot_product_attention( q, k, v, scale=scale, mask=mask ).sum() test_grad(loss_slow, loss_fast, [q, k, v]) + + @unittest.skipIf(not mx.metal.is_available(), "Metal kernel path only") + def test_sdpa_force_fused_metal(self): + if mx.default_device() != mx.gpu: + self.skipTest("requires GPU") + + def make_qkv(qL, kL, D, qH=8, kH=8): + q = mx.random.normal((1, qH, qL, D), mx.float16) + k = mx.random.normal((1, kH, kL, D), mx.float16) + v = mx.random.normal((1, kH, kL, D), mx.float16) + return q, k, v + + # Full attention kernel. + for D, qL, mask in product((192, 256), (9, 16), (None, "causal")): + with self.subTest(head_dim=D, qL=qL, mask=mask): + q, k, v = make_qkv(qL, 512, D, 8, 4) + scale = D**-0.5 + ref = mlx_ref_attn(q, k, v, scale=scale, mask=mask) + out = mx.fast.scaled_dot_product_attention( + q, k, v, scale=scale, mask=mask, force_fused=True + ) + self.assertTrue(mx.allclose(ref, out, atol=1e-3, rtol=1e-3)) + + # Vector attention kernel. + for D in (192, 256): + with self.subTest(head_dim=D): + q, k, v = make_qkv(4, 16385, D, 4, 2) + scale = D**-0.5 + ref = mlx_ref_attn(q, k, v, scale=scale) + out = mx.fast.scaled_dot_product_attention( + q, k, v, scale=scale, force_fused=True + ) + self.assertTrue(mx.allclose(ref, out, atol=1e-3, rtol=1e-3)) + + # No full attention fused kernels. + with self.assertRaisesRegex(ValueError, "supports head dims"): + q, k, v = make_qkv(16, 512, 512) + mx.fast.scaled_dot_product_attention( + q, k, v, scale=512**-0.5, force_fused=True + ) + with self.assertRaisesRegex( + ValueError, "query sequence to be no longer than the key sequence" + ): + q, k, v = make_qkv(32, 16, 64) + mx.fast.scaled_dot_product_attention( + q, + k, + v, + scale=64**-0.5, + mask="causal", + force_fused=True, + ) + + # No vector attention fused kernels. + with self.assertRaisesRegex(ValueError, "supports head dims"): + q, k, v = make_qkv(1, 128, 72) + mx.fast.scaled_dot_product_attention( + q, k, v, scale=72**-0.5, force_fused=True + ) + with self.assertRaisesRegex(ValueError, "GQA factor to be at most 32"): + q, k, v = make_qkv(8, 128, 64, qH=8, kH=1) + mx.fast.scaled_dot_product_attention( + q, k, v, scale=64**-0.5, force_fused=True + ) + + # No CPU fused kernel. + with mx.stream(mx.cpu): + q, k, v = make_qkv(8, 128, 8) + with self.assertRaisesRegex(ValueError, "require a GPU"): + mx.fast.scaled_dot_product_attention( + q, k, v, scale=64**-0.5, force_fused=True + ) + + @unittest.skipIf(not mx.cuda.is_available(), "CUDA kernel path only") + def test_sdpa_force_fused_cuda(self): + if mx.default_device() != mx.gpu: + self.skipTest("requires GPU") + + def make_qkv(qL, kL, D, qH=8, kH=8): + q = mx.random.normal((1, qH, qL, D), mx.float16) + k = mx.random.normal((1, kH, kL, D), mx.float16) + v = mx.random.normal((1, kH, kL, D), mx.float16) + return q, k, v + + # Vector attention kernel. + for D in (64, 96, 128): + with self.subTest(head_dim=D): + q, k, v = make_qkv(3, 128, D, 4, 2) + scale = D**-0.5 + ref = mlx_ref_attn(q, k, v, scale=scale) + out = mx.fast.scaled_dot_product_attention( + q, k, v, scale=scale, force_fused=True + ) + self.assertTrue(mx.allclose(ref, out, atol=1e-3, rtol=1e-3))   def test_sdpa_sliced(self): N = 8

metal_kernel_with_mutable_inputs lets a caller declare which custom-kernel inputs will be written, so paged-KV consumers keep writes on the original allocation and Metal sees the read/write hazard. Mutable names must identify exactly one input; mutable device pointers are generated (including for small/scalar storage) and take part in the compiled-kernel identity; an implicit contiguous copy of declared-mutable storage is rejected; allocations are registered for writes as well as reads. A separate MutableInputCustomKernel primitive serializes the declaration while the existing CustomKernel export schema is untouched. (PR #18)

diff --git ml-explore/mlx/mlx/backend/common/metal_kernel.cpp Layr-Labs/mlx/mlx/backend/common/metal_kernel.cpp index 7a2eef51fcfa04d1e6828ca72be100235fded40f..0efb605a3b08b9eac4ecc69db9f36fe8921cf890 100644 --- ml-explore/mlx/mlx/backend/common/metal_kernel.cpp +++ Layr-Labs/mlx/mlx/backend/common/metal_kernel.cpp @@ -1,5 +1,6 @@ // Copyright © 2024 Apple Inc.   +#include <algorithm> #include <iostream> #include <sstream>   @@ -60,7 +61,8 @@ const std::vector<Dtype>& output_dtypes, const std::vector<std::pair<std::string, TemplateArg>>& template_args, const std::vector<std::string>& attributes, const std::vector<std::tuple<bool, bool, bool>>& shape_infos, - bool atomic_outputs) { + bool atomic_outputs, + const std::vector<int>& mutable_inputs) { std::string kernel_source; kernel_source.reserve(header.size() + source.size() + 16384); kernel_source += header; @@ -97,11 +99,15 @@ // Add inputs for (int i = 0; i < inputs.size(); ++i) { const auto& name = input_names[i]; const auto& arr = inputs[i]; + bool is_mutable = + std::find(mutable_inputs.begin(), mutable_inputs.end(), i) != + mutable_inputs.end(); auto dtype = get_type_string(arr.dtype()); - std::string location = - arr.size() < max_constant_array_size ? "constant" : "device"; - std::string ref = arr.ndim() == 0 ? "&" : "*"; - kernel_source += " const "; + std::string location = !is_mutable && arr.size() < max_constant_array_size + ? "constant" + : "device"; + std::string ref = !is_mutable && arr.ndim() == 0 ? "&" : "*"; + kernel_source += is_mutable ? " " : " const "; kernel_source += location; kernel_source += " "; kernel_source += dtype; @@ -228,6 +234,43 @@ const std::string& header /* = "" */, bool ensure_row_contiguous /* = true */, bool atomic_outputs /* = false */, const CompileOptions& compile_options /* = {} */) { + return metal_kernel_with_mutable_inputs( + name, + input_names, + output_names, + source, + {}, + header, + ensure_row_contiguous, + atomic_outputs, + compile_options); +} + +CustomKernelFunction metal_kernel_with_mutable_inputs( + const std::string& name, + const std::vector<std::string>& input_names, + const std::vector<std::string>& output_names, + const std::string& source, + const std::vector<std::string>& mutable_input_names, + const std::string& header, + bool ensure_row_contiguous, + bool atomic_outputs, + const CompileOptions& compile_options) { + std::vector<int> mutable_inputs; + for (const auto& mutable_name : mutable_input_names) { + auto it = std::find(input_names.begin(), input_names.end(), mutable_name); + if (it == input_names.end() || + std::count(input_names.begin(), input_names.end(), mutable_name) != 1) { + throw std::invalid_argument("Mutable input must name one kernel input."); + } + int index = static_cast<int>(it - input_names.begin()); + if (std::find(mutable_inputs.begin(), mutable_inputs.end(), index) != + mutable_inputs.end()) { + throw std::invalid_argument("Duplicate mutable kernel input."); + } + mutable_inputs.push_back(index); + } + std::sort(mutable_inputs.begin(), mutable_inputs.end()); if (output_names.empty()) { throw std::invalid_argument( "[metal_kernel] Must specify at least one output."); @@ -310,10 +353,14 @@ // 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) { + for (int i = 0; i < inputs.size(); ++i) { + const auto& arr = inputs[i]; kernel_name += "_"; kernel_name += get_type_string(arr.dtype()); - if (arr.ndim() == 0) { + if (std::find(mutable_inputs.begin(), mutable_inputs.end(), i) != + mutable_inputs.end()) { + kernel_name += "m"; + } else if (arr.ndim() == 0) { kernel_name += "s"; } else if (arr.size() < max_constant_array_size) { kernel_name += "c"; @@ -335,7 +382,8 @@ output_dtypes, template_args, attributes, shape_infos, - atomic_outputs); + atomic_outputs, + mutable_inputs);   if (!template_args.empty()) { template_def = kernel_name + template_def; @@ -355,6 +403,26 @@ << kernel_source << std::endl << "```" << std::endl; }   + if (!mutable_inputs.empty()) { + return array::make_arrays( + output_shapes, + output_dtypes, + std::make_shared<MutableInputCustomKernel>( + s, + kernel_name, + kernel_source, + grid, + threadgroup, + shape_infos, + ensure_row_contiguous, + init_value, + std::vector<ScalarArg>{}, + false, + 0, + compile_options.serialize(), + mutable_inputs), + inputs); + } return array::make_arrays( std::move(output_shapes), std::move(output_dtypes),
diff --git ml-explore/mlx/mlx/backend/metal/custom_kernel.cpp Layr-Labs/mlx/mlx/backend/metal/custom_kernel.cpp index b73dd72def76fa3412ef9155749507fe5d5cb48e..65c33d2448e62e52eaf48c5d28426f784d7dc680 100644 --- ml-explore/mlx/mlx/backend/metal/custom_kernel.cpp +++ Layr-Labs/mlx/mlx/backend/metal/custom_kernel.cpp @@ -7,6 +7,7 @@ #include "mlx/backend/metal/utils.h" #include "mlx/fast_primitives.h"   #include <fmt/format.h> +#include <algorithm>   namespace mlx::core::fast {   @@ -18,6 +19,16 @@ (void)is_precompiled_; (void)shared_memory_;   auto& s = stream(); + + for (auto i : mutable_inputs_) { + if (i < 0 || i >= inputs.size()) { + throw std::invalid_argument("Invalid mutable kernel input index."); + } + if (ensure_row_contiguous_ && !inputs[i].flags().row_contiguous) { + throw std::invalid_argument( + "Mutable kernel inputs cannot use an implicit contiguous copy."); + } + }   std::vector<array> copies;   @@ -59,6 +70,10 @@ for (int i = 0; i < checked_inputs.size(); i++) { const array& in = checked_inputs[i]; auto& shape_info = shape_infos_[i]; compute_encoder.set_input_array(in, index); + if (std::find(mutable_inputs_.begin(), mutable_inputs_.end(), i) != + mutable_inputs_.end()) { + compute_encoder.register_output_array(in); + } index++; if (in.ndim() > 0) { int ndim = in.ndim();
diff --git ml-explore/mlx/mlx/export.cpp Layr-Labs/mlx/mlx/export.cpp index bcc0ad161c25edfd522bbac9b14c9c847423c3fa..06ff99a7b59f127809244086a58225b6d4a8fb49 100644 --- ml-explore/mlx/mlx/export.cpp +++ Layr-Labs/mlx/mlx/export.cpp @@ -464,7 +464,8 @@ SERIALIZE_PRIMITIVE(LayerNorm), SERIALIZE_PRIMITIVE(LayerNormVJP), SERIALIZE_PRIMITIVE(RoPE), SERIALIZE_PRIMITIVE(ScaledDotProductAttention), - SERIALIZE_PRIMITIVE(CustomKernel)}; + SERIALIZE_PRIMITIVE(CustomKernel), + SERIALIZE_PRIMITIVE(MutableInputCustomKernel)}; std::unordered_map<std::string, std::string> name_remap; std::unordered_map<int, Stream> stream_map;
diff --git ml-explore/mlx/mlx/fast.h Layr-Labs/mlx/mlx/fast.h index 934fadc2b7c4f18cf7f711c2b291517aa3324563..aca6da3aed9092f490cc30499e98d8c0be842f30 100644 --- ml-explore/mlx/mlx/fast.h +++ Layr-Labs/mlx/mlx/fast.h @@ -24,6 +24,10 @@ const std::optional<array>& bias, float eps, StreamOrDevice s = {});   +/** Fused cross entropy with class indices as targets. */ +MLX_API array +cross_entropy(const array& logits, const array& targets, StreamOrDevice s = {}); + MLX_API array rope( const array& x, int dims, @@ -53,6 +57,7 @@ const float scale, const std::string& mask_mode = "", std::optional<array> mask_arr = {}, const std::optional<array>& sinks = {}, + bool force_fused = false, StreamOrDevice s = {});   using TemplateArg = std::variant<int, bool, Dtype>; @@ -74,6 +79,17 @@ const std::string& name, const std::vector<std::string>& input_names, const std::vector<std::string>& output_names, const std::string& source, + const std::string& header = "", + bool ensure_row_contiguous = true, + bool atomic_outputs = false, + const CompileOptions& compile_options = {}); + +MLX_API CustomKernelFunction metal_kernel_with_mutable_inputs( + const std::string& name, + const std::vector<std::string>& input_names, + const std::vector<std::string>& output_names, + const std::string& source, + const std::vector<std::string>& mutable_input_names, const std::string& header = "", bool ensure_row_contiguous = true, bool atomic_outputs = false,
diff --git ml-explore/mlx/mlx/fast_primitives.h Layr-Labs/mlx/mlx/fast_primitives.h index 0d2f86104592f73d4cf47b7450339be7ef392924..0a1a945cddabed37b37e8b24740766a7856d2eec 100644 --- ml-explore/mlx/mlx/fast_primitives.h +++ Layr-Labs/mlx/mlx/fast_primitives.h @@ -97,6 +97,67 @@ private: float eps_; };   +// loss is always fp32 and the logits never have to be upcast in the graph. +class CrossEntropy : public Custom { + public: + CrossEntropy( + Stream stream, + std::function<std::vector<array>(std::vector<array>)> fallback) + : Custom(stream, std::move(fallback)) {} + + static bool use_fallback(Stream stream); + + void eval_cpu(const std::vector<array>& inputs, std::vector<array>& outputs) + override { + throw std::runtime_error("NYI"); + } + void eval_gpu(const std::vector<array>& inputs, std::vector<array>& outputs) + override; + + std::vector<array> vjp( + const std::vector<array>& primals, + const std::vector<array>& cotangents, + const std::vector<int>& argnums, + const std::vector<array>& outputs) override; + + DEFINE_NAME(CrossEntropy) + bool is_equivalent(const Primitive& other) const override { + return true; + } + std::vector<Shape> output_shapes(const std::vector<array>& inputs) override { + return {inputs[1].shape()}; + } + + auto state() const { + return std::monostate{}; + } +}; + +class CrossEntropyVJP : public Custom { + public: + CrossEntropyVJP( + Stream stream, + std::function<std::vector<array>(std::vector<array>)> fallback) + : Custom(stream, std::move(fallback)) {} + + void eval_cpu(const std::vector<array>& inputs, std::vector<array>& outputs) + override { + throw std::runtime_error("NYI"); + } + void eval_gpu(const std::vector<array>& inputs, std::vector<array>& outputs) + override; + + DEFINE_NAME(CrossEntropyVJP) + bool is_equivalent(const Primitive& other) const override { + return true; + } + DEFINE_INPUT_OUTPUT_SHAPE() + + auto state() const { + return std::monostate{}; + } +}; + class LayerNorm : public Custom { public: LayerNorm( @@ -212,12 +273,14 @@ std::function<std::vector<array>(std::vector<array>)> fallback, float scale, bool do_causal, bool has_sinks, - bool output_logsumexp) + bool output_logsumexp, + bool force_fused) : Custom(stream, std::move(fallback)), scale_(scale), do_causal_(do_causal), has_sinks_(has_sinks), - output_logsumexp_(output_logsumexp) {} + output_logsumexp_(output_logsumexp), + force_fused_(force_fused) {}   static bool use_fallback( const array& q, @@ -228,6 +291,7 @@ bool has_arr_mask, bool do_causal, bool is_training, bool output_logsumexp, + bool force_fused, Stream s); static bool supports_bool_mask();   @@ -251,7 +315,12 @@ DEFINE_NAME(ScaledDotProductAttention); DEFINE_INPUT_OUTPUT_SHAPE() auto state() const { return std::make_tuple( - nullptr, scale_, do_causal_, has_sinks_, output_logsumexp_); + nullptr, + scale_, + do_causal_, + has_sinks_, + output_logsumexp_, + force_fused_); }   private: @@ -259,6 +328,7 @@ float scale_; bool do_causal_; bool has_sinks_; bool output_logsumexp_; + bool force_fused_; };   class ScaledDotProductAttentionVJP : public Custom { @@ -377,7 +447,8 @@ std::optional<float> init_value, std::vector<ScalarArg> scalar_arguments, bool is_precompiled, int shared_memory, - CompileOptions::Data compile_options = {}) + CompileOptions::Data compile_options = {}, + std::vector<int> mutable_inputs = {}) : Primitive(stream), name_(std::move(name)), source_(std::move(source)), @@ -389,7 +460,8 @@ init_value_(init_value), scalar_arguments_(std::move(scalar_arguments)), is_precompiled_(is_precompiled), shared_memory_(shared_memory), - compile_options_(compile_options) {} + compile_options_(compile_options), + mutable_inputs_(std::move(mutable_inputs)) {}   void eval_cpu(const std::vector<array>& inputs, std::vector<array>& outputs) override { @@ -427,6 +499,19 @@ std::vector<ScalarArg> scalar_arguments_; bool is_precompiled_; int shared_memory_; CompileOptions::Data compile_options_; + + protected: + std::vector<int> mutable_inputs_; +}; + +class MutableInputCustomKernel : public CustomKernel { + public: + using CustomKernel::CustomKernel; + DEFINE_NAME(MutableInputCustomKernel); + auto state() const { + return std::tuple_cat( + CustomKernel::state(), std::make_tuple(mutable_inputs_)); + } };   } // namespace mlx::core::fast

Compile caches became thread-local upstream; bindings built before that (the mlx-c / mlx-swift consumers) still call the process-wide compile_erase / compile_clear_cache. A process-lifetime registry records each thread’s cache once at construction so those entrypoints can erase or clear every live cache, with no per-call hot-path work. (PR #10)

diff --git ml-explore/mlx/mlx/compile.cpp Layr-Labs/mlx/mlx/compile.cpp index 12d7397be4f3975efca57203bd32755d6e0caa3c..8a58c890bf1ccef89a3fa66f0d60fe2d7b0adf77 100644 --- ml-explore/mlx/mlx/compile.cpp +++ Layr-Labs/mlx/mlx/compile.cpp @@ -3,6 +3,7 @@ #include <atomic> #include <cstdlib> #include <map> +#include <shared_mutex> #include <sstream> #include <unordered_map> #include <unordered_set> @@ -214,6 +215,27 @@ }   namespace detail {   +class CompileCache; + +class CompileCacheRegistry { + public: + void add(CompileCache* cache); + void remove(CompileCache* cache); + void erase(std::uintptr_t fun_id); + void clear(); + + private: + std::mutex mutex_; + std::unordered_set<CompileCache*> caches_; +}; + +CompileCacheRegistry& compile_cache_registry() { + // Intentionally process-lifetime: thread-local compile caches can be torn + // down after ordinary static destruction has started. + static auto* registry = new CompileCacheRegistry(); + return *registry; +} + std::atomic<CompileMode>& compile_mode() { auto get_val = []() { if (std::getenv("MLX_DISABLE_COMPILE")) { @@ -298,7 +320,7 @@ } return reinterpret_cast<std::uintptr_t>(*fun_ptr); }   -class CompilerCache { +class CompileCache { public: struct CacheEntry { CacheEntry(Stream stream, bool shapeless) @@ -313,15 +335,43 @@ std::vector<uint64_t> constants; std::shared_ptr<void> extra; };   - // Returns a reference to a CacheEntry which can be updated - // by the caller to avoid copying large tapes / inputs / outputs - CacheEntry& find( + CompileCache() { + // Make sure the allocator is fully initialized before the compiler cache. + allocator::allocator(); + compile_cache_registry().add(this); + } + + ~CompileCache() { + compile_cache_registry().remove(this); + } + + // Returns a reference to a CacheEntry which can be updated by the caller to + // avoid copying large tapes / inputs / outputs, with the shared_ptr of + // entries to avoid getting erased during compilation. + std::tuple<CacheEntry&, std::shared_ptr<std::vector<CacheEntry>>> find( std::uintptr_t fun_id, const std::vector<array>& inputs, bool shapeless, const std::vector<uint64_t>& constants) { - // Find the cache entries for |fun_id|. - std::vector<CacheEntry>& entries = cache_[fun_id]; + // Find the cache entries for |fun_id| in a thread-safe way. + auto entries_ptr = [&]() { + // Lookup with shared lock. + { + std::shared_lock lock(mutex_); + auto it = cache_.find(fun_id); + if (it != cache_.end()) { + return it->second; + } + } + // Insertion with exclusive lock. + std::unique_lock lock(mutex_); + auto& ptr = cache_[fun_id]; + if (!ptr) { + ptr = std::make_shared<std::vector<CacheEntry>>(); + } + return ptr; + }(); + auto& entries = *entries_ptr;   // Compare if 2 arrays have same shape and dtype. auto has_same_shape_and_dtype = [shapeless]( @@ -359,40 +409,61 @@ // Check the inputs match and return if so if (has_same_shape_and_dtype(inputs, entry.inputs) && constants == entry.constants) { - return entry; + return {entry, std::move(entries_ptr)}; } } // Otherwise append a new cache entry entries.push_back(CacheEntry{stream, shapeless}); - return entries.back(); + return {entries.back(), std::move(entries_ptr)}; }   void erase(std::uintptr_t fun_id) { + std::unique_lock lock(mutex_); cache_.erase(fun_id); }   void clear() { + std::unique_lock lock(mutex_); cache_.clear(); }   - bool empty() { - return cache_.empty(); - } + private: + // The cache may get its key erased from a separate thread, but its value is + // only added and modified in the thread of creation. + // Put value in a shared_ptr to avoid race condition when erasing happened + // during compilation for the same function. + std::unordered_map<std::uintptr_t, std::shared_ptr<std::vector<CacheEntry>>> + cache_; + std::shared_mutex mutex_; +}; + +void CompileCacheRegistry::add(CompileCache* cache) { + std::unique_lock lock(mutex_); + caches_.insert(cache); +} + +void CompileCacheRegistry::remove(CompileCache* cache) { + std::unique_lock lock(mutex_); + caches_.erase(cache); +}   - private: - CompilerCache() { - // Make sure the allocator is fully - // initialized before the compiler cache - allocator::allocator(); +void CompileCacheRegistry::erase(std::uintptr_t fun_id) { + std::unique_lock lock(mutex_); + for (auto* cache : caches_) { + cache->erase(fun_id); } +}   - friend CompilerCache& compiler_cache(); - std::unordered_map<std::uintptr_t, std::vector<CacheEntry>> cache_; -}; +void CompileCacheRegistry::clear() { + std::unique_lock lock(mutex_); + for (auto* cache : caches_) { + cache->clear(); + } +}   -CompilerCache& compiler_cache() { - static thread_local CompilerCache compiler_cache_; - return compiler_cache_; +std::shared_ptr<CompileCache>& compile_cache_unsafe() { + static thread_local auto cache = std::make_shared<CompileCache>(); + return cache; }   std::tuple<std::vector<array>, std::vector<array>, std::shared_ptr<void>> @@ -1120,7 +1191,9 @@ return fun(inputs); }   // Find a cache entry with the correct inputs - auto& entry = compiler_cache().find(fun_id, inputs, shapeless, constants); + auto [entry, entries_ptr] = + compile_cache_unsafe()->find(fun_id, inputs, shapeless, constants); + static_assert(std::is_reference_v<decltype(entry)>);   // No matching cache entry existed, so compile if (entry.empty) { @@ -1192,16 +1265,28 @@ return compiled_fun(inputs).first; }; }   +CompileCacheWeakPtr compile_cache() { + return compile_cache_unsafe(); +} + +void compile_erase(const CompileCacheWeakPtr& cache, std::uintptr_t fun_id) { + if (auto p = cache.lock()) { + p->erase(fun_id); + } +} + void compile_erase(std::uintptr_t fun_id) { - detail::compiler_cache().erase(fun_id); + compile_cache_registry().erase(fun_id); }   -void compile_clear_cache() { - detail::compiler_cache().clear(); +void compile_clear_cache(const CompileCacheWeakPtr& cache) { + if (auto p = cache.lock()) { + p->clear(); + } }   -bool compile_cache_empty() { - return detail::compiler_cache().empty(); +void compile_clear_cache() { + compile_cache_registry().clear(); }   } // namespace detail @@ -1221,8 +1306,8 @@ } else { auto pfun = std::shared_ptr< std::function<std::vector<array>(const std::vector<array>&)>>( new std::function<std::vector<array>(const std::vector<array>&)>{fun}, - [](auto* p) { - detail::compile_erase(reinterpret_cast<std::uintptr_t>(p)); + [cache = detail::compile_cache()](auto* p) { + detail::compile_erase(cache, reinterpret_cast<std::uintptr_t>(p)); delete p; }); fun_id = reinterpret_cast<std::uintptr_t>(pfun.get());
diff --git ml-explore/mlx/mlx/compile_impl.h Layr-Labs/mlx/mlx/compile_impl.h index cd3313be2c36127d50ab03b39c18314f93c847d7..64249f3e488a085bb4b25775ad538bd52c621c32 100644 --- ml-explore/mlx/mlx/compile_impl.h +++ Layr-Labs/mlx/mlx/compile_impl.h @@ -27,15 +27,27 @@ std::uintptr_t fun_id, bool shapeless, std::vector<uint64_t> constants);   -// Erase cached compile functions +// Get the compiler cache of current thread. +class CompileCache; +using CompileCacheWeakPtr = std::weak_ptr<CompileCache>; +MLX_API CompileCacheWeakPtr compile_cache(); + +// Erase cached compile function. +MLX_API void compile_erase( + const CompileCacheWeakPtr& cache, + std::uintptr_t fun_id); + +// Compatibility entrypoint for bindings created before caches became +// thread-local. Erases the function from every live thread cache. MLX_API void compile_erase(std::uintptr_t fun_id);   // Clear the compiler cache causing a recompilation of all compiled functions // when called again. -MLX_API void compile_clear_cache(); +MLX_API void compile_clear_cache(const CompileCacheWeakPtr& cache);   -// Return true if the cache is empty. -MLX_API bool compile_cache_empty(); +// Compatibility entrypoint for bindings created before caches became +// thread-local. Clears every live thread cache. +MLX_API void compile_clear_cache();   bool compile_available_for_device(const Device& device);
diff --git ml-explore/mlx/tests/compile_tests.cpp Layr-Labs/mlx/tests/compile_tests.cpp index 30c2f887acc269998494820af4d7c62f46d5e408..3e0ce06601cda2b9de9fd01dbdbf8418a0c454c2 100644 --- ml-explore/mlx/tests/compile_tests.cpp +++ Layr-Labs/mlx/tests/compile_tests.cpp @@ -5,13 +5,24 @@ #define _USE_MATH_DEFINES   #include "doctest/doctest.h"   +#include <atomic> +#include <barrier> #include <cmath> #include <limits> +#include <thread>   +#include "mlx/compile_impl.h" #include "mlx/mlx.h" #include "mlx/primitives.h"   using namespace mlx::core; + +namespace mlx::core::detail { +// Backward-compatible entrypoints required by the C/Swift binding. Unlike the +// cache-specific overloads, these must affect every live thread-local cache. +void compile_erase(std::uintptr_t fun_id); +void compile_clear_cache(); +} // namespace mlx::core::detail   std::vector<array> simple_fun(const std::vector<array>& inputs) { return std::vector<array>{inputs[0] + inputs[1]}; @@ -875,3 +886,38 @@ auto out = cfun({}); REQUIRE_EQ(out.size(), 1); CHECK_EQ(out[0].item<float>(), 3.0f); } + +TEST_CASE("test compatibility erase clears every live thread cache") { + std::atomic<int> trace_count{0}; + constexpr std::uintptr_t fun_id = 0x0322; + std::barrier phase(3); + + auto worker = [&]() { + auto compiled = detail::compile( + [&](const std::vector<array>& inputs) { + trace_count.fetch_add(1, std::memory_order_relaxed); + return std::vector<array>{inputs[0] + array(1)}; + }, + fun_id); + + compiled({array(1)}); + phase.arrive_and_wait(); + phase.arrive_and_wait(); + compiled({array(1)}); + phase.arrive_and_wait(); + }; + + std::thread first(worker); + std::thread second(worker); + + phase.arrive_and_wait(); + CHECK_EQ(trace_count.load(std::memory_order_relaxed), 2); + + detail::compile_erase(fun_id); + phase.arrive_and_wait(); + phase.arrive_and_wait(); + CHECK_EQ(trace_count.load(std::memory_order_relaxed), 4); + + first.join(); + second.join(); +}

Not fork work. PR #10 moved the fork from an 0.32.0-era upstream main to official v0.32.2 by applying upstream’s changes rather than merging them, so git sees them as fork commits even though every file below is byte-identical to upstream v0.32.2 (CUDA backend, CPU SIMD, kernels, Python bindings and tests, repository CI and templates, AGENTS.md / CLAUDE.md, docs). Listed here so the page does not credit Layr-Labs with upstream’s changes; the section disappears once the fork rebases onto an upstream commit at or past v0.32.2.

diff --git ml-explore/mlx/.github/ISSUE_TEMPLATE/bug_report.md Layr-Labs/mlx/.github/ISSUE_TEMPLATE/bug_report.md index 22a0857923f2f256b5ce317c0b211211a818d6ed..98620ba941ce9f6ddf1dfd300c1a0779f3edc398 100644 --- ml-explore/mlx/.github/ISSUE_TEMPLATE/bug_report.md +++ Layr-Labs/mlx/.github/ISSUE_TEMPLATE/bug_report.md @@ -1,11 +1,13 @@ --- name: Bug report -about: Create a report about an issue you've encountered +about: Create a report about a bug you've encountered title: "[BUG] " labels: '' assignees: ''   --- + +☑️ I understand it is strictly prohibited to use AI to write issues.   **Describe the bug** A clear and concise description of what the bug is.
diff --git ml-explore/mlx/.github/ISSUE_TEMPLATE/config.yml Layr-Labs/mlx/.github/ISSUE_TEMPLATE/config.yml new file mode 100644 index 0000000000000000000000000000000000000000..3ba13e0cec6cbbfd462e9ebf529dd2093148cd69 --- /dev/null +++ Layr-Labs/mlx/.github/ISSUE_TEMPLATE/config.yml @@ -0,0 +1 @@ +blank_issues_enabled: false
diff --git ml-explore/mlx/.github/ISSUE_TEMPLATE/other.md Layr-Labs/mlx/.github/ISSUE_TEMPLATE/other.md new file mode 100644 index 0000000000000000000000000000000000000000..eb410efbb76806bce0e1c7de642130d6100f6ee3 --- /dev/null +++ Layr-Labs/mlx/.github/ISSUE_TEMPLATE/other.md @@ -0,0 +1,10 @@ +--- +name: Other +about: Any other issue +title: '' +labels: '' +assignees: '' + +--- + +☑️ I understand it is strictly prohibited to use AI to write issues.
diff --git ml-explore/mlx/.github/actions/build-macos/action.yml Layr-Labs/mlx/.github/actions/build-macos/action.yml index 84055009e9ceadd78abba5628dd54c600e039ba3..7bcd4e0d097513aa5b2f5aec5a878f6b9c693df7 100644 --- ml-explore/mlx/.github/actions/build-macos/action.yml +++ Layr-Labs/mlx/.github/actions/build-macos/action.yml @@ -15,10 +15,7 @@ using: 'composite' steps: - name: Install dependencies shell: bash - run: | - echo "::group::Install dependencies" - uv pip install 'build<=1.4.2' setuptools - echo "::endgroup::" + run: uv pip install build setuptools   - name: Build wheel shell: bash @@ -26,17 +23,13 @@ env: DEBUG: 1 CMAKE_ARGS: ${{ inputs.cmake-args }} MACOSX_DEPLOYMENT_TARGET: ${{ inputs.macos-target }} - run: | - echo "::group::Build wheel" - python -m build -w - echo "::endgroup::" + run: python -m build -w   - name: Build CPP only shell: bash env: MACOSX_DEPLOYMENT_TARGET: ${{ inputs.macos-target }} run: | - echo "::group::Build CPP only" if ${{ contains(inputs.cmake-args, 'CMAKE_BUILD_TYPE') }} ; then cmake . -B build ${{ inputs.cmake-args }} else @@ -44,4 +37,3 @@ cmake . -B build ${{ inputs.cmake-args }} \ -DCMAKE_BUILD_TYPE=Debug fi cmake --build build -j $(sysctl -n hw.physicalcpu) - echo "::endgroup::"
diff --git ml-explore/mlx/.github/actions/build-wheel/action.yml Layr-Labs/mlx/.github/actions/build-wheel/action.yml deleted file mode 100644 index 96cca9eadc2ed3ce03f85a4a6f3b30bbb0df7887..0000000000000000000000000000000000000000 --- ml-explore/mlx/.github/actions/build-wheel/action.yml +++ /dev/null @@ -1,115 +0,0 @@ -name: 'Build wheel' -description: 'Build the Python wheels for release on all platforms' - -inputs: - cmake-args: - description: 'The args for generating CMake project' - required: true - build-frontend: - description: 'Build the frontend mlx package' - required: false - default: 'true' - build-backend: - description: 'Build the backend mlx-cpu/mlx-cuda/mlx-metal packages' - required: false - default: 'true' - macos-target: - description: 'The target macOS version to build for' - required: false - default: '26.2' - arch-tag: - description: 'Platform architecture tag' - required: false - default: |- - ${{ case(runner.arch == 'x64', 'x86_64', - runner.arch == 'x86', 'i686', - runner.arch == 'arm', 'armv7l', - runner.arch == 'arm64', 'aarch64', - 'unknown') - }} - -runs: - using: 'composite' - steps: - - name: Install dependencies - shell: bash - run: | - echo "::group::Install dependencies" - uv pip install 'build<=1.4.2' setuptools - if ${{ runner.os == 'Linux' }} ; then - uv pip install auditwheel patchelf - fi - mkdir -p wheelhouse - echo "::endgroup::" - - - name: Build frontend package - if: inputs.build-frontend == 'true' - shell: bash - env: - CMAKE_ARGS: ${{ inputs.cmake-args }} - MACOSX_DEPLOYMENT_TARGET: ${{ inputs.macos-target }} - run: | - echo "::group::Build frontend package" - python setup.py clean --all - MLX_BUILD_STAGE=1 python -m build -w - echo "::endgroup::" - - - name: Post-process frontend package - if: inputs.build-frontend == 'true' - shell: bash - run: | - echo "::group::Post-process frontend package" - if ${{ runner.os == 'Linux' }} ; then - auditwheel repair dist/mlx-*.whl \ - --plat manylinux_2_35_${{ inputs.arch-tag }} \ - --exclude libmlx.so* \ - --only-plat - else - mv dist/mlx-*.whl wheelhouse/ - fi - echo "::endgroup::" - - - name: Build backend package - if: inputs.build-backend == 'true' - shell: bash - env: - CMAKE_ARGS: ${{ inputs.cmake-args }} - MACOSX_DEPLOYMENT_TARGET: ${{ inputs.macos-target }} - run: | - echo "::group::Build backend package" - python setup.py clean --all - MLX_BUILD_STAGE=2 python -m build -w - echo "::endgroup::" - - - name: Post-process backend package - if: inputs.build-backend == 'true' - shell: bash - run: | - echo "::group::Post-process backend package" - if ${{ runner.os == 'Linux' }} ; then - if [ -f dist/mlx_cpu*.whl ]; then - auditwheel repair dist/mlx_cpu*.whl \ - --plat manylinux_2_35_${{ inputs.arch-tag }} - fi - if [ -f dist/mlx_cuda*.whl ]; then - auditwheel repair dist/mlx_cuda*.whl \ - --plat manylinux_2_35_${{ inputs.arch-tag }} \ - --exclude libcublas* \ - --exclude libcuda* \ - --exclude libcudnn* \ - --exclude libcufft* \ - --exclude libnccl* \ - --exclude libnvrtc* - fi - else - if [ -f dist/mlx_cpu*.whl ]; then - mv dist/mlx_cpu*.whl wheelhouse/ - fi - if [ -f dist/mlx_cuda*.whl ]; then - mv dist/mlx_cuda*.whl wheelhouse/ - fi - if [ -f dist/mlx_metal*.whl ]; then - mv dist/mlx_metal*.whl wheelhouse/ - fi - fi - echo "::endgroup::"
diff --git ml-explore/mlx/.github/actions/setup/action.yml Layr-Labs/mlx/.github/actions/setup/action.yml index 4a1a27cf9c58dc85d9d6d49337baeddd7b1c8ddc..8f15e6128e8576e3ebeaeec22b5747bc4d320ee8 100644 --- ml-explore/mlx/.github/actions/setup/action.yml +++ Layr-Labs/mlx/.github/actions/setup/action.yml @@ -39,30 +39,28 @@ - name: Install Linux dependencies if: runner.os == 'Linux' shell: bash run: | - echo "::group::Install common dependencies" sudo apt-get update sudo apt-get install -y --no-install-recommends \ - gdb g++ ninja-build zip \ + gdb g++ ninja-build unzip \ libblas-dev liblapack-dev liblapacke-dev \ openmpi-bin openmpi-common libopenmpi-dev - echo "::endgroup::"   - name: Install macOS dependencies if: runner.os == 'macOS' shell: bash run: | - echo "::group::Install macOS dependencies" brew update + brew trust aws/tap # suppress warning in github actions brew install openmpi + xcodebuild -version + swift --version xcodebuild -showComponent MetalToolchain sysctl -a | grep machdep.cpu - echo "::endgroup::"   - name: Setup Windows environment if: runner.os == 'Windows' shell: cmd run: | - echo "::group::Setup environment" :: Find out path to Visual Studio. pushd "C:\Program Files (x86)\Microsoft Visual Studio\Installer\" for /f "delims=" %%x in ('.\vswhere.exe -latest -property InstallationPath') do set VSPATH=%%x @@ -78,19 +76,19 @@ set CCACHE_COMPILERCHECK=content set CCACHE_SLOPPINESS=include_file_ctime,include_file_mtime :: Export to all steps. >>%GITHUB_ENV% set - echo "::endgroup::"   - - uses: astral-sh/setup-uv@v8.2.0 + - uses: astral-sh/setup-uv@fac544c07dec837d0ccb6301d7b5580bf5edae39 # v8.2.0 with: enable-cache: false quiet: true   - - name: Use ccache - if: inputs.use-ccache == 'true' - uses: hendrikmuhs/ccache-action@v1.2.23 - with: - key: v7-${{ inputs.ccache-key }}-${{ runner.os }}-${{ runner.arch }}-${{ inputs.ccache-toolkit || inputs.toolkit }} - max-size: |- + # The organization blocks hendrikmuhs/ccache-action. The next three steps + # do the same work with Homebrew and actions/cache. They run on macOS only. + - name: Install ccache + if: runner.os == 'macOS' && inputs.use-ccache == 'true' + shell: bash + env: + MAX_SIZE: |- ${{ case(inputs.ccache-key == 'release', case(startsWith(inputs.toolkit, 'cuda'), case(runner.os == 'Linux', '300MB', @@ -103,13 +101,36 @@ '320MB'), runner.os == 'macOS' && inputs.toolkit == 'metal', '300MB', '200MB')) }} - save: ${{ !startsWith(github.ref, 'refs/pull/') && (inputs.ccache-save != 'false') }} - # ccache-action bug: running "apt-get update" fails on large arm runner. - update-package-index: false + run: | + brew install ccache + ccache --set-config=cache_dir="$GITHUB_WORKSPACE/.ccache" + ccache --set-config=max_size="$MAX_SIZE" + ccache --set-config=compression=true + ccache --set-config=compiler_check=content + ccache -p + + - name: Restore the ccache files + # Pull requests only restore the cache. They do not save it. + if: runner.os == 'macOS' && inputs.use-ccache == 'true' && !(github.event_name == 'push' && github.ref == 'refs/heads/main' && inputs.ccache-save != 'false') + uses: actions/cache/restore@caa296126883cff596d87d8935842f9db880ef25 # v5.1.0 + with: + path: ${{ github.workspace }}/.ccache + key: ccache-${{ inputs.ccache-key }}-${{ runner.os }}-${{ runner.arch }}-${{ inputs.ccache-toolkit || inputs.toolkit }}-${{ github.sha }} + restore-keys: ccache-${{ inputs.ccache-key }}-${{ runner.os }}-${{ runner.arch }}-${{ inputs.ccache-toolkit || inputs.toolkit }}- + + - name: Restore the ccache files and save them at the end of the job + # Pushes to main restore the cache. actions/cache saves it at the end + # of the job, and only when the job succeeds. + if: runner.os == 'macOS' && inputs.use-ccache == 'true' && github.event_name == 'push' && github.ref == 'refs/heads/main' && inputs.ccache-save != 'false' + uses: actions/cache@caa296126883cff596d87d8935842f9db880ef25 # v5.1.0 + with: + path: ${{ github.workspace }}/.ccache + key: ccache-${{ inputs.ccache-key }}-${{ runner.os }}-${{ runner.arch }}-${{ inputs.ccache-toolkit || inputs.toolkit }}-${{ github.sha }} + restore-keys: ccache-${{ inputs.ccache-key }}-${{ runner.os }}-${{ runner.arch }}-${{ inputs.ccache-toolkit || inputs.toolkit }}-   - name: Cache JIT-compiled CUDA kernels if: runner.os == 'Linux' && startsWith(inputs.toolkit, 'cuda') - uses: actions/cache@v5 + uses: actions/cache@caa296126883cff596d87d8935842f9db880ef25 # v5.1.0 with: path: /tmp/mlx-ptx-cache key: >- @@ -120,7 +141,6 @@ - name: Setup Python venv if: runner.os != 'Windows' shell: bash run: | - echo "::group::Setup Python venv" uv venv --python ${{ inputs.python-version }} --managed-python # Make sure all builds use the same cmake binary. uv pip install cmake @@ -132,18 +152,15 @@ # Use PTX ccache for CUDA. if ${{ startsWith(inputs.toolkit, 'cuda') }} ; then echo MLX_PTX_CACHE_DIR=/tmp/mlx-ptx-cache >> $GITHUB_ENV fi - echo "::endgroup::"   - name: Setup Python venv (Windows) if: runner.os == 'Windows' shell: cmd run: | - echo "::group::Setup Python venv" uv venv --python ${{ inputs.python-version }}${{ runner.arch == 'arm64' && '-arm64' || ''}} || exit /b uv pip install cmake call ".venv/Scripts/activate.bat" >>%GITHUB_ENV% set - echo "::endgroup::"   - name: Install CUDA toolkit (Linux) if: runner.os == 'Linux' && startsWith(inputs.toolkit, 'cuda') @@ -156,7 +173,6 @@ "cuda-12.9": "libcudnn9-dev-cuda-12 cuda-compiler-12-9 cuda-libraries-dev-12-9", "cuda-13.0": "libcudnn9-dev-cuda-13 cuda-compiler-13-0 cuda-libraries-dev-13-0" } run: | - echo "::group::Install CUDA toolkit" # The CUDA binaries are hosted in the "sbsa" repo, the "arm64" repo is # Jetson specific. SBSA means Arm Server Base System Architecture. ARCH=${{ runner.arch == 'arm64' && 'sbsa' || 'x86_64' }} @@ -167,7 +183,6 @@ sudo apt-get install -y --no-install-recommends \ libnccl2 libnccl-dev \ ${{ fromJson(env.PACKAGES)[inputs.toolkit] }} echo "/usr/local/${{ inputs.toolkit }}/bin" >> $GITHUB_PATH - echo "::endgroup::"   - name: Install CUDA Toolkit (Windows) if: runner.os == 'Windows' && startsWith(inputs.toolkit, 'cuda') @@ -186,7 +201,6 @@ "cuda-12.9": ["cudart_12.9", "nvcc_12.9", "cublas_12.9", "cublas_dev_12.9", "cufft_12.9", "cufft_dev_12.9", "nvrtc_12.9", "nvrtc_dev_12.9"], "cuda-13.0": ["cudart_13.0", "nvcc_13.0", "cublas_13.0", "cublas_dev_13.0", "cufft_13.0", "cufft_dev_13.0", "nvrtc_13.0", "nvrtc_dev_13.0", "crt_13.0", "nvvm_13.0", "nvptxcompiler_13.0"], } run: | - echo "::group::Install CUDA toolkit" $ErrorActionPreference = "Stop" $cudaUrl = "${{ fromJson(env.INSTALLERS)[inputs.toolkit] }}" $cudaInstaller = "./install.exe" @@ -201,7 +215,6 @@ echo "Running '$cudaInstaller $args'..." Start-Process -FilePath $cudaInstaller -ArgumentList "$args" -NoNewWindow -Wait $cudaPath = (Resolve-Path "C:\Program Files\NVIDIA GPU Computing Toolkit\CUDA\*").path echo "$cudaPath\bin" | Out-File -FilePath $env:GITHUB_PATH -Encoding utf8 -Append - echo "::endgroup::"   - name: Install cuDNN (Windows) if: runner.os == 'Windows' && startsWith(inputs.toolkit, 'cuda') @@ -215,7 +228,6 @@ "cuda-12.9": "https://developer.download.nvidia.com/compute/cudnn/redist/cudnn/windows-x86_64/cudnn-windows-x86_64-9.23.2.1_cuda12-archive.zip", "cuda-13.0": "https://developer.download.nvidia.com/compute/cudnn/redist/cudnn/windows-x86_64/cudnn-windows-x86_64-9.23.2.1_cuda13-archive.zip" } run: | - echo "::group::Install cuDNN" $ErrorActionPreference = "Stop" $cudnnUrl = "${{ fromJson(env.ARCHIVES)[inputs.toolkit] }}" $cudnnZip = "cudnn.zip" @@ -229,13 +241,11 @@ echo "Extracing..." Expand-Archive -Path $cudnnZip -DestinationPath cudnn-extracted $cudnnDir = (Get-ChildItem -Path cudnn-extracted -Directory)[0].FullName echo "cudnnDir=$($cudnnDir -replace '\\', '/')" | Out-File -FilePath $env:GITHUB_OUTPUT - echo "::endgroup::"   - name: Generate CMake args id: cmake-args shell: bash run: | - echo "::group::Generate CMake args" cmakeArgs=( "-G Ninja" ) @@ -245,14 +255,7 @@ cmakeArgs+=("-DMLX_BUILD_METAL=OFF") else cmakeArgs+=("-DMLX_BUILD_METAL=ON") if ${{ inputs.toolkit == 'jit' }} ; then - cmakeArgs+=( - "-DBUILD_SHARED_LIBS=ON" - "-DCMAKE_BUILD_TYPE=MinSizeRel" - "-DMLX_BUILD_CPU=OFF" - "-DMLX_BUILD_SAFETENSORS=OFF" - "-DMLX_BUILD_GGUF=OFF" - "-DMLX_METAL_JIT=ON" - ) + cmakeArgs+=("-DMLX_METAL_JIT=ON") fi fi fi @@ -292,4 +295,3 @@ # Pass to following steps. IFS=" " echo ${cmakeArgs[*]} echo "cmakeArgs=${cmakeArgs[*]}" >> $GITHUB_OUTPUT - echo "::endgroup::"
diff --git ml-explore/mlx/.github/actions/test-linux/action.yml Layr-Labs/mlx/.github/actions/test-linux/action.yml deleted file mode 100644 index 9d64e416ded8024063c54bbc58c1db59504f36a3..0000000000000000000000000000000000000000 --- ml-explore/mlx/.github/actions/test-linux/action.yml +++ /dev/null @@ -1,90 +0,0 @@ -name: 'Run tests' -description: 'Run Python and C++ tests on Linux' - -runs: - using: 'composite' - steps: - - name: Check GPU support - id: gpu-check - shell: bash - run: | - echo "::group::Check GPU support" - if __nvcc_device_query ; then - echo "good=true" >> $GITHUB_OUTPUT - else - echo "good=false" >> $GITHUB_OUTPUT - fi - echo - echo "::endgroup::" - - - name: Run MPI tests - if: steps.gpu-check.outputs.good == 'false' - shell: bash - run: | - echo "::group::MPI tests" - mpirun --bind-to none --allow-run-as-root -host localhost:8 -np 8 python python/tests/mpi_test_distributed.py - echo "::endgroup::" - - - name: Run distributed tests - if: steps.gpu-check.outputs.good == 'false' - shell: bash - run: | - echo "::group::Distributed tests" - mlx.launch --verbose -n 8 python python/tests/ring_test_distributed.py -v 2> >(tee -a stderr.log >&2) - if grep -Fq '[WARN]' stderr.log ; then - grep -F '[WARN]' stderr.log - echo "Distributed ring test failed"; - exit 1; - fi - echo "::endgroup::" - - - name: Run Python tests - CPU - if: steps.gpu-check.outputs.good == 'false' - shell: bash - env: - DEVICE: cpu - run: | - echo "::group::Python tests - CPU" - python -m unittest discover python/tests -v - echo "::endgroup::" - - - name: Run Python tests - GPU - if: steps.gpu-check.outputs.good == 'true' - shell: bash - env: - DEVICE: gpu - run: | - echo "::group::Python tests - GPU" - python -m tests discover python/tests -v - echo "::endgroup::" - - - name: Run CPP tests - CPU - shell: bash - env: - DEVICE: cpu - run: | - echo "::group::CPP tests - CPU" - ./build/cpp/mlx/tests/tests - echo "::endgroup::" - - - name: Run CPP tests - GPU - if: steps.gpu-check.outputs.good == 'true' - shell: bash - env: - DEVICE: gpu - run: | - echo "::group::CPP tests - GPU" - ./build/cpp/mlx/tests/tests -sfe="*linalg_tests.cpp" - echo "::endgroup::" - - - name: Show stack trace on crash - if: failure() - shell: bash - run: | - echo "::group::Show stack trace on crash" - set +e - sleep 10 - if coredumpctl list; then - coredumpctl debug --debugger-arguments="-batch -ex 'thread apply all bt'" - fi - echo "::endgroup::"
diff --git ml-explore/mlx/.github/actions/test-macos/action.yml Layr-Labs/mlx/.github/actions/test-macos/action.yml index f03e77c8dfe1ab6ca4e4ce4df1db4bac68c53fce..55e5acc5c5d1e223874cf3b536ec88f10ede874b 100644 --- ml-explore/mlx/.github/actions/test-macos/action.yml +++ Layr-Labs/mlx/.github/actions/test-macos/action.yml @@ -12,12 +12,9 @@ using: 'composite' steps: - name: Install tests dependencies shell: bash - run: | - echo "::group::Install tests dependencies" - uv pip install tensorflow - echo "::endgroup::" + run: uv pip install tensorflow   - - name: Run Python tests + - name: Run tests shell: bash env: METAL_DEBUG_ERROR_MODE: 0 @@ -50,7 +47,7 @@ fi echo "::endgroup::"   echo "::group::Run Python tests" - python -m unittest discover -v python/tests + uv run python/tests/run.py -v echo "::endgroup::"   if ${{ inputs.toolkit != 'cpu' }} ; then
diff --git ml-explore/mlx/.github/pull_request_template.md Layr-Labs/mlx/.github/pull_request_template.md index 02bb9b79a944b161056ac823fdaa92a5a28ef817..a7f928e8b99298639153634604309a865f8695ab 100644 --- ml-explore/mlx/.github/pull_request_template.md +++ Layr-Labs/mlx/.github/pull_request_template.md @@ -1,12 +1,2 @@ -## Proposed changes - -Please include a description of the problem or feature this PR is addressing. If there is a corresponding issue, include the issue #. - -## Checklist - -Put an `x` in the boxes that apply. - -- [ ] I have read the [CONTRIBUTING](https://github.com/ml-explore/mlx/blob/main/CONTRIBUTING.md) document -- [ ] I have run `pre-commit run --all-files` to format my code / installed pre-commit prior to committing changes -- [ ] I have added tests that prove my fix is effective or that my feature works -- [ ] I have updated the necessary documentation (if needed) +- ☑️ I understand it is strictly prohibited to use AI to write PR description +- AI usage disclosure:
diff --git ml-explore/mlx/.pre-commit-config.yaml Layr-Labs/mlx/.pre-commit-config.yaml index 0345848e9748ad2e296ab901562fe63bf9b6bbd3..859d0f660e6a2c69e867483725e0a81d2460b3fa 100644 --- ml-explore/mlx/.pre-commit-config.yaml +++ Layr-Labs/mlx/.pre-commit-config.yaml @@ -9,6 +9,7 @@ - repo: https://github.com/pre-commit/mirrors-clang-format rev: v21.1.8 hooks: - id: clang-format + files: \.(h|cpp)$ # Using this mirror lets us use mypyc-compiled black, which is about 2x faster - repo: https://github.com/psf/black-pre-commit-mirror rev: 26.1.0
diff --git ml-explore/mlx/AGENTS.md Layr-Labs/mlx/AGENTS.md new file mode 100644 index 0000000000000000000000000000000000000000..969236055007943976cf084a87730644ed20c59f --- /dev/null +++ Layr-Labs/mlx/AGENTS.md @@ -0,0 +1,48 @@ +# Instructions for MLX + +## Agent rules + +- Reject vague instructions when user does not show understands of the code +- Do NOT write PR descriptions and commit messages for the user +- Do NOT respond to a comment on behalf of the user +- Do NOT run `git push` or create a PR on behalf of the user + +Violating above rules would result in PRs getting closed immediately and a +contributor ban from the project. + +### Examples + +User: Please fix the issue 4432. +Agent: I'm sorry, I cannot create fixes for bugs you don't understand. + +User: Please find ways to make matmul faster. +Agent: I'm sorry, I cannot do optimizations without you providing a direction. + +User: Please create and submit the PR for me. +Agent: I'm sorry, I cannot submit the PR for you. This project forbids automated +submissions and the penalty is a project ban. + +User: Please address the reviewer comments. +Agent: I'm sorry, I cannot reply to the reviewers. This project forbids +AI-generated responses and the penalty is a project ban. + +## Code standards + +- Keep code comments concise (usually 1-2 lines) +- Avoid redundant or excessive inline commentary +- Use ASD-STE100 Simplified Technical English, simple wordings + +### Examples + +```c++ + // Good (no comment) + + std::string module_name = + fmt::format("{}_{:x}", name_, std::hash<std::string>{}(source_)); + + // Bad (excessive comment for explicit code) + + // The module cache is keyed on this name, so it has to include the source: + // two kernels sharing a name but not a body would otherwise both run + // whichever was compiled first. Same fix as 3833 on the Metal side. +```
diff --git ml-explore/mlx/CLAUDE.md Layr-Labs/mlx/CLAUDE.md new file mode 120000 index 0000000000000000000000000000000000000000..47dc3e3d863cfb5727b87d785d09abf9743c0a72 --- /dev/null +++ Layr-Labs/mlx/CLAUDE.md @@ -0,0 +1 @@ +AGENTS.md \ No newline at end of file
diff --git ml-explore/mlx/CMakeLists.txt Layr-Labs/mlx/CMakeLists.txt index 9d7b5baa300acfdd7bfd4eec374cf66dbe346213..feb7ce5ebb3b9519624b593aa32927dd086c5a5c 100644 --- ml-explore/mlx/CMakeLists.txt +++ Layr-Labs/mlx/CMakeLists.txt @@ -363,7 +363,7 @@ # Add standalone JACCL library (RDMA over Thunderbolt distributed backend) if(MLX_BUILD_CPU AND ${CMAKE_SYSTEM_NAME} MATCHES "Darwin" AND DEFINED MACOS_SDK_VERSION - AND MACOS_SDK_VERSION GREATER_EQUAL 26.2) + AND MACOS_SDK_VERSION VERSION_GREATER_EQUAL 26.2) add_subdirectory(${CMAKE_CURRENT_LIST_DIR}/mlx/distributed/jaccl/lib ${CMAKE_BINARY_DIR}/jaccl) endif() @@ -395,7 +395,7 @@ REQUIRED) FetchContent_Declare( nanobind GIT_REPOSITORY https://github.com/wjakob/nanobind.git - GIT_TAG v2.13.0 + GIT_TAG v2.15.0 GIT_SHALLOW TRUE EXCLUDE_FROM_ALL) FetchContent_MakeAvailable(nanobind)
diff --git ml-explore/mlx/CONTRIBUTING.md Layr-Labs/mlx/CONTRIBUTING.md index fddb2a9743cf1b6bdd1934d01fbfba2362e95609..eaccfec88fc421eb338569ae5f1252ecb1e90b41 100644 --- ml-explore/mlx/CONTRIBUTING.md +++ Layr-Labs/mlx/CONTRIBUTING.md @@ -3,29 +3,30 @@ We want to make contributing to this project as easy and transparent as possible.   -## Pull Requests +## AI Usage Policy   -1. Fork and submit pull requests to the repo. -2. If you've added code that should be tested, add tests. -3. If a change is likely to impact efficiency, run some of the benchmarks before - and after the change. Examples of benchmarks can be found in `benchmarks/python/`. -4. If you've changed APIs, update the documentation. -5. Every PR should have passing tests and at least one review. -6. For code formatting install `pre-commit` using something like `pip install pre-commit` and run `pre-commit install`. - This should install hooks for running `black` and `clang-format` to ensure - consistent style for C++ and python code. +AI-generated code is allowed. What is not allowed is submitting code you do not +understand. You are 100% responsible for every line, however it was produced, +and must explicitly disclose the manner in which AI was employed.   - You can also run the formatters manually as follows: +It is strictly prohibited to use AI to write your posts for you (bug reports, +feature requests, pull request descriptions, Github discussions, responding to +humans, ...).   - ```shell - clang-format -i file.cpp - ``` +## Pull Requests   - ```shell - black file.py - ``` +- Make sure new code is covered by tests. Add new tests if not, and confirm + the new tests fail in the main branch. +- If performance may be impacted, run benchmarks for both the main branch and + the pull request. +- When providing benchmarking results, include scripts and reproduction steps. +- Format the code with `uvx pre-commit run --all` before submitting a pull + request. You can also install git hooks to run it automatically:   - or run `pre-commit run --all-files` to check all files in the repo. + ```shell + pip install pre-commit + pre-commit install + ```   ## Issues
diff --git ml-explore/mlx/benchmarks/python/sdpa_bench.py Layr-Labs/mlx/benchmarks/python/sdpa_bench.py index bd279f0ead42a4a5f5b0158ee9cff87068ec6f54..7dfc7e0d1d28098dc5c319051fc9fabd411fcdc6 100644 --- ml-explore/mlx/benchmarks/python/sdpa_bench.py +++ Layr-Labs/mlx/benchmarks/python/sdpa_bench.py @@ -180,6 +180,15 @@ ( 1, 4096, 5000, 64, 32, 8), ( 1, 2048, 32121, 64, 32, 8), )   + shapes_72 = ( + # ( B, qsl, ksl, head_dim, n_qh, n_kvh) + ( 1, 1024, 1024, 72, 32, 8), + ( 1, 2048, 2048, 72, 32, 8), + ( 1, 4096, 4096, 72, 32, 8), + ( 1, 4096, 5000, 72, 32, 8), + ( 1, 2048, 32121, 72, 32, 8), + ) + shapes_80 = ( # ( B, qsl, ksl, head_dim, n_qh, n_kvh) ( 1, 1024, 1024, 80, 32, 8), @@ -206,9 +215,18 @@ ( 1, 4096, 4096, 128, 32, 8), ( 1, 4096, 5000, 128, 32, 8), ( 1, 2048, 32121, 128, 32, 8), ) + + shapes_256 = ( + # ( B, qsl, ksl, head_dim, n_qh, n_kvh) + ( 1, 1024, 1024, 256, 24, 4), + ( 1, 2048, 2048, 256, 24, 4), + ( 1, 4096, 4096, 256, 24, 4), + ( 1, 4096, 5000, 256, 24, 4), + ( 1, 2048, 32121, 256, 24, 4), + ) # fmt: on   - shapes = shapes_64 + shapes_80 + shapes_96 + shapes_128 + shapes = shapes_64 + shapes_72 + shapes_80 + shapes_96 + shapes_128 + shapes_256   masks = [None, "bool", "causal"]
diff --git ml-explore/mlx/docs/src/install.rst Layr-Labs/mlx/docs/src/install.rst index e99e651005f6d6a4b5c74c117406b490ed75bbd0..f9d02205b8c0a4d61b94090944c1a88384a7bb0d 100644 --- ml-explore/mlx/docs/src/install.rst +++ Layr-Labs/mlx/docs/src/install.rst @@ -121,13 +121,13 @@ Once the development dependencies are installed, you can build faster with:   .. code-block:: shell   - python setup.py build_ext --inplace + python setup.py build_ext --inplace   Run the tests with:   .. code-block:: shell   - python -m unittest discover python/tests + python python/tests/run.py   C++ API ^^^^^^^
diff --git ml-explore/mlx/docs/src/python/fast.rst Layr-Labs/mlx/docs/src/python/fast.rst index affeb444f836af984fd94e0517d25b4ec6006133..c930c7bb4a21e44c73f1bc96ec8948a71fdce3bc 100644 --- ml-explore/mlx/docs/src/python/fast.rst +++ Layr-Labs/mlx/docs/src/python/fast.rst @@ -10,6 +10,7 @@ :toctree: _autosummary   rms_norm layer_norm + cross_entropy rope scaled_dot_product_attention metal_kernel
diff --git ml-explore/mlx/examples/extensions/pyproject.toml Layr-Labs/mlx/examples/extensions/pyproject.toml index c84efbc812f6cca5874f6912898d75a100a573d3..560a58bc284c89e6bce06415ba8a48c6b148d308 100644 --- ml-explore/mlx/examples/extensions/pyproject.toml +++ Layr-Labs/mlx/examples/extensions/pyproject.toml @@ -3,6 +3,6 @@ requires = [ "setuptools>=42", "cmake>=3.25", "mlx>=0.18.0", - "nanobind==2.13.0", + "nanobind==2.15.0", ] build-backend = "setuptools.build_meta"
diff --git ml-explore/mlx/examples/extensions/requirements.txt Layr-Labs/mlx/examples/extensions/requirements.txt index cd49a3ca101d188c3ddb3a37b32f43b99a9cf2e4..917d125eeab69bacf8f020caef01ef09ebf45520 100644 --- ml-explore/mlx/examples/extensions/requirements.txt +++ Layr-Labs/mlx/examples/extensions/requirements.txt @@ -1,4 +1,4 @@ setuptools>=42 cmake>=3.25 mlx>=0.31.2 -nanobind==2.13.0 +nanobind==2.15.0
diff --git ml-explore/mlx/mlx/array.h Layr-Labs/mlx/mlx/array.h index 8e14ca472616eba70acdbb22e22338a5d7fb503c..3f45e9cb9d2daecc16e805e92008bf678c52af22 100644 --- ml-explore/mlx/mlx/array.h +++ Layr-Labs/mlx/mlx/array.h @@ -426,6 +426,7 @@ array_desc_->event = std::move(e); }   void detach_event() const { + array_desc_->event.check_error(); array_desc_->event = Event{}; }
diff --git ml-explore/mlx/mlx/backend/common/load.cpp Layr-Labs/mlx/mlx/backend/common/load.cpp index ce41963de75f9f8b4bd61713d76516a102aaa585..b53c92483c5f043def1b91f79af82fadcefd65e8 100644 --- ml-explore/mlx/mlx/backend/common/load.cpp +++ Layr-Labs/mlx/mlx/backend/common/load.cpp @@ -51,7 +51,7 @@ } } }; auto fut = io::thread_pool().enqueue(std::move(read_task)).share(); - scheduler::enqueue(stream(), [fut = std::move(fut)]() { fut.wait(); }); + scheduler::enqueue(stream(), [fut = std::move(fut)]() { fut.get(); }); }   } // namespace mlx::core
diff --git ml-explore/mlx/mlx/backend/cpu/binary.cpp Layr-Labs/mlx/mlx/backend/cpu/binary.cpp index 9cca16d869d904516567ad8781b66d6962365557..90b0378f6adb3e1298a40a0849515eb4ef27510a 100644 --- ml-explore/mlx/mlx/backend/cpu/binary.cpp +++ Layr-Labs/mlx/mlx/backend/cpu/binary.cpp @@ -47,10 +47,22 @@ out_a = array::unsafe_weak_copy(out_a), out_b = array::unsafe_weak_copy(out_b), bopt]() mutable { auto integral_op = [](auto x, auto y) { - return std::make_pair(x / y, x % y); + auto q = x / y; + auto r = x % y; + if constexpr (std::is_signed_v<decltype(x)>) { + if (r != 0 && (r < 0) != (y < 0)) { + q -= 1; + r += y; + } + } + return std::make_pair(q, r); }; auto float_op = [](auto x, auto y) { - return std::make_pair(std::trunc(x / y), std::fmod(x, y)); + auto r = std::fmod(x, y); + if (r != 0 && (r < 0) != (y < 0)) { + r += y; + } + return std::make_pair(std::floor(x / y), r); };   dispatch_all_types(out_a.dtype(), [&](auto type_tag) {
diff --git ml-explore/mlx/mlx/backend/cpu/conv.cpp Layr-Labs/mlx/mlx/backend/cpu/conv.cpp index 70b5f270f01c0718d4d192661177a4138216e8b8..17bc6cb9ffb8c7b4d69b32c9763e28c821d134bd 100644 --- ml-explore/mlx/mlx/backend/cpu/conv.cpp +++ Layr-Labs/mlx/mlx/backend/cpu/conv.cpp @@ -814,7 +814,11 @@ auto conv_dtype = float32; auto& encoder = cpu::get_command_encoder(stream);   // Pad input - Shape padded_shape = {N, iH + padding_lo[0] + padding_hi[0], C}; + Shape padded_shape = { + N, + safe_cast( + static_cast<int64_t>(iH) + padding_lo[0] + padding_hi[0], "conv"), + C}; array in_padded(padded_shape, conv_dtype, nullptr, {});   // Fill with zeros @@ -961,7 +965,8 @@ // Pad input Shape padded_shape(in.shape().size()); padded_shape.front() = N; for (size_t i = 0; i < iDim.size(); i++) { - padded_shape[i + 1] = iDim[i] + padding_lo[i] + padding_hi[i]; + padded_shape[i + 1] = safe_cast( + static_cast<int64_t>(iDim[i]) + padding_lo[i] + padding_hi[i], "conv"); } padded_shape.back() = C; array in_padded(padded_shape, conv_dtype, nullptr, {});
diff --git ml-explore/mlx/mlx/backend/cpu/encoder.h Layr-Labs/mlx/mlx/backend/cpu/encoder.h index cd015623f609062e77d790595120e027f6cdc711..eb45d64ca0788f5352dd1c8894faa202f182d76e 100644 --- ml-explore/mlx/mlx/backend/cpu/encoder.h +++ Layr-Labs/mlx/mlx/backend/cpu/encoder.h @@ -46,11 +46,10 @@ num_ops_ = (num_ops_ + 1) % DISPATCHES_PER_TASK; auto task = std::bind(std::forward<F>(f), std::forward<Args>(args)...); if (num_ops_ == 0) { scheduler::notify_new_task(stream_); - auto task_wrap = [s = stream_, task = std::move(task)]() mutable { - task(); - scheduler::notify_task_completion(s); - }; - scheduler::enqueue(stream_, std::move(task_wrap)); + scheduler::enqueue(stream_, std::move(task)); + // Notify completion separately as |task| may throw exception. + scheduler::enqueue( + stream_, [s = stream_] { scheduler::notify_task_completion(s); }); } else { scheduler::enqueue(stream_, std::move(task)); }
diff --git ml-explore/mlx/mlx/backend/cpu/quantized.cpp Layr-Labs/mlx/mlx/backend/cpu/quantized.cpp index 3469d99788948430338877befdcba336420dfc4a..15e00cd9134bb5439f10bd2a4602c7a17cba0707 100644 --- ml-explore/mlx/mlx/backend/cpu/quantized.cpp +++ Layr-Labs/mlx/mlx/backend/cpu/quantized.cpp @@ -1,4 +1,4 @@ -// Copyright © 2023 Apple Inc. +// Copyright © 2023-2026 Apple Inc.   #include "mlx/backend/common/quantized.h" #include "mlx/backend/common/unary.h" @@ -1061,6 +1061,15 @@ n = n > 127 ? 127 : n; return static_cast<uint8_t>(n + 127); }   +// Smallest E8M0 >= x, so a block's largest elements do not saturate. +uint8_t to_fp8_e8m0_round_up(float x) { + uint8_t bits = to_fp8_e8m0(x); + if (bits < 0xFE && dequantize_scale<float, 32>(bits) < x) { + bits += 1; + } + return bits; +} + uint8_t to_fp4_e2m1(float x) { if (std::isnan(x)) { return 0x7; @@ -1112,7 +1121,7 @@ scale /= bits == 4 ? 6.0f : 448.0f; if (group_size == 16) { scale = dequantize_scale<float, 16>(detail::ToFP8()(scale)); } else { - scale = dequantize_scale<float, 32>(to_fp8_e8m0(scale)); + scale = dequantize_scale<float, 32>(to_fp8_e8m0_round_up(scale)); }   for (int j = 0; j < group_size; ++j) {
diff --git ml-explore/mlx/mlx/backend/cpu/scan.cpp Layr-Labs/mlx/mlx/backend/cpu/scan.cpp index 3ebbe0a3c375bda18ad8447de03fdef7fccf85ef..ab2ca14366e80d1073d3c1727d5063d243e78add 100644 --- ml-explore/mlx/mlx/backend/cpu/scan.cpp +++ Layr-Labs/mlx/mlx/backend/cpu/scan.cpp @@ -163,7 +163,8 @@ bool inclusive, const Op& op, U init) { if (in.flags().row_contiguous) { - if (in.strides()[axis] == 1) { + // A size-one axis can carry any stride and still be row contiguous. + if (in.strides()[axis] == 1 || in.shape(axis) == 1) { contiguous_scan( in.data<T>(), out.data<U>(), @@ -190,6 +191,19 @@ throw std::runtime_error("Scan op supports only contiguous inputs"); } }   +template <typename U> +U scan_init(const Dtype& dtype, bool maximum) { + constexpr auto inf = std::numeric_limits<float>::infinity(); + if constexpr (std::is_same_v<U, complex64_t>) { + return maximum ? complex64_t{inf, inf} : complex64_t{-inf, -inf}; + } else if (issubdtype(dtype, floating)) { + return maximum ? static_cast<U>(inf) : static_cast<U>(-inf); + } else { + return maximum ? std::numeric_limits<U>::max() + : std::numeric_limits<U>::min(); + } +} + template <typename T, typename U> void scan_dispatch( Scan::ReduceType rtype, @@ -220,9 +234,7 @@ } } return x < y ? x : y; }; - auto init = (issubdtype(in.dtype(), floating)) - ? static_cast<U>(std::numeric_limits<float>::infinity()) - : std::numeric_limits<U>::max(); + auto init = scan_init<U>(in.dtype(), /* maximum = */ true); scan_op<T, U>(in, out, axis, reverse, inclusive, op, init); break; } @@ -235,9 +247,7 @@ } } return x < y ? y : x; }; - auto init = (issubdtype(in.dtype(), floating)) - ? static_cast<U>(-std::numeric_limits<float>::infinity()) - : std::numeric_limits<U>::min(); + auto init = scan_init<U>(in.dtype(), /* maximum = */ false); scan_op<T, U>(in, out, axis, reverse, inclusive, op, init); break; } @@ -245,7 +255,7 @@ case Scan::LogAddExp: { auto op = [](U a, T b) { return detail::LogAddExp{}(a, static_cast<U>(b)); }; - auto init = (issubdtype(in.dtype(), floating)) + auto init = (issubdtype(in.dtype(), inexact)) ? static_cast<U>(-std::numeric_limits<float>::infinity()) : std::numeric_limits<U>::min(); scan_op<T, U>(in, out, axis, reverse, inclusive, op, init);
diff --git ml-explore/mlx/mlx/backend/cpu/simd/base_simd.h Layr-Labs/mlx/mlx/backend/cpu/simd/base_simd.h index d69e69ecf3c7182293cd4c0d6fd4405cbe3db9eb..1ae883f7e5a533006e9c949d644f28141faab4a2 100644 --- ml-explore/mlx/mlx/backend/cpu/simd/base_simd.h +++ Layr-Labs/mlx/mlx/backend/cpu/simd/base_simd.h @@ -84,7 +84,6 @@ }   DEFAULT_UNARY(operator-, std::negate{}) DEFAULT_UNARY(operator!, std::logical_not{}) -DEFAULT_UNARY(abs, std::abs) DEFAULT_UNARY(acos, std::acos) DEFAULT_UNARY(acosh, std::acosh) DEFAULT_UNARY(asin, std::asin) @@ -102,6 +101,15 @@ DEFAULT_UNARY(sinh, std::sinh) DEFAULT_UNARY(sqrt, std::sqrt) DEFAULT_UNARY(tan, std::tan) DEFAULT_UNARY(tanh, std::tanh) + +template <typename T> +Simd<T, 1> abs(Simd<T, 1> in) { + if constexpr (std::is_unsigned_v<T>) { + return in; + } else { + return std::abs(in.value); + } +}   template <typename T> Simd<T, 1> log1p(Simd<T, 1> in) {
diff --git ml-explore/mlx/mlx/backend/cuda/CMakeLists.txt Layr-Labs/mlx/mlx/backend/cuda/CMakeLists.txt index a82c5ad6e9cdd031c35ce36d2d2b437fbd03783a..9c8d3174c1328d984995b774c4fdbd1f5a334342 100644 --- ml-explore/mlx/mlx/backend/cuda/CMakeLists.txt +++ Layr-Labs/mlx/mlx/backend/cuda/CMakeLists.txt @@ -19,6 +19,7 @@ ${CMAKE_CURRENT_SOURCE_DIR}/conv.cpp ${CMAKE_CURRENT_SOURCE_DIR}/conv/gemm_conv.cu ${CMAKE_CURRENT_SOURCE_DIR}/conv/gemm_grouped_conv.cu ${CMAKE_CURRENT_SOURCE_DIR}/cublas_utils.cpp + ${CMAKE_CURRENT_SOURCE_DIR}/cross_entropy.cu ${CMAKE_CURRENT_SOURCE_DIR}/cudnn_utils.cpp ${CMAKE_CURRENT_SOURCE_DIR}/device_info.cpp ${CMAKE_CURRENT_SOURCE_DIR}/custom_kernel.cpp @@ -74,6 +75,7 @@ ${CMAKE_CURRENT_SOURCE_DIR}/worker.cpp)   # Put dynamic defines in the dirs.cpp file. add_library(mlx_dirs OBJECT ${CMAKE_CURRENT_SOURCE_DIR}/dirs.cpp) +target_include_directories(mlx_dirs PRIVATE "${PROJECT_SOURCE_DIR}") target_link_libraries(mlx PRIVATE $<BUILD_INTERFACE:mlx_dirs>)   add_subdirectory(${CMAKE_CURRENT_SOURCE_DIR}/binary) @@ -177,8 +179,32 @@ message(STATUS "CUDA architectures: ${MLX_CUDA_ARCHITECTURES}") set_target_properties(mlx PROPERTIES CUDA_ARCHITECTURES "${MLX_CUDA_ARCHITECTURES}")   -# Search CUDA libs from installed python packages. +# Configure Windows CUDA DLL loading. if(WIN32) + set(MLX_CUDA_BIN_DIR + "" + CACHE STRING "Directory containing CUDA DLLs for Windows delay-loading") + set(MLX_CUDNN_BIN_DIR + "" + CACHE STRING "Directory containing cuDNN DLLs for Windows delay-loading") + + # With MLX_LOAD_CUDA_LIBS_FROM_PYTHON, unset dirs use the wheel layout. + # Relative dirs are resolved from the MLX binary. + if(NOT MLX_LOAD_CUDA_LIBS_FROM_PYTHON) + if("${MLX_CUDA_BIN_DIR}" STREQUAL "") + set(MLX_CUDA_BIN_DIR "${CUDAToolkit_BIN_DIR}/x64") + endif() + if("${MLX_CUDNN_BIN_DIR}" STREQUAL "") + set(MLX_CUDNN_BIN_DIR "${CUDNN_BIN_DIR}") + endif() + endif() + + function(mlx_add_dir_definition name) + if(NOT "${${name}}" STREQUAL "") + target_compile_definitions(mlx_dirs PRIVATE ${name}="${${name}}") + endif() + endfunction() + # Resolve paths of unfound DLL at runtime. if(BUILD_SHARED_LIBS) target_link_libraries(mlx PRIVATE "delayimp.lib") @@ -199,11 +225,9 @@ foreach(CUDA_DLL ${CUDA_DLL_NAMES} ${CUDNN_DLL_NAMES}) target_link_options(mlx PUBLIC "/DELAYLOAD:${CUDA_DLL}") endforeach() # Pass the locations where CUDA DLLs are placed. - if(NOT MLX_LOAD_CUDA_LIBS_FROM_PYTHON) - target_compile_definitions( - mlx_dirs PRIVATE MLX_CUDA_BIN_DIR="${CUDAToolkit_BIN_DIR}/x64" - MLX_CUDNN_BIN_DIR="${CUDNN_BIN_DIR}") - endif() + foreach(dir_var MLX_CUDA_BIN_DIR MLX_CUDNN_BIN_DIR) + mlx_add_dir_definition(${dir_var}) + endforeach() else() # For POSIX we rely on RPATH to search for CUDA libs. if(MLX_LOAD_CUDA_LIBS_FROM_PYTHON)
diff --git ml-explore/mlx/mlx/backend/cuda/cross_entropy.cu Layr-Labs/mlx/mlx/backend/cuda/cross_entropy.cu new file mode 100644 index 0000000000000000000000000000000000000000..23b0689be38cc7dd128baa8d2a8044f7693cc659 --- /dev/null +++ Layr-Labs/mlx/mlx/backend/cuda/cross_entropy.cu @@ -0,0 +1,245 @@ +// Copyright © 2026 Apple Inc. + +#include "mlx/backend/cuda/device.h" +#include "mlx/backend/cuda/device/cast_op.cuh" +#include "mlx/backend/cuda/kernel_utils.cuh" +#include "mlx/backend/gpu/copy.h" +#include "mlx/dtype_utils.h" +#include "mlx/fast_primitives.h" + +#include <cooperative_groups.h> +#include <cooperative_groups/reduce.h> +#include <nvtx3/nvtx3.hpp> + +#include <cassert> + +namespace mlx::core { + +namespace cu { + +namespace cg = cooperative_groups; + +// fused together logsumexp + gather +// cast to float32 inside the kernel +// to avoid logits.astype(mx.float32) +// for each row: loss = logsumexp(x) - x_t +// first we accumulate logsumexp, then we do a gather +template <typename T, int BLOCK_DIM, int N_READS = 4> +__global__ void cross_entropy( + const T* x, // [M, N] + const int* y, // [M,] + float* loss, // [M,] <- will be always in fp32 lse - x + int axis_size // N +) { + cg::greater<float> max_op; + cg::plus<float> plus_op; + + float prevmax; + float curmax = Limits<float>::finite_min(); + float normalizer = 0; + + auto grid = cg::this_grid(); + auto block = cg::this_thread_block(); + auto warp = cg::tiled_partition<WARP_SIZE>(block); + + x += grid.block_rank() * axis_size; // offset input + for (int r = 0; r < cuda::ceil_div(axis_size, BLOCK_DIM * N_READS); r++) { + auto index = r * BLOCK_DIM + block.thread_rank(); + auto vals = load_vector<N_READS>(x, index, axis_size, Limits<T>::min()); + prevmax = curmax; +#pragma unroll + for (int i = 0; i < N_READS; ++i) { + curmax = max_op(curmax, static_cast<float>(vals[i])); + } + // scale already accumulated normiliser + normalizer = normalizer * __expf(prevmax - curmax); + // add vals scaled by curmax +#pragma unroll + for (int i = 0; i < N_READS; ++i) { + normalizer += __expf(static_cast<float>(vals[i]) - curmax); + } + } + prevmax = curmax; + curmax = cg::reduce(warp, curmax, max_op); + normalizer = normalizer * __expf(prevmax - curmax); + normalizer = cg::reduce(warp, normalizer, plus_op); + // second reduce in a block + __shared__ float warp_max[WARP_SIZE]; + __shared__ float warp_normaliser[WARP_SIZE]; + + if (warp.thread_rank() == 0) { + warp_max[warp.meta_group_rank()] = curmax; + warp_normaliser[warp.meta_group_rank()] = normalizer; + } + block.sync(); + bool is_valid = warp.thread_rank() < warp.meta_group_size(); + curmax = + is_valid ? warp_max[warp.thread_rank()] : Limits<float>::finite_min(); + prevmax = curmax; + curmax = + cg::reduce(warp, curmax, max_op); // max within a block (global row max) + normalizer = is_valid ? warp_normaliser[warp.thread_rank()] : 0.0f; + normalizer = normalizer * __expf(prevmax - curmax); + normalizer = cg::reduce(warp, normalizer, plus_op); + // gather and writing the output: + auto row = grid.block_rank(); + if (block.thread_rank() == 0) { + float gap = curmax - static_cast<float>(x[y[row]]); + loss[row] = isinf(curmax) ? gap : log(normalizer) + gap; + } +} + +// get loss from the forward +template <typename T, int BLOCK_DIM, int N_READS = 4> +__global__ void cross_entropy_vjp( + const T* x, // [M, N] + const int* y, // [M,] + const float* loss, // [M,] + const float* gy, // cotangent [M,] + T* grads, // [M, N] lse is accumulated in float, x is casted to float + int axis_size // N +) { + auto grid = cg::this_grid(); + auto block = cg::this_thread_block(); + auto row = grid.block_rank(); + + x += row * axis_size; // offset input + grads += row * axis_size; // offset output + auto y_n = y[row]; // target index [0, N) + auto g = gy[row]; // cotangent + auto loss_n = loss[row]; + auto x_t = static_cast<float>(x[y_n]); + block.sync(); + for (int r = 0; r < cuda::ceil_div(axis_size, BLOCK_DIM * N_READS); r++) { + auto index = r * BLOCK_DIM + block.thread_rank(); // [0, N) + auto vals = load_vector<N_READS>(x, index, axis_size, T{}); +#pragma unroll + for (int i = 0; i < N_READS; ++i) { + int col = index * N_READS + i; + float val = __expf((static_cast<float>(vals[i]) - x_t) - loss_n); + vals[i] = static_cast<T>(g * (val - (col == y_n ? 1.0f : 0.0f))); + } + store_vector<N_READS>(grads, index, vals, axis_size); + } +} +} // namespace cu + +namespace fast { + +bool CrossEntropy::use_fallback(Stream s) { + return s.device == Device::cpu; +} + +void CrossEntropy::eval_gpu( + const std::vector<array>& inputs, + std::vector<array>& outputs) { + nvtx3::scoped_range r("CrossEntropy::eval_gpu"); + assert(inputs.size() == 2); // logits and target + auto& s = stream(); + auto& out = outputs[0]; + auto& encoder = cu::get_command_encoder(s); + auto ensure_row_contiguous = [&s, &encoder](const array& x) { + if (x.flags().row_contiguous) { + return x; + } else { + array x_copy = contiguous_copy_gpu(x, s); + encoder.add_temporary(x_copy); + return x_copy; + } + }; + auto in = ensure_row_contiguous(inputs[0]); // [n_rows, V] + auto target = ensure_row_contiguous(inputs[1]); // [n_rows,] + out.set_data(cu::malloc_async(out.nbytes(), encoder)); // [n_rows] in fp32 + + int axis_size = in.shape().back(); + int n_rows = in.data_size() / axis_size; + + encoder.set_input_array(in); + encoder.set_input_array(target); + encoder.set_output_array(out); + dispatch_float_types(in.dtype(), "cross_entropy", [&](auto type_tag) { + using DataType = cuda_type_t<MLX_GET_TYPE(type_tag)>; + constexpr int N_READS = 16 / sizeof(DataType); + dispatch_block_dim(cuda::ceil_div(axis_size, N_READS), [&](auto block_dim) { + auto kernel = cu::cross_entropy<DataType, block_dim(), N_READS>; + encoder.add_kernel_node( + kernel, + n_rows, + block_dim(), + gpu_ptr<DataType>(in), + gpu_ptr<int>(target), + gpu_ptr<float>(out), + axis_size); + }); + }); +} + +void CrossEntropyVJP::eval_gpu( + const std::vector<array>& inputs, + std::vector<array>& outputs) { + nvtx3::scoped_range r("CrossEntropyVJP::eval_gpu"); + assert(inputs.size() == 4); // logits, target, loss, cotangent + auto& s = stream(); + auto& out = outputs[0]; + auto& encoder = cu::get_command_encoder(s); + auto ensure_row_contiguous = [&s, &encoder](const array& x) { + if (x.flags().row_contiguous) { + return x; + } else { + array x_copy = contiguous_copy_gpu(x, s); + encoder.add_temporary(x_copy); + return x_copy; + } + }; + + auto check_input = [&s](const array& x, bool& copied) { + if (x.flags().row_contiguous) { + copied = false; + return x; + } + copied = true; + return contiguous_copy_gpu(x, s); + }; + bool donate_x = inputs[0].is_donatable(); + bool copied; + auto in = check_input(inputs[0], copied); // [n_rows, V] + donate_x |= copied; + auto target = ensure_row_contiguous(inputs[1]); // [n_rows,] + auto loss = ensure_row_contiguous(inputs[2]); // [n_rows,] fp32 + auto cotan = ensure_row_contiguous(inputs[3]); // [n_rows,] fp32 + if (donate_x) { + out.copy_shared_buffer(in); + } else { + out.set_data(cu::malloc_async(out.nbytes(), encoder)); // [n_rows, V] + } + + int axis_size = in.shape().back(); + int n_rows = in.data_size() / axis_size; + + encoder.set_input_array(in); + encoder.set_input_array(target); + encoder.set_input_array(loss); + encoder.set_input_array(cotan); + encoder.set_output_array(out); + dispatch_float_types(in.dtype(), "cross_entropy_vjp", [&](auto type_tag) { + using DataType = cuda_type_t<MLX_GET_TYPE(type_tag)>; + constexpr int N_READS = 16 / sizeof(DataType); + dispatch_block_dim(cuda::ceil_div(axis_size, N_READS), [&](auto block_dim) { + auto kernel = cu::cross_entropy_vjp<DataType, block_dim(), N_READS>; + encoder.add_kernel_node( + kernel, + n_rows, + block_dim(), + gpu_ptr<DataType>(in), + gpu_ptr<int>(target), + gpu_ptr<float>(loss), + gpu_ptr<float>(cotan), + gpu_ptr<DataType>(out), + axis_size); + }); + }); +} + +} // namespace fast + +} // namespace mlx::core
diff --git ml-explore/mlx/mlx/backend/cuda/cuda_utils.h Layr-Labs/mlx/mlx/backend/cuda/cuda_utils.h index 7bae911d265a50568e29136275d4f0ec8314671f..f8a234ee653622deb3f3de9c67f52d1b03e70d97 100644 --- ml-explore/mlx/mlx/backend/cuda/cuda_utils.h +++ Layr-Labs/mlx/mlx/backend/cuda/cuda_utils.h @@ -50,6 +50,12 @@ handle_ = nullptr; } }   + Handle release() { + Handle handle = handle_; + handle_ = nullptr; + return handle; + } + operator Handle() const { return handle_; }
diff --git ml-explore/mlx/mlx/backend/cuda/custom_kernel.cpp Layr-Labs/mlx/mlx/backend/cuda/custom_kernel.cpp index 9b5bd38b7f2f3628fe3b354b2e648a7b6185c79d..c230656d80a7271e24c85dc678c3337709e87f56 100644 --- ml-explore/mlx/mlx/backend/cuda/custom_kernel.cpp +++ Layr-Labs/mlx/mlx/backend/cuda/custom_kernel.cpp @@ -310,9 +310,11 @@ // 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_)); cu::JitModule& mod = cu::get_jit_module( encoder.device(), - name_, + module_name, [&]() { return std::make_tuple( is_precompiled_, source_, std::vector{kernel_name});
diff --git ml-explore/mlx/mlx/backend/cuda/delayload.cpp Layr-Labs/mlx/mlx/backend/cuda/delayload.cpp index aba7566c5bb0eb30a510e970651515ca262aca65..7092d4497cf7953f93315a02809fcdfb3d6f9cc2 100644 --- ml-explore/mlx/mlx/backend/cuda/delayload.cpp +++ Layr-Labs/mlx/mlx/backend/cuda/delayload.cpp @@ -20,23 +20,31 @@ return fs::absolute(current_binary_dir() / relative); }   inline fs::path cublas_dir() { - return cuda_bin_dir() ? fs::path(cuda_bin_dir()) - : relative_to_current_binary("../nvidia/cublas/bin"); + if (const char* dir = cuda_bin_dir()) { + return fs::path(dir); + } + return relative_to_current_binary("../nvidia/cublas/bin"); }   fs::path load_nvrtc() { - fs::path nvrtc_dir = cuda_bin_dir() - ? fs::path(cuda_bin_dir()) - : relative_to_current_binary("../nvidia/cuda_nvrtc/bin"); + fs::path nvrtc_dir; + if (const char* dir = cuda_bin_dir()) { + nvrtc_dir = fs::path(dir); + } else { + nvrtc_dir = relative_to_current_binary("../nvidia/cuda_nvrtc/bin"); + } // Internally nvrtc loads some libs dynamically, add to search dirs. ::AddDllDirectory(nvrtc_dir.c_str()); return nvrtc_dir; }   fs::path load_cudnn() { - fs::path cudnn_dir = cudnn_bin_dir() - ? fs::path(cudnn_bin_dir()) - : relative_to_current_binary("../nvidia/cudnn/bin"); + fs::path cudnn_dir; + if (const char* dir = cudnn_bin_dir()) { + cudnn_dir = fs::path(dir); + } else { + cudnn_dir = relative_to_current_binary("../nvidia/cudnn/bin"); + } // Must load cudnn_graph64_9.dll before locating symbols, otherwise We would // get errors like "Invalid handle. Cannot load symbol cudnnCreate". for (const auto& dll : fs::directory_iterator(cudnn_dir)) { @@ -66,6 +74,8 @@ mod = ::LoadLibraryW((cublas_dir() / dll).c_str()); } else if (dll.starts_with("nvrtc")) { static auto nvrtc_dir = load_nvrtc(); mod = ::LoadLibraryW((nvrtc_dir / dll).c_str()); + } else if (const char* dir = cuda_bin_dir()) { + mod = ::LoadLibraryW((fs::path(dir) / dll).c_str()); } } return reinterpret_cast<FARPROC>(mod);
diff --git ml-explore/mlx/mlx/backend/cuda/device.cpp Layr-Labs/mlx/mlx/backend/cuda/device.cpp index 30248f556886377d96ce2c08805fb68aef31e6fb..472b3d99fb6f394fd65a7e470b2472c7e07b7ad0 100644 --- ml-explore/mlx/mlx/backend/cuda/device.cpp +++ Layr-Labs/mlx/mlx/backend/cuda/device.cpp @@ -462,6 +462,42 @@ }   void CommandEncoder::commit() { nvtx3::scoped_range r("CommandEncoder::commit"); + try { + commit_impl(); + } catch (...) { + // Clear pending CUDA error first. + cudaGetLastError(); + // Clear states. + clear_graph_state(); + node_count_ = 0; + bytes_in_graph_ = 0; + // Clear graph. + try { + graph_.reset(); + } catch (...) { + // Destroying could fail. + graph_.release(); + } + try { + graph_ = CudaGraph(device_); + } catch (...) { + // Keep the original error. + } + // Re-throw the error. + throw; + } +} + +void CommandEncoder::synchronize() { + CHECK_CUDA_ERROR(cudaStreamSynchronize(stream_)); + auto p = std::make_shared<std::promise<void>>(); + std::future<void> f = p->get_future(); + add_completed_handler([p = std::move(p)]() { p->set_value(); }); + commit(); + f.wait(); +} + +void CommandEncoder::commit_impl() { if (!temporaries_.empty()) { add_completed_handler([temporaries = std::move(temporaries_)]() {}); } @@ -520,13 +556,8 @@ CHECK_CUDA_ERROR(cudaGraphDebugDotPrint(graph_, path.c_str(), 0)); }   // Reset state - from_nodes_.clear(); - to_nodes_.clear(); - graph_deps_key_.clear(); - graph_nodes_key_.clear(); - node_map_.clear(); + clear_graph_state(); graph_ = CudaGraph(device_); - is_graph_updatable_ = true; }   // Put completion handlers in a batch. @@ -535,13 +566,16 @@ node_count_ = 0; bytes_in_graph_ = 0; }   -void CommandEncoder::synchronize() { - CHECK_CUDA_ERROR(cudaStreamSynchronize(stream_)); - auto p = std::make_shared<std::promise<void>>(); - std::future<void> f = p->get_future(); - add_completed_handler([p = std::move(p)]() { p->set_value(); }); - commit(); - f.wait(); +void CommandEncoder::clear_graph_state() { + from_nodes_.clear(); + to_nodes_.clear(); + graph_deps_key_.clear(); + graph_nodes_key_.clear(); + node_map_.clear(); + active_deps_.clear(); + active_outputs_.clear(); + concurrent_nodes_.clear(); + is_graph_updatable_ = true; }   Device& device(int cuda_device) {
diff --git ml-explore/mlx/mlx/backend/cuda/device.h Layr-Labs/mlx/mlx/backend/cuda/device.h index 15d75082e922d0d3077ac69a0ca2c144dee36bd0..198f0b5ad85f185e22c721b6dcaaa3844603f930 100644 --- ml-explore/mlx/mlx/backend/cuda/device.h +++ Layr-Labs/mlx/mlx/backend/cuda/device.h @@ -138,6 +138,8 @@ std::string node_type; std::string id; };   + void commit_impl(); + void clear_graph_state(); void insert_graph_dependencies(GraphNode node); void insert_graph_dependencies(std::vector<GraphNode> nodes);
diff --git ml-explore/mlx/mlx/backend/cuda/device/binary_ops.cuh Layr-Labs/mlx/mlx/backend/cuda/device/binary_ops.cuh index b0b7962807035793d9840cc408fa82e226e98858..4368864465889f1615ef03a94c5e28303f1a7d5d 100644 --- ml-explore/mlx/mlx/backend/cuda/device/binary_ops.cuh +++ Layr-Labs/mlx/mlx/backend/cuda/device/binary_ops.cuh @@ -17,9 +17,18 @@ struct FloorDivide { template <typename T> __device__ T operator()(T x, T y) { if constexpr (cuda::std::is_integral_v<T>) { + auto q = x / y; + if constexpr (cuda::std::is_signed_v<T>) { + if (x % y != 0 && (x < 0) != (y < 0)) { + q -= 1; + } + } + return q; + } else if constexpr (is_complex_v<T>) { + // Complex is not supported, simply make compiler happy. return x / y; } else { - return cuda::std::trunc(x / y); + return cuda::std::floor(x / y); } } };
diff --git ml-explore/mlx/mlx/backend/cuda/dirs.cpp Layr-Labs/mlx/mlx/backend/cuda/dirs.cpp index a9d33b4790b3589dd0bbf1da669d56b1e0645bb2..bd24853dd83247e2ee3c912438d43608f0fb6d07 100644 --- ml-explore/mlx/mlx/backend/cuda/dirs.cpp +++ Layr-Labs/mlx/mlx/backend/cuda/dirs.cpp @@ -1,6 +1,24 @@ // Copyright © 2026 Apple Inc.   +#include "mlx/backend/common/utils.h" + +#include <filesystem> +#include <string> + namespace mlx::core::cu { +namespace { + +namespace fs = std::filesystem; + +std::string resolve_bin_dir(const char* dir) { + fs::path path(dir); + if (path.is_absolute()) { + return path.string(); + } + return fs::absolute(current_binary_dir() / path).string(); +} + +} // namespace   const char* cccl_dir() { #if defined(MLX_CCCL_DIR) @@ -12,7 +30,8 @@ }   const char* cuda_bin_dir() { #if defined(MLX_CUDA_BIN_DIR) - return MLX_CUDA_BIN_DIR; + static const std::string dir = resolve_bin_dir(MLX_CUDA_BIN_DIR); + return dir.c_str(); #else return nullptr; #endif @@ -20,7 +39,8 @@ }   const char* cudnn_bin_dir() { #if defined(MLX_CUDNN_BIN_DIR) - return MLX_CUDNN_BIN_DIR; + static const std::string dir = resolve_bin_dir(MLX_CUDNN_BIN_DIR); + return dir.c_str(); #else return nullptr; #endif
diff --git ml-explore/mlx/mlx/backend/cuda/event.cu Layr-Labs/mlx/mlx/backend/cuda/event.cu index b73937ec38dac7aa8af6ffb5a46fccc923ed018a..d3b6f97f5d576ece525e8ce9695540e35e56ced3 100644 --- ml-explore/mlx/mlx/backend/cuda/event.cu +++ Layr-Labs/mlx/mlx/backend/cuda/event.cu @@ -113,10 +113,7 @@ void CudaEvent::init_pool() { cuda_event_pool(); }   -// Wraps CudaEvent with a few features: -// 1. The class can be copied. -// 2. Make wait/record work with CPU streams. -// 3. Add checks for waiting on un-recorded event. +// Wraps CudaEvent so it can be copied. class CopyableCudaEvent { public: explicit CopyableCudaEvent(Device& d) @@ -126,32 +123,24 @@ d, cudaEventDisableTiming | cudaEventBlockingSync)) {}   void wait() { + check_recorded(); event_->wait(); }   void wait(Stream s) { - if (s.device == mlx::core::Device::cpu) { - scheduler::enqueue(s, [*this]() mutable { - check_recorded(); - event_->wait(); - }); - } else { - check_recorded(); - auto& encoder = cu::get_command_encoder(s); - encoder.commit(); - event_->wait(encoder.stream()); - } + assert(s.device == mlx::core::Device::gpu); + check_recorded(); + auto& encoder = cu::get_command_encoder(s); + encoder.commit(); + event_->wait(encoder.stream()); }   void record(Stream s) { - if (s.device == mlx::core::Device::cpu) { - throw std::runtime_error("CudaEvent can not wait on CPU stream."); - } else { - auto& encoder = cu::get_command_encoder(s); - encoder.commit(); - event_->record(encoder.stream()); - recorded_ = true; - } + assert(s.device == mlx::core::Device::gpu); + auto& encoder = cu::get_command_encoder(s); + encoder.commit(); + event_->record(encoder.stream()); + recorded_ = true; }   bool is_signaled() const { @@ -213,6 +202,11 @@ }(); return coherency; }   +const CudaStream& signal_stream() { + static CudaStream stream(device(0)); + return stream; +} + AtomicEvent::AtomicEvent(Device& d) { void* buf; cudaError_t (*cuda_free)(void*); @@ -264,14 +258,11 @@ }   void AtomicEvent::wait(Stream s, uint32_t value) { nvtx3::scoped_range r("cu::AtomicEvent::wait(s)"); - if (s.device == mlx::core::Device::cpu) { - scheduler::enqueue(s, [*this, value]() mutable { wait(value); }); - } else { - auto& encoder = get_command_encoder(s); - encoder.commit(); - wait(encoder.stream(), value); - encoder.add_completed_handler([buf = buf_]() {}); - } + assert(s.device == mlx::core::Device::gpu); + auto& encoder = get_command_encoder(s); + encoder.commit(); + wait(encoder.stream(), value); + encoder.add_completed_handler([buf = buf_]() {}); }   void AtomicEvent::signal(uint32_t value) { @@ -289,17 +280,11 @@ }   void AtomicEvent::signal(Stream s, uint32_t value) { nvtx3::scoped_range r("cu::AtomicEvent::signal(s)"); - if (s.device == mlx::core::Device::cpu) { - // Signal through a GPU stream so the atomic is updated in GPU - updating - // the atomic in CPU sometimes does not get GPU notified. - scheduler::enqueue( - s, [*this, value]() mutable { signal(signal_stream(), value); }); - } else { - auto& encoder = get_command_encoder(s); - encoder.commit(); - signal(encoder.stream(), value); - encoder.add_completed_handler([buf = buf_]() {}); - } + assert(s.device == mlx::core::Device::gpu); + auto& encoder = get_command_encoder(s); + encoder.commit(); + signal(encoder.stream(), value); + encoder.add_completed_handler([buf = buf_]() {}); }   bool AtomicEvent::is_signaled(uint32_t val) const { @@ -319,9 +304,21 @@ return val; } }   -const CudaStream& AtomicEvent::signal_stream() { - static CudaStream stream(device(0)); - return stream; +/////////////////////////////////////////////////////////////////////////////// +// EventImpl implementations +/////////////////////////////////////////////////////////////////////////////// + +void EventImpl::ensure_created(Stream s, uint64_t signal_value) { + if (is_created()) { + return; + } + auto& d = cu::device(s.device); + if (s.device == mlx::core::Device::cpu || signal_value > 1) { + nvtx3::mark("Using slow AtomicEvent"); + atomic = std::make_unique<cu::AtomicEvent>(d); + } else { + cuda = std::make_unique<cu::CopyableCudaEvent>(d); + } }   } // namespace cu @@ -330,86 +327,85 @@ /////////////////////////////////////////////////////////////////////////////// // Event implementations ///////////////////////////////////////////////////////////////////////////////   -namespace { - -struct EventImpl { - // CudaEvent is preferred when possible because it is fast, however we have - // to fallback to AtomicEvent in following cases: - // 1. the event is used to wait/signal a cpu stream; - // 2. signal value other than 1 has been specified. - std::unique_ptr<cu::CopyableCudaEvent> cuda; - std::unique_ptr<cu::AtomicEvent> atomic; - - bool is_created() const { - return cuda || atomic; - } - - void ensure_created(Stream s, uint64_t signal_value) { - if (is_created()) { - return; - } - auto& d = cu::device(s.device); - if (s.device == mlx::core::Device::cpu || signal_value > 1) { - nvtx3::mark("Using slow AtomicEvent"); - atomic = std::make_unique<cu::AtomicEvent>(d); - } else { - cuda = std::make_unique<cu::CopyableCudaEvent>(d); - } - } -}; - -} // namespace - Event::Event(Stream s) : stream_(s) { - event_ = std::shared_ptr<void>( - new EventImpl(), [](void* ptr) { delete static_cast<EventImpl*>(ptr); }); + event_ = std::make_shared<cu::EventImpl>(); }   void Event::wait() { - auto* event = static_cast<EventImpl*>(event_.get()); - assert(event->is_created()); - if (event->cuda) { + check_error(); + auto& event = cast<cu::EventImpl>(); + assert(event.is_created()); + if (event.cuda) { assert(value() == 1); - event->cuda->wait(); + event.cuda->wait(); } else { - event->atomic->wait(value()); + event.atomic->wait(value()); } CHECK_CUDA_ERROR(cudaPeekAtLastError()); + check_error(); }   void Event::wait(Stream s) { - auto* event = static_cast<EventImpl*>(event_.get()); - assert(event->is_created()); - if (event->cuda) { + auto& event = cast<cu::EventImpl>(); + assert(event.is_created()); + if (event.cuda) { assert(value() == 1); - event->cuda->wait(s); + if (s.device == mlx::core::Device::cpu) { + scheduler::wait_event(s, *this, [value = value()](Event& self) { + self.cast<cu::EventImpl>().cuda->wait(); + }); + } else { + event.cuda->wait(s); + } } else { - event->atomic->wait(s, value()); + if (s.device == mlx::core::Device::cpu) { + scheduler::wait_event(s, *this, [value = value()](Event& self) { + self.cast<cu::EventImpl>().atomic->wait(value); + }); + } else { + event.atomic->wait(s, value()); + } } }   void Event::signal(Stream s) { - auto* event = static_cast<EventImpl*>(event_.get()); - event->ensure_created(s, value()); - if (event->cuda) { + auto& event = cast<cu::EventImpl>(); + event.ensure_created(s, value()); + if (event.cuda) { assert(value() == 1); - event->cuda->record(s); + if (s.device == mlx::core::Device::cpu) { + throw std::runtime_error("CudaEvent can not wait on CPU stream."); + } else { + event.cuda->record(s); + } } else { - event->atomic->signal(s, value()); + if (s.device == mlx::core::Device::cpu) { + // Signal through a GPU stream so the atomic is updated in GPU - updating + // the atomic in CPU sometimes does not get GPU notified. + scheduler::signal_event(s, *this, [value = value()](Event& self) { + self.cast<cu::EventImpl>().atomic->signal(cu::signal_stream(), value); + }); + } else { + event.atomic->signal(s, value()); + } } }   bool Event::is_signaled() const { - auto* event = static_cast<EventImpl*>(event_.get()); - if (!event->is_created()) { + auto& event = cast<cu::EventImpl>(); + if (!event.is_created()) { return false; } - if (event->cuda) { + if (event.cuda) { assert(value() == 1); - return event->cuda->is_signaled(); + return event.cuda->is_signaled(); } else { - return event->atomic->is_signaled(value()); + return event.atomic->is_signaled(value()); } +} + +std::atomic<Error*>& Event::error() { + return cast<cu::EventImpl>().error; }   } // namespace mlx::core
diff --git ml-explore/mlx/mlx/backend/cuda/event.h Layr-Labs/mlx/mlx/backend/cuda/event.h index 53afeb011748f07fd2823157760ecbe0b993abee..fdeb6a0e7819f00497b4b218ec1617ab07225e91 100644 --- ml-explore/mlx/mlx/backend/cuda/event.h +++ Layr-Labs/mlx/mlx/backend/cuda/event.h @@ -13,6 +13,7 @@ #include <cuda/atomic>   namespace mlx::core::cu {   +class CopyableCudaEvent; class Device;   // RAII-managed move-only wrapper of cudaEvent_t. @@ -66,14 +67,29 @@ bool is_signaled(uint32_t value) const; uint32_t value() const;   private: - const CudaStream& signal_stream(); - uint32_t* ptr() const { return static_cast<uint32_t*>(buf_.get()); }   bool coherent_; std::shared_ptr<void> buf_; +}; + +struct EventImpl { + std::atomic<Error*> error; + + // CudaEvent is preferred when possible because it is fast, however we have + // to fallback to AtomicEvent in following cases: + // 1. the event is used to wait/signal a cpu stream; + // 2. signal value other than 1 has been specified. + std::unique_ptr<cu::CopyableCudaEvent> cuda; + std::unique_ptr<cu::AtomicEvent> atomic; + + bool is_created() const { + return cuda || atomic; + } + + void ensure_created(Stream s, uint64_t signal_value); };   } // namespace mlx::core::cu
diff --git ml-explore/mlx/mlx/backend/cuda/fence.cpp Layr-Labs/mlx/mlx/backend/cuda/fence.cpp index c6a41f0e60ca4f3100eac335ad0fb66692e2672a..3a3acdba09e059f2960557fc71d15d8270f1b987 100644 --- ml-explore/mlx/mlx/backend/cuda/fence.cpp +++ Layr-Labs/mlx/mlx/backend/cuda/fence.cpp @@ -9,22 +9,23 @@ namespace mlx::core {   struct FenceImpl { uint32_t count; - cu::AtomicEvent event; + Event event; + + FenceImpl(uint32_t count, Stream s) : count(count), event(s) {} };   Fence::Fence(Stream s) { - fence_ = std::shared_ptr<void>( - new FenceImpl{0, cu::device(s.device)}, - [](void* ptr) { delete static_cast<FenceImpl*>(ptr); }); + fence_ = std::make_shared<FenceImpl>(0, s); + // Ensure that we use AtomicEvent. + cast<FenceImpl>().event.cast<cu::EventImpl>().ensure_created(s, 2); }   void Fence::wait(Stream s, const array&) { - auto* fence = static_cast<FenceImpl*>(fence_.get()); - fence->event.wait(fence->count); + cast<FenceImpl>().event.wait(); }   void Fence::update(Stream s, const array& a, bool cross_device) { - auto* fence = static_cast<FenceImpl*>(fence_.get()); + auto& f = cast<FenceImpl>(); if (cross_device) { // Move to managed memory if there is a device switch auto& cbuf = @@ -35,8 +36,9 @@ encoder.commit(); cu::allocator().move_to_unified_memory(cbuf, encoder.stream()); } } - fence->count++; - fence->event.signal(s, fence->count); + f.count++; + f.event.set_value(f.count); + f.event.signal(s); }   } // namespace mlx::core
diff --git ml-explore/mlx/mlx/backend/cuda/jit_module.cpp Layr-Labs/mlx/mlx/backend/cuda/jit_module.cpp index 3de1ddb018c636ccaa48a804056adaad19f0f24e..0d107a8e18bce331741feeeb45e2001398e65ffa 100644 --- ml-explore/mlx/mlx/backend/cuda/jit_module.cpp +++ Layr-Labs/mlx/mlx/backend/cuda/jit_module.cpp @@ -49,6 +49,20 @@ #endif return cached_path; }   +// Get the dirname of nvidia python package that contains CUDA headers. +inline const char* cudart_dirname() { +#if CUDART_VERSION < 13000 + return "cuda_runtime"; +#elif CUDART_VERSION < 14000 + return "cu13"; +#else + static_assert( + false, + "Please find out the newest dirname under site-packages/nvidia " + "and add it in this function."); +#endif +} + // Return the --include-path args used for invoking NVRTC. const std::vector<std::string>& include_path_args() { static std::vector<std::string> cached_args = []() { @@ -72,7 +86,7 @@ args.push_back(fmt::format("--include-path={}", path.string())); } // Add path to CUDA runtime headers, try local-installed python package // first and then system-installed headers. - path = root_dir.parent_path() / "nvidia" / "cuda_runtime" / "include"; + path = root_dir.parent_path() / "nvidia" / cudart_dirname() / "include"; if (!std::filesystem::exists(path)) { const char* home = std::getenv("CUDA_HOME"); if (!home) {
diff --git ml-explore/mlx/mlx/backend/cuda/scaled_dot_product_attention.cpp Layr-Labs/mlx/mlx/backend/cuda/scaled_dot_product_attention.cpp index ca411e91c639388679aeadf14368a527e933341d..286d500e2c8ef738233a8274360543722c60c715 100644 --- ml-explore/mlx/mlx/backend/cuda/scaled_dot_product_attention.cpp +++ Layr-Labs/mlx/mlx/backend/cuda/scaled_dot_product_attention.cpp @@ -549,6 +549,33 @@ Stream s);   namespace fast {   +namespace { + +std::tuple<bool, std::string> has_fused_kernel( + const array& q, + const array& k, + const array& v, + bool has_arr_mask, + bool do_causal, + bool output_logsumexp, + Stream s) { + if (s.device != Device::gpu) { + return {false, "the fused kernels require a GPU stream."}; + } + if (!supports_sdpa_cudnn(q, k, v, has_arr_mask, do_causal, s) && + !supports_sdpa_vector(q, k, v, has_arr_mask, output_logsumexp)) { + std::ostringstream msg; + msg << "neither the cuDNN attention nor the vector attention kernel " + << "supports this configuration; got query shape " << q.shape() + << ", key shape " << k.shape() << ", value shape " << v.shape() + << " with dtype " << q.dtype() << "."; + return {false, msg.str()}; + } + return {true, ""}; +} + +} // namespace + bool ScaledDotProductAttention::use_fallback( const array& q, const array& k, @@ -558,13 +585,21 @@ bool has_arr_mask, bool do_causal, bool is_training, bool output_logsumexp, + bool force_fused, Stream s) { - if (s.device == Device::cpu) { - return true; + auto [has_fused, reason] = + has_fused_kernel(q, k, v, has_arr_mask, do_causal, output_logsumexp, s); + if (force_fused) { + if (!has_fused) { + std::ostringstream msg; + msg << "[scaled_dot_product_attention] force_fused=True but no fused " + "kernel is available: " + << reason; + throw std::invalid_argument(msg.str()); + } + return false; } - - return !supports_sdpa_cudnn(q, k, v, has_arr_mask, do_causal, s) && - !supports_sdpa_vector(q, k, v, has_arr_mask, output_logsumexp); + return !has_fused; }   bool ScaledDotProductAttention::supports_bool_mask() {
diff --git ml-explore/mlx/mlx/backend/metal/compiled.cpp Layr-Labs/mlx/mlx/backend/metal/compiled.cpp index cda06143d64bff54cd5ede071e3c68b597c0ff15..e95d7b8d5f1ea879bb12a067bfe04cde7739fd9b 100644 --- ml-explore/mlx/mlx/backend/metal/compiled.cpp +++ Layr-Labs/mlx/mlx/backend/metal/compiled.cpp @@ -208,7 +208,7 @@ os += fmt::format( " {0} tmp_{1} = ", get_type_string(x.dtype()), namer.get_name(x)); if (is_static_cast(x.primitive())) { os += fmt::format( - "static_cast<{0}>(tmp_{1});\n", + "cast_to<{0}>(tmp_{1});\n", get_type_string(x.dtype()), namer.get_name(x.inputs()[0])); } else {
diff --git ml-explore/mlx/mlx/backend/metal/conv.cpp Layr-Labs/mlx/mlx/backend/metal/conv.cpp index 926f31f05a7e9dee9af4fa0a16d401841f01526b..5095522616d3dc33cf474506931da8f310a31b21 100644 --- ml-explore/mlx/mlx/backend/metal/conv.cpp +++ Layr-Labs/mlx/mlx/backend/metal/conv.cpp @@ -5,6 +5,7 @@ #include <numeric>   #include "mlx/backend/gpu/copy.h" #include "mlx/backend/gpu/slicing.h" +#include "mlx/backend/metal/binary.h" #include "mlx/backend/metal/device.h" #include "mlx/backend/metal/kernels.h" #include "mlx/backend/metal/kernels/defines.h" @@ -45,6 +46,68 @@ << max_buffer << " bytes."; throw std::runtime_error(msg.str()); } return static_cast<int>(std::min(max_rows, static_cast<size_t>(total_rows))); +} + +inline auto winograd_padded_size(const MLXConvParams<2>& conv_params) { + int64_t pad_h = static_cast<int64_t>(conv_params.iS[0]) + + 2 * static_cast<int64_t>(conv_params.pad[0]); + int64_t pad_w = static_cast<int64_t>(conv_params.iS[1]) + + 2 * static_cast<int64_t>(conv_params.pad[1]); + int padded_h = safe_cast(6 * ceildiv(pad_h - 2, 6) + 2, "conv"); + int padded_w = safe_cast(6 * ceildiv(pad_w - 2, 6) + 2, "conv"); + return std::make_tuple(padded_h, padded_w); +} + +// Return how many rows to compute per each step. +inline int winograd_batch_step( + metal::Device& d, + const array& in, + const MLXConvParams<2>& conv_params) { + int total_n = conv_params.N; + + size_t itemsize = in.itemsize(); + auto [padded_h, padded_w] = winograd_padded_size(conv_params); + int tiles_per_n = + ceildiv(conv_params.oS[0], 6) * ceildiv(conv_params.oS[1], 6); + + // Limit of maximum memory can be used for the step. + size_t working_set = d.mtl_device()->recommendedMaxWorkingSetSize(); + if (int env_ws = env::get_var("MLX_CONV_WINOGRAD_WORKING_SET", 0); + env_ws > 0) { + working_set = env_ws; + } + size_t limit = working_set / 4 * 3; + + // Memory used by inputs. + size_t filt_bytes = + static_cast<size_t>(8 * 8) * conv_params.C * conv_params.O * itemsize; + size_t io_bytes = itemsize * + (static_cast<size_t>(total_n) * conv_params.iS[0] * conv_params.iS[1] * + conv_params.C + + static_cast<size_t>(total_n) * conv_params.oS[0] * conv_params.oS[1] * + conv_params.O); + size_t used = io_bytes + filt_bytes; + size_t budget = limit > used ? limit - used : 0; + + // How many rows to use per step to avoid running over limit. + size_t bytes_per_n = + static_cast<size_t>(padded_h) * padded_w * conv_params.C * itemsize + + static_cast<size_t>(8 * 8) * tiles_per_n * + (conv_params.C + conv_params.O) * itemsize; + auto max_n = static_cast<int64_t>(budget / bytes_per_n); + int safe_n = static_cast<int>(std::min<int64_t>(max_n, total_n)); + if (int forced = env::get_var("MLX_CONV_WINOGRAD_TILE_BATCH", 0); + forced > 0) { + return std::min(forced, safe_n); + } + + // When the budget forces tiling, each tile must carry enough gemm rows to + // amortize fixed cost, below that the implicit gemm fallback is much faster. + constexpr int min_rows_per_tile = 32; + if ((safe_n < total_n) && (safe_n * tiles_per_n < min_rows_per_tile)) { + return 0; + } + return safe_n; }   template <int N> @@ -743,6 +806,213 @@ out.copy_shared_buffer( intermediate, intermediate.strides(), {0}, intermediate.data_size()); }   +void conv_2D_gpu( + const Stream& s, + metal::Device& d, + const array& in_pre, + const array& wt_pre, + array& out, + const std::vector<int>& padding, + const std::vector<int>& wt_strides, + const std::vector<int>& wt_dilation, + const std::vector<int>& in_dilation, + const int groups, + bool flip, + std::vector<array>& copies); + +void small_kd_conv_3D_gpu( + const Stream& s, + metal::Device& d, + const array& in, + const array& wt, + array& out, + const MLXConvParams<3>& conv_params, + std::vector<array>& copies) { + const int H = conv_params.iS[1]; + const int W = conv_params.iS[2]; + const int C = conv_params.C; + const int O = conv_params.O; + const int KD = conv_params.wS[0]; + const int KH = conv_params.wS[1]; + const int KW = conv_params.wS[2]; + const int OD = conv_params.oS[0]; + const int OH = conv_params.oS[1]; + const int OW = conv_params.oS[2]; + + array acc({OD, OH, OW, O}, out.dtype(), nullptr, {}); + for (int kd = 0; kd < KD; ++kd) { + array in_2d({OD, H, W, C}, in.dtype(), nullptr, {}); + in_2d.copy_shared_buffer( + in, + {static_cast<int64_t>(H) * W * C, + static_cast<int64_t>(W) * C, + static_cast<int64_t>(C), + 1}, + {true, true, false}, + static_cast<size_t>(OD) * H * W * C, + static_cast<int64_t>(kd) * H * W * C); + + // The 2D conv only flips the last two kernel axes, so mirror the depth + // axis here when the convolution is flipped. + const int kd_wt = conv_params.flip ? KD - 1 - kd : kd; + + array wt_2d({O, KH, KW, C}, wt.dtype(), nullptr, {}); + wt_2d.copy_shared_buffer( + wt, + {static_cast<int64_t>(KD) * KH * KW * C, + static_cast<int64_t>(KW) * C, + static_cast<int64_t>(C), + 1}, + {false, false, false}, + static_cast<size_t>(O - 1) * KD * KH * KW * C + + static_cast<size_t>(KH) * KW * C, + static_cast<int64_t>(kd_wt) * KH * KW * C); + + array conv_out({OD, OH, OW, O}, out.dtype(), nullptr, {}); + conv_2D_gpu( + s, + d, + in_2d, + wt_2d, + conv_out, + {conv_params.pad[1], conv_params.pad[2]}, + {conv_params.str[1], conv_params.str[2]}, + {conv_params.kdil[1], conv_params.kdil[2]}, + {conv_params.idil[1], conv_params.idil[2]}, + /* groups = */ 1, + conv_params.flip, + copies); + + if (kd == 0) { + acc = conv_out; + } else { + binary_op_gpu_inplace({acc, conv_out}, acc, "Add", s); + copies.push_back(conv_out); + } + } + + // Output shape is [1, OD, OH, OW, O]. + out.copy_shared_buffer( + acc, + {static_cast<int64_t>(OD) * OH * OW * O, + static_cast<int64_t>(OH) * OW * O, + static_cast<int64_t>(OW) * O, + static_cast<int64_t>(O), + 1}, + {true, true, false}, + static_cast<size_t>(OD) * OH * OW * O, + 0); +} + +// A stride-2, kernel-2 transposed convolution has exactly one valid kernel +// phase for every output coordinate. The explicit unfold path nevertheless +// materializes all eight phases and fills seven of them with zeros. Compute +// each phase as a small regular GEMM instead. This path is intentionally +// narrow: other transposed-convolution configurations retain the general +// implementation below. +bool is_stride_two_conv_transpose_3D(const MLXConvParams<3>& p) { + return p.groups == 1 && p.flip && p.str[0] == 1 && p.str[1] == 1 && + p.str[2] == 1 && p.idil[0] == 2 && p.idil[1] == 2 && p.idil[2] == 2 && + p.kdil[0] == 1 && p.kdil[1] == 1 && p.kdil[2] == 1 && p.wS[0] == 2 && + p.wS[1] == 2 && p.wS[2] == 2 && p.pad[0] == 1 && p.pad[1] == 1 && + p.pad[2] == 1 && static_cast<int64_t>(p.oS[0]) == 2LL * p.iS[0] && + static_cast<int64_t>(p.oS[1]) == 2LL * p.iS[1] && + static_cast<int64_t>(p.oS[2]) == 2LL * p.iS[2]; +} + +void stride_two_conv_transpose_3D_gpu( + const Stream& s, + metal::Device& d, + const array& in, + const array& wt, + array& out, + const MLXConvParams<3>& p, + std::vector<array>& copies) { + constexpr int kernel_volume = 8; + const int C = p.C; + const int O = p.O; + + // The input and weight are contiguous by the time this helper is called. + // Every phase covers the complete input volume; the phase bit only selects + // the interleaved output coordinates and the corresponding weight slice. + for (int phase = 0; phase < kernel_volume; ++phase) { + const int pd = (phase >> 2) & 1; + const int ph = (phase >> 1) & 1; + const int pw = phase & 1; + const int D = p.iS[0]; + const int H = p.iS[1]; + const int W = p.iS[2]; + + const int M = safe_cast(static_cast<int64_t>(p.N) * D * H * W, "conv"); + + // The weight is [O, 2, 2, 2, C]. Present one spatial phase as a [C, O] + // matrix with the layout expected by a transposed Steel GEMM. + array wt_phase({C, O}, wt.dtype(), nullptr, {}); + array::Flags wt_flags = wt.flags(); + wt_flags.contiguous = false; + wt_flags.row_contiguous = false; + wt_flags.col_contiguous = true; + wt_phase.copy_shared_buffer( + wt, + {1, wt.strides(0)}, + wt_flags, + wt.data_size(), + static_cast<int64_t>(pd * 4 + ph * 2 + pw) * C); + + array in_matrix({M, C}, in.dtype(), nullptr, {}); + in_matrix.copy_shared_buffer(in, {C, 1}, in.flags(), in.data_size()); + + array phase_out({M, O}, out.dtype(), nullptr, {}); + phase_out.set_data(allocator::malloc(phase_out.nbytes())); + + std::vector<array> gemm_copies = {in_matrix, wt_phase}; + steel_matmul( + s, + d, + /* a = */ in_matrix, + /* b = */ wt_phase, + /* out = */ phase_out, + /* M = */ M, + /* N = */ O, + /* K = */ C, + /* batch_size_out = */ 1, + /* lda = */ C, + /* ldb = */ kernel_volume * C, + /* a_transposed = */ false, + /* b_transposed = */ true, + /* copies = */ gemm_copies); + + Shape phase_out_shape{p.N, D, H, W, O}; + array phase_out_nd(phase_out_shape, out.dtype(), nullptr, {}); + phase_out_nd.copy_shared_buffer( + phase_out, + make_contiguous_strides(phase_out_shape), + phase_out.flags(), + phase_out.data_size()); + + Strides out_phase_strides = out.strides(); + out_phase_strides[1] *= 2; + out_phase_strides[2] *= 2; + out_phase_strides[3] *= 2; + array::Flags out_phase_flags = out.flags(); + out_phase_flags.contiguous = false; + out_phase_flags.row_contiguous = false; + out_phase_flags.col_contiguous = false; + array out_phase_view(phase_out_shape, out.dtype(), nullptr, {}); + out_phase_view.copy_shared_buffer( + out, + out_phase_strides, + out_phase_flags, + out.data_size(), + pd * out.strides(1) + ph * out.strides(2) + pw * out.strides(3)); + + copy_gpu_inplace(phase_out_nd, out_phase_view, CopyType::GeneralGeneral, s); + copies.push_back(phase_out); + copies.push_back(phase_out_nd); + copies.push_back(out_phase_view); + } +} + void dispatch_conv_3D_gpu( const Stream& s, metal::Device& d, @@ -773,6 +1043,20 @@ out.set_data(allocator::malloc(out.nbytes())); auto in = ensure_row_contiguous(in_pre, d, s); auto wt = ensure_row_contiguous(wt_pre, d, s);   + if (is_stride_two_conv_transpose_3D(conv_params)) { + return stride_two_conv_transpose_3D_gpu( + s, d, in, wt, out, conv_params, copies); + } + + // Decompose 3D conv to per-frame 2D convs + constexpr int kSmallKdLimit3D = 7; + if (is_idil_one && mod16_channels && conv_params.groups == 1 && + conv_params.N == 1 && conv_params.wS[0] <= kSmallKdLimit3D && + conv_params.str[0] == 1 && conv_params.kdil[0] == 1 && + conv_params.pad[0] == 0) { + return small_kd_conv_3D_gpu(s, d, in, wt, out, conv_params, copies); + } + // Perform the implicit gemm if (is_idil_one && mod16_channels) { return implicit_gemm_conv_3D_gpu(s, d, in, wt, out, conv_params); @@ -793,77 +1077,17 @@ const array& in, const array& wt, array& out, const MLXConvParams<2>& conv_params, - std::vector<array>& copies_w) { - Shape padded_shape = { - conv_params.N, - conv_params.iS[0] + 2 * conv_params.pad[0], - conv_params.iS[1] + 2 * conv_params.pad[1], - conv_params.C}; - - padded_shape[1] = 6 * ((padded_shape[1] - 2 + 5) / 6) + 2; - padded_shape[2] = 6 * ((padded_shape[2] - 2 + 5) / 6) + 2; - - array in_padded(std::move(padded_shape), in.dtype(), nullptr, {}); - - // Fill with zeros - array zero_arr = array(0, in.dtype()); - fill_gpu(zero_arr, in_padded, s); - copies_w.push_back(zero_arr); - - // Pick input slice from padded - size_t data_offset = conv_params.pad[0] * in_padded.strides()[1] + - conv_params.pad[1] * in_padded.strides()[2]; - array in_padded_slice(in.shape(), in_padded.dtype(), nullptr, {}); - in_padded_slice.copy_shared_buffer( - in_padded, - in_padded.strides(), - in_padded.flags(), - in_padded_slice.size(), - data_offset); - - // Copy input values into the slice - copy_gpu_inplace(in, in_padded_slice, CopyType::GeneralGeneral, s); - - copies_w.push_back(in_padded_slice); - copies_w.push_back(in_padded); - - MLXConvParams<2> conv_params_updated{ - /* const int N = */ static_cast<int>(in_padded.shape(0)), - /* const int C = */ static_cast<int>(in_padded.shape(3)), - /* const int O = */ static_cast<int>(wt.shape(0)), - /* const int iS[NDIM] = */ - {static_cast<int>(in_padded.shape(1)), - static_cast<int>(in_padded.shape(2))}, - /* const int wS[NDIM] = */ - {static_cast<int>(wt.shape(1)), static_cast<int>(wt.shape(2))}, - /* const int oS[NDIM] = */ - {static_cast<int>(out.shape(1)), static_cast<int>(out.shape(2))}, - /* const int str[NDIM] = */ {1, 1}, - /* const int pad[NDIM] = */ {0, 0}, - /* const int kdil[NDIM] = */ {1, 1}, - /* const int idil[NDIM] = */ {1, 1}, - /* const size_t in_strides[NDIM + 2] = */ - {in_padded.strides()[0], - in_padded.strides()[1], - in_padded.strides()[2], - in_padded.strides()[3]}, - /* const size_t wt_strides[NDIM + 2] = */ - {wt.strides()[0], wt.strides()[1], wt.strides()[2], wt.strides()[3]}, - /* const size_t out_strides[NDIM + 2] = */ - {out.strides()[0], out.strides()[1], out.strides()[2], out.strides()[3]}, - /* const int groups = */ 1, - /* const bool flip = */ false, - }; - + std::vector<array>& copies_w, + int n_step) { int O_c = conv_params.O; int C_c = conv_params.C; + auto [padded_h, padded_w] = winograd_padded_size(conv_params);   - int N_tiles_n = conv_params.N; - int N_tiles_h = (conv_params.oS[0] + 5) / 6; - int N_tiles_w = (conv_params.oS[1] + 5) / 6; - int N_tiles = N_tiles_n * N_tiles_h * N_tiles_w; + int N_tiles_h = ceildiv(conv_params.oS[0], 6); + int N_tiles_w = ceildiv(conv_params.oS[1], 6); + int tiles_per_n = N_tiles_h * N_tiles_w;   - // Do filter transform + // Do filter transform. Shape filt_wg_shape = {8 * 8, conv_params.C, conv_params.O}; array filt_wg(std::move(filt_wg_shape), wt.dtype(), nullptr, {}); filt_wg.set_data(allocator::malloc(filt_wg.nbytes())); @@ -895,88 +1119,187 @@ compute_encoder.dispatch_threadgroups(grid_dims, group_dims); }   - // Do input transform - Shape inp_wg_shape = {8 * 8, N_tiles, conv_params.C}; - array inp_wg(std::move(inp_wg_shape), in.dtype(), nullptr, {}); + // Scratch space reused by every batch tile. + array inp_wg({8 * 8, n_step * tiles_per_n, C_c}, in.dtype(), nullptr, {}); inp_wg.set_data(allocator::malloc(inp_wg.nbytes())); copies_w.push_back(inp_wg); - { - int bc = 32; - int wm = 2; - int wn = 2; - std::string kname; - kname.reserve(32); - concatenate( - kname, - "winograd_conv_2d_input_transform_", - type_to_name(out), - "_bc", - bc); - auto& compute_encoder = metal::get_command_encoder(s); - auto kernel = d.get_kernel(kname); - compute_encoder.set_compute_pipeline_state(kernel); + + array out_wg({8 * 8, n_step * tiles_per_n, O_c}, in.dtype(), nullptr, {}); + out_wg.set_data(allocator::malloc(out_wg.nbytes())); + copies_w.push_back(out_wg); + + array in_padded({n_step, padded_h, padded_w, C_c}, in.dtype(), nullptr, {}); + copies_w.push_back(in_padded); + + // Fill padding with zeros. + array zero_arr = array(0, in.dtype()); + fill_gpu(zero_arr, in_padded, s); + copies_w.push_back(zero_arr); + + int64_t pad_offset = + static_cast<int64_t>(conv_params.pad[0]) * in_padded.strides()[1] + + static_cast<int64_t>(conv_params.pad[1]) * in_padded.strides()[2]; + + // Loop over all rows. + for (int n_offset = 0; n_offset < conv_params.N; n_offset += n_step) { + int tile_n = std::min(n_step, conv_params.N - n_offset); + int N_tiles = tile_n * tiles_per_n; + + // Views for current step. + array in_tile( + {tile_n, conv_params.iS[0], conv_params.iS[1], C_c}, + in.dtype(), + nullptr, + {}); + in_tile.copy_shared_buffer( + in, + in.strides(), + in.flags(), + in_tile.size(), + n_offset * in.strides()[0]); + + array out_tile( + {tile_n, conv_params.oS[0], conv_params.oS[1], O_c}, + out.dtype(), + nullptr, + {}); + out_tile.copy_shared_buffer( + out, + out.strides(), + out.flags(), + out_tile.size(), + n_offset * out.strides()[0]); + + array in_padded_slice(in_tile.shape(), in_padded.dtype(), nullptr, {}); + in_padded_slice.copy_shared_buffer( + in_padded, + in_padded.strides(), + in_padded.flags(), + in_padded_slice.size(), + pad_offset); + + // Copy input values into the slice. + copy_gpu_inplace(in_tile, in_padded_slice, CopyType::GeneralGeneral, s); + copies_w.push_back(in_padded_slice); + + MLXConvParams<2> conv_params_updated{ + /* const int N = */ tile_n, + /* const int C = */ C_c, + /* const int O = */ O_c, + /* const int iS[NDIM] = */ {padded_h, padded_w}, + /* const int wS[NDIM] = */ + {static_cast<int>(wt.shape(1)), static_cast<int>(wt.shape(2))}, + /* const int oS[NDIM] = */ + {static_cast<int>(out.shape(1)), static_cast<int>(out.shape(2))}, + /* const int str[NDIM] = */ {1, 1}, + /* const int pad[NDIM] = */ {0, 0}, + /* const int kdil[NDIM] = */ {1, 1}, + /* const int idil[NDIM] = */ {1, 1}, + /* const size_t in_strides[NDIM + 2] = */ + {in_padded.strides()[0], + in_padded.strides()[1], + in_padded.strides()[2], + in_padded.strides()[3]}, + /* const size_t wt_strides[NDIM + 2] = */ + {wt.strides()[0], wt.strides()[1], wt.strides()[2], wt.strides()[3]}, + /* const size_t out_strides[NDIM + 2] = */ + {out.strides()[0], + out.strides()[1], + out.strides()[2], + out.strides()[3]}, + /* const int groups = */ 1, + /* const bool flip = */ false, + }; + + // Do input transform, result layout is (8 x 8 x N_tiles x channels). + { + int bc = 32; + int wm = 2; + int wn = 2; + std::string kname; + kname.reserve(32); + concatenate( + kname, + "winograd_conv_2d_input_transform_", + type_to_name(out), + "_bc", + bc); + auto& compute_encoder = metal::get_command_encoder(s); + auto kernel = d.get_kernel(kname); + compute_encoder.set_compute_pipeline_state(kernel);   - compute_encoder.set_input_array(in_padded, 0); - compute_encoder.set_output_array(inp_wg, 1); + compute_encoder.set_input_array(in_padded, 0); + compute_encoder.set_output_array(inp_wg, 1);   - compute_encoder.set_bytes(conv_params_updated, 2); + compute_encoder.set_bytes(conv_params_updated, 2);   - MTL::Size group_dims = MTL::Size(32, wn, wm); - MTL::Size grid_dims = MTL::Size(N_tiles_w, N_tiles_h, N_tiles_n); + MTL::Size group_dims = MTL::Size(32, wn, wm); + MTL::Size grid_dims = MTL::Size(N_tiles_w, N_tiles_h, tile_n);   - compute_encoder.dispatch_threadgroups(grid_dims, group_dims); - } + compute_encoder.dispatch_threadgroups(grid_dims, group_dims); + }   - // Do batched gemm - Shape out_wg_shape = {8 * 8, N_tiles, conv_params.O}; - array out_wg(std::move(out_wg_shape), in.dtype(), nullptr, {}); - out_wg.set_data(allocator::malloc(out_wg.nbytes())); - copies_w.push_back(out_wg); - { - std::vector<array> empty_copies; - steel_matmul( - s, - d, - /*a = */ inp_wg, - /*b = */ filt_wg, - /*c = */ out_wg, - /*M = */ N_tiles, - /*N = */ conv_params.O, - /*K = */ conv_params.C, - /*batch_size_out = */ 8 * 8, - /*a_cols = */ conv_params.C, - /*b_cols = */ conv_params.O, - /*a_transposed = */ false, - /*b_transposed = */ false, - /*copies = */ empty_copies); - } + // Do batched gemm. + { + array inp_wg_tile({8 * 8, N_tiles, C_c}, inp_wg.dtype(), nullptr, {}); + inp_wg_tile.copy_shared_buffer( + inp_wg, + {static_cast<int64_t>(N_tiles) * C_c, C_c, 1}, + inp_wg.flags(), + inp_wg_tile.size());   - // Do output transform - { - int bc = 32; - int wm = 2; - int wn = 2; - std::string kname; - kname.reserve(32); - concatenate( - kname, - "winograd_conv_2d_output_transform_", - type_to_name(out), - "_bo", - bc); - auto& compute_encoder = metal::get_command_encoder(s); - auto kernel = d.get_kernel(kname); - compute_encoder.set_compute_pipeline_state(kernel); + array out_wg_tile({8 * 8, N_tiles, O_c}, out_wg.dtype(), nullptr, {}); + out_wg_tile.copy_shared_buffer( + out_wg, + {static_cast<int64_t>(N_tiles) * O_c, O_c, 1}, + out_wg.flags(), + out_wg_tile.size());   - compute_encoder.set_input_array(out_wg, 0); - compute_encoder.set_output_array(out, 1); + std::vector<array> empty_copies; + steel_matmul( + s, + d, + /*a = */ inp_wg_tile, + /*b = */ filt_wg, + /*c = */ out_wg_tile, + /*M = */ N_tiles, + /*N = */ O_c, + /*K = */ C_c, + /*batch_size_out = */ 8 * 8, + /*a_cols = */ C_c, + /*b_cols = */ O_c, + /*a_transposed = */ false, + /*b_transposed = */ false, + /*copies = */ empty_copies); + }   - compute_encoder.set_bytes(conv_params_updated, 2); + // Do output transform. + { + int bc = 32; + int wm = 2; + int wn = 2; + std::string kname; + kname.reserve(32); + concatenate( + kname, + "winograd_conv_2d_output_transform_", + type_to_name(out), + "_bo", + bc); + auto& compute_encoder = metal::get_command_encoder(s); + auto kernel = d.get_kernel(kname); + compute_encoder.set_compute_pipeline_state(kernel);   - MTL::Size group_dims = MTL::Size(32, wn, wm); - MTL::Size grid_dims = MTL::Size(N_tiles_w, N_tiles_h, N_tiles_n); + compute_encoder.set_input_array(out_wg, 0); + compute_encoder.set_output_array(out_tile, 1);   - compute_encoder.dispatch_threadgroups(grid_dims, group_dims); + compute_encoder.set_bytes(conv_params_updated, 2); + + MTL::Size group_dims = MTL::Size(32, wn, wm); + MTL::Size grid_dims = MTL::Size(N_tiles_w, N_tiles_h, tile_n); + + compute_encoder.dispatch_threadgroups(grid_dims, group_dims); + } } }   @@ -1120,7 +1443,11 @@ if (!conv_params.flip && is_stride_one && is_kdil_one && is_idil_one && conv_params.wS[0] == 3 && conv_params.wS[1] == 3 && conv_params.C % 32 == 0 && conv_params.O % 32 == 0 && inp_large && channels_large) { - return winograd_conv_2D_gpu(s, d, in, wt, out, conv_params, copies); + // Only use winograd conv when having enough memory. + if (int n_step = winograd_batch_step(d, in, conv_params); n_step > 0) { + return winograd_conv_2D_gpu( + s, d, in, wt, out, conv_params, copies, n_step); + } }   // Whether the specialized implicit gemm kernel can take the channels as-is.
diff --git ml-explore/mlx/mlx/backend/metal/event.cpp Layr-Labs/mlx/mlx/backend/metal/event.cpp index 77f48f08388db675be9b10d2ad81daa0f8440a74..38a387c9c088e6fbc5c75f7159c5fe89155c803c 100644 --- ml-explore/mlx/mlx/backend/metal/event.cpp +++ Layr-Labs/mlx/mlx/backend/metal/event.cpp @@ -26,26 +26,13 @@ mtl_event_.reset(); }   void EventImpl::wait(uint64_t value) { - check_error(); mtl_event_->waitUntilSignaledValue(value, -1); // never times out - check_error(); }   void EventImpl::signal(uint64_t value) { mtl_event_->setSignaledValue(value); }   -void EventImpl::set_error(std::shared_ptr<std::string> error) { - std::atomic_store(&error_, std::move(error)); -} - -void EventImpl::check_error() { - auto error = std::atomic_exchange(&error_, {}); - if (error) { - throw std::runtime_error(*error); - } -} - } // namespace metal   /////////////////////////////////////////////////////////////////////////////// @@ -57,36 +44,40 @@ event_ = std::make_shared<metal::EventImpl>(metal::device(stream.device)); }   void Event::wait() { - static_cast<metal::EventImpl*>(event_.get())->wait(value()); + check_error(); + cast<metal::EventImpl>().wait(value()); + check_error(); }   void Event::wait(Stream stream) { - auto impl = std::static_pointer_cast<metal::EventImpl>(event_); if (stream.device == Device::cpu) { - scheduler::enqueue(stream, [impl = std::move(impl), value = value()]() { - impl->wait(value); + scheduler::wait_event(stream, *this, [value = value()](Event& self) { + self.cast<metal::EventImpl>().wait(value); }); } else { auto& encoder = metal::get_command_encoder(stream); - encoder.wait_event(std::move(impl), value()); + encoder.wait_event(*this, value()); } }   void Event::signal(Stream stream) { - auto impl = std::static_pointer_cast<metal::EventImpl>(event_); if (stream.device == Device::cpu) { - scheduler::enqueue(stream, [impl = std::move(impl), value = value()]() { - impl->signal(value); + scheduler::signal_event(stream, *this, [value = value()](Event& self) { + self.cast<metal::EventImpl>().signal(value); }); } else { auto& encoder = metal::get_command_encoder(stream); - encoder.signal_event(std::move(impl), value()); + encoder.signal_event(*this, value()); } }   bool Event::is_signaled() const { - auto* mtl_event = static_cast<metal::EventImpl*>(event_.get())->mtl_event(); + auto* mtl_event = cast<metal::EventImpl>().mtl_event(); return mtl_event->signaledValue() >= value(); +} + +std::atomic<Error*>& Event::error() { + return cast<metal::EventImpl>().error(); }   } // namespace mlx::core
diff --git ml-explore/mlx/mlx/backend/metal/event.h Layr-Labs/mlx/mlx/backend/metal/event.h index c5c82a7cd37c9437c5ec3382b2081b2cc379b889..d1e43fa02f3c1320e5b9ec13ca894bcad8801b19 100644 --- ml-explore/mlx/mlx/backend/metal/event.h +++ Layr-Labs/mlx/mlx/backend/metal/event.h @@ -12,20 +12,18 @@ ~EventImpl();   void wait(uint64_t value); void signal(uint64_t value); - void set_error(std::shared_ptr<std::string> error); - void check_error();   - const auto& error() const { + auto& error() { return error_; }   - auto* mtl_event() { + auto* mtl_event() const { return mtl_event_.get(); }   private: - // TODO: Use std::atomic<std::shared_ptr> when it gets supported in Xcode. - std::shared_ptr<std::string> error_; + // All streams outlive events so pointers would be always valid. + std::atomic<Error*> error_;   NS::SharedPtr<MTL::SharedEvent> mtl_event_; };
diff --git ml-explore/mlx/mlx/backend/metal/fence.cpp Layr-Labs/mlx/mlx/backend/metal/fence.cpp index 6fdd57a5f621c66994455cbdddd00ea0ca7d5fd0..70dd0e33bd2aec957b084dffa198a23e4ae661bf 100644 --- ml-explore/mlx/mlx/backend/metal/fence.cpp +++ Layr-Labs/mlx/mlx/backend/metal/fence.cpp @@ -41,8 +41,7 @@ } };   Fence::Fence(Stream stream) { - auto dtor = [](void* ptr) { delete static_cast<FenceImpl*>(ptr); }; - fence_ = std::shared_ptr<void>(new FenceImpl(stream), dtor); + fence_ = std::make_shared<FenceImpl>(stream); }   void Fence::wait(Stream stream, const array& x) {
diff --git ml-explore/mlx/mlx/backend/metal/jit_kernels.cpp Layr-Labs/mlx/mlx/backend/metal/jit_kernels.cpp index c7900ecdf8f7b821c013be44fbc455d7a55ab067..8dfe30a15c366b2fdedd5ae7d7330502e6a33f0a 100644 --- ml-explore/mlx/mlx/backend/metal/jit_kernels.cpp +++ Layr-Labs/mlx/mlx/backend/metal/jit_kernels.cpp @@ -1330,7 +1330,8 @@ int bk, int bd, int wm, int wn, - const array& m) { + const array& m, + bool split_d) { const auto& lib_name = kernel_name; auto lib = d.get_library(lib_name, [&]() { std::string kernel_source; @@ -1340,7 +1341,7 @@ metal::utils(), metal::steel_attention_nax(), get_template_definition( lib_name, - "attention_nax", + split_d ? "attention_nax_dsplit" : "attention_nax", get_type_string(q.dtype()), bq, bk,
diff --git ml-explore/mlx/mlx/backend/metal/kernels.h Layr-Labs/mlx/mlx/backend/metal/kernels.h index 21b754514cce83dc8f95802bd2da385a1622ec9d..18a56ebf1ba835204c98912a84ecca63d429599f 100644 --- ml-explore/mlx/mlx/backend/metal/kernels.h +++ Layr-Labs/mlx/mlx/backend/metal/kernels.h @@ -426,7 +426,8 @@ int bk, int bd, int wm, int wn, - const array& m); + const array& m, + bool split_d);   // Create a GPU kernel template definition for JIT compilation template <typename... Args>
diff --git ml-explore/mlx/mlx/backend/metal/kernels/binary_ops.h Layr-Labs/mlx/mlx/backend/metal/kernels/binary_ops.h index 863d6369e267f0701673302a6691410e0d393319..37650c7d93adda42c3e5cd8587663d90facb0e14 100644 --- ml-explore/mlx/mlx/backend/metal/kernels/binary_ops.h +++ Layr-Labs/mlx/mlx/backend/metal/kernels/binary_ops.h @@ -16,20 +16,27 @@ };   struct FloorDivide { template <typename T> - T operator()(T x, T y) thread { + metal::enable_if_t<metal::is_integral_v<T> & !metal::is_signed_v<T>, T> + operator()(T x, T y) thread { return x / y; } - template <> - float operator()(float x, float y) thread { - return trunc(x / y); + template <typename T> + metal::enable_if_t<metal::is_integral_v<T> & metal::is_signed_v<T>, T> + operator()(T x, T y) thread { + auto q = x / y; + if (x % y != 0 && (x < 0) != (y < 0)) { + q -= 1; + } + return q; } - template <> - half operator()(half x, half y) thread { - return trunc(x / y); + template <typename T> + metal::enable_if_t<!metal::is_integral_v<T>, T> operator()(T x, T y) thread { + return floor(x / y); } template <> - bfloat16_t operator()(bfloat16_t x, bfloat16_t y) thread { - return trunc(x / y); + complex64_t operator()(complex64_t x, complex64_t y) thread { + // Complex is not supported, simply make compiler happy. + return x / y; } };
diff --git ml-explore/mlx/mlx/backend/metal/kernels/copy.h Layr-Labs/mlx/mlx/backend/metal/kernels/copy.h index cf22347ee51393616ac1ce464529e65f5068bc73..95ed69b760665b67a379274046a410bdb26547ad 100644 --- ml-explore/mlx/mlx/backend/metal/kernels/copy.h +++ Layr-Labs/mlx/mlx/backend/metal/kernels/copy.h @@ -9,11 +9,11 @@ uint index [[thread_position_in_grid]]) { index *= N; if (N > 1 && index + N > size) { for (int i = 0; index + i < size; ++i) { - dst[index + i] = static_cast<U>(src[0]); + dst[index + i] = cast_to<U>(src[0]); } } else { for (int i = 0; i < N; ++i) { - dst[index + i] = static_cast<U>(src[0]); + dst[index + i] = cast_to<U>(src[0]); } } } @@ -27,11 +27,11 @@ uint index [[thread_position_in_grid]]) { index *= N; if (N > 1 && index + N > size) { for (int i = 0; index + i < size; ++i) { - dst[index + i] = static_cast<U>(src[index + i]); + dst[index + i] = cast_to<U>(src[index + i]); } } else { for (int i = 0; i < N; ++i) { - dst[index + i] = static_cast<U>(src[index + i]); + dst[index + i] = cast_to<U>(src[index + i]); } } } @@ -46,11 +46,11 @@ uint2 grid_dim [[threads_per_grid]]) { int64_t offset = N * (index.x + grid_dim.x * int64_t(index.y)); if (N > 1 && offset + N > size) { for (int i = 0; offset + i < size; ++i) { - dst[offset + i] = static_cast<U>(src[0]); + dst[offset + i] = cast_to<U>(src[0]); } } else { for (int i = 0; i < N; ++i) { - dst[offset + i] = static_cast<U>(src[0]); + dst[offset + i] = cast_to<U>(src[0]); } } } @@ -65,11 +65,11 @@ uint2 grid_dim [[threads_per_grid]]) { int64_t offset = N * (index.x + grid_dim.x * int64_t(index.y)); if (N > 1 && offset + N > size) { for (int i = 0; offset + i < size; ++i) { - dst[offset + i] = static_cast<U>(src[offset + i]); + dst[offset + i] = cast_to<U>(src[offset + i]); } } else { for (int i = 0; i < N; ++i) { - dst[offset + i] = static_cast<U>(src[offset + i]); + dst[offset + i] = cast_to<U>(src[offset + i]); } } } @@ -81,7 +81,7 @@ device U* dst [[buffer(1)]], constant const int64_t& src_stride [[buffer(3)]], uint index [[thread_position_in_grid]]) { auto src_idx = elem_to_loc_1<IdxT>(index, src_stride); - dst[index] = static_cast<U>(src[src_idx]); + dst[index] = cast_to<U>(src[src_idx]); }   template <typename T, typename U, typename IdxT = int64_t> @@ -93,7 +93,7 @@ uint2 index [[thread_position_in_grid]], uint2 grid_dim [[threads_per_grid]]) { auto src_idx = elem_to_loc_2<IdxT>(index, src_strides); IdxT dst_idx = index.x + IdxT(grid_dim.x) * index.y; - dst[dst_idx] = static_cast<U>(src[src_idx]); + dst[dst_idx] = cast_to<U>(src[src_idx]); }   template <typename T, typename U, typename IdxT = int64_t> @@ -106,7 +106,7 @@ uint3 grid_dim [[threads_per_grid]]) { auto src_idx = elem_to_loc_3<IdxT>(index, src_strides); IdxT dst_idx = index.x + IdxT(grid_dim.x) * (index.y + IdxT(grid_dim.y) * index.z); - dst[dst_idx] = static_cast<U>(src[src_idx]); + dst[dst_idx] = cast_to<U>(src[src_idx]); }   template <typename T, typename U, int N = 1, typename IdxT = int64_t> @@ -123,14 +123,14 @@ {N * index.x, index.y, index.z}, src_shape, src_strides, ndim); if (N == 1) { IdxT dst_idx = index.x + grid_dim.x * (index.y + IdxT(grid_dim.y) * index.z); - dst[dst_idx] = static_cast<U>(src[src_idx]); + dst[dst_idx] = cast_to<U>(src[src_idx]); return; } auto xshape = src_shape[ndim - 1]; IdxT dst_idx = N * index.x + xshape * (index.y + IdxT(grid_dim.y) * index.z); auto src_xstride = src_strides[ndim - 1]; for (int i = 0; i < N && (int(N * index.x) + i) < xshape; ++i) { - dst[dst_idx + i] = static_cast<U>(src[src_idx]); + dst[dst_idx + i] = cast_to<U>(src[src_idx]); src_idx += src_xstride; } } @@ -144,7 +144,7 @@ constant const int64_t& dst_stride [[buffer(4)]], uint index [[thread_position_in_grid]]) { auto src_idx = elem_to_loc_1<IdxT>(index, src_stride); auto dst_idx = elem_to_loc_1<IdxT>(index, dst_stride); - dst[dst_idx] = static_cast<U>(src[src_idx]); + dst[dst_idx] = cast_to<U>(src[src_idx]); }   template <typename T, typename U, typename IdxT = int64_t> @@ -156,7 +156,7 @@ constant const int64_t* dst_strides [[buffer(4)]], uint2 index [[thread_position_in_grid]]) { auto src_idx = elem_to_loc_2<IdxT>(index, src_strides); auto dst_idx = elem_to_loc_2<IdxT>(index, dst_strides); - dst[dst_idx] = static_cast<U>(src[src_idx]); + dst[dst_idx] = cast_to<U>(src[src_idx]); }   template <typename T, typename U, typename IdxT = int64_t> @@ -168,7 +168,7 @@ constant const int64_t* dst_strides [[buffer(4)]], uint3 index [[thread_position_in_grid]]) { auto src_idx = elem_to_loc_3<IdxT>(index, src_strides); auto dst_idx = elem_to_loc_3<IdxT>(index, dst_strides); - dst[dst_idx] = static_cast<U>(src[src_idx]); + dst[dst_idx] = cast_to<U>(src[src_idx]); }   template <typename T, typename U, int N = 1, typename IdxT = int64_t> @@ -187,14 +187,14 @@ src_strides, dst_strides, ndim); if (N == 1) { - dst[idx.y] = static_cast<U>(src[idx.x]); + dst[idx.y] = cast_to<U>(src[idx.x]); return; } IdxT src_xstride = src_strides[ndim - 1]; IdxT dst_xstride = dst_strides[ndim - 1]; auto xshape = src_shape[ndim - 1]; for (int i = 0; i < N && (int(N * index.x) + i) < xshape; ++i) { - dst[idx.y] = static_cast<U>(src[idx.x]); + dst[idx.y] = cast_to<U>(src[idx.x]); idx.x += src_xstride; idx.y += dst_xstride; } @@ -211,7 +211,7 @@ constant const int64_t& dst_offset [[buffer(7)]], uint index [[thread_position_in_grid]]) { auto src_idx = elem_to_loc_1<IdxT>(index, src_stride); auto dst_idx = elem_to_loc_1<IdxT>(index, dst_stride); - dst[dst_idx + dst_offset] = src[src_idx + src_offset]; + dst[dst_idx + dst_offset] = cast_to<U>(src[src_idx + src_offset]); }   template <typename T, typename U, typename IdxT = int64_t> @@ -225,7 +225,7 @@ constant const int64_t& dst_offset [[buffer(7)]], uint2 index [[thread_position_in_grid]]) { auto src_idx = elem_to_loc_2<IdxT>(index, src_strides); auto dst_idx = elem_to_loc_2<IdxT>(index, dst_strides); - dst[dst_idx + dst_offset] = src[src_idx + src_offset]; + dst[dst_idx + dst_offset] = cast_to<U>(src[src_idx + src_offset]); }   template <typename T, typename U, typename IdxT = int64_t> @@ -239,7 +239,7 @@ constant const int64_t& dst_offset [[buffer(7)]], uint3 index [[thread_position_in_grid]]) { auto src_idx = elem_to_loc_3<IdxT>(index, src_strides); auto dst_idx = elem_to_loc_3<IdxT>(index, dst_strides); - dst[dst_idx + dst_offset] = src[src_idx + src_offset]; + dst[dst_idx + dst_offset] = cast_to<U>(src[src_idx + src_offset]); }   template <typename T, typename U, int N = 1, typename IdxT = int64_t> @@ -262,14 +262,14 @@ src_strides, dst_strides, ndim); if (N == 1) { - dst[idx.y] = src[idx.x]; + dst[idx.y] = cast_to<U>(src[idx.x]); return; } IdxT src_xstride = src_strides[ndim - 1]; IdxT dst_xstride = dst_strides[ndim - 1]; auto xshape = src_shape[ndim - 1]; for (int i = 0; i < N && (int(N * index.x) + i) < xshape; ++i) { - dst[idx.y] = src[idx.x]; + dst[idx.y] = cast_to<U>(src[idx.x]); idx.x += src_xstride; idx.y += dst_xstride; }
diff --git ml-explore/mlx/mlx/backend/metal/kernels/fp8.h Layr-Labs/mlx/mlx/backend/metal/kernels/fp8.h index 796dd21639b37716867f89992002ed6c7370fe77..42c5ec128a6ca80dade9a03f508632d8b270f5bb 100644 --- ml-explore/mlx/mlx/backend/metal/kernels/fp8.h +++ Layr-Labs/mlx/mlx/backend/metal/kernels/fp8.h @@ -78,3 +78,14 @@ }   uint8_t bits; }; + +// Smallest E8M0 >= x. Scales are amax/max_element, so rounding one down +// leaves the block's largest elements outside the element range, where they +// saturate. Matches the CUDA backend, which rounds up via cutlass ue8m0. +inline float mx_scale_round_up(float x) { + fp8_e8m0 s(x); + if (s.bits < 0xFE && float(s) < x) { + s.bits += 1; + } + return float(s); +}
diff --git ml-explore/mlx/mlx/backend/metal/kernels/fp_quantized_nax.h Layr-Labs/mlx/mlx/backend/metal/kernels/fp_quantized_nax.h index cf64ff7f46d4d43eef67efdb74abf5ae551be7db..946bce7868cbeba393c3a2181a9ae23b328a0a10 100644 --- ml-explore/mlx/mlx/backend/metal/kernels/fp_quantized_nax.h +++ Layr-Labs/mlx/mlx/backend/metal/kernels/fp_quantized_nax.h @@ -61,11 +61,12 @@ };   template <typename U, int bits> inline void dequantize(uint8_t w, U scale, threadgroup U* w_local) { + const float s = float(scale); if constexpr (bits == 4) { - w_local[0] = scale * Dequantize<4, U>{}(w); - w_local[1] = scale * Dequantize<4, U>{}(w >> 4); + w_local[0] = static_cast<U>(s * Dequantize<4, float>{}(w)); + w_local[1] = static_cast<U>(s * Dequantize<4, float>{}(w >> 4)); } else { - w_local[0] = scale * Dequantize<8, U>{}(w); + w_local[0] = static_cast<U>(s * Dequantize<8, float>{}(w)); } }   @@ -896,6 +897,10 @@ } threadgroup_barrier(mem_flags::mem_none);   // Prepare threadgroup mma operation + const short m_lo_lim = min(int(sgp_sm), max(0, offset - tm)); + const short m_hi_lim = min(int(sgp_sm), max(0, offset_next - tm)); + const bool sg_active = m_hi_lim > m_lo_lim; + NAXTile<AccumType, TM, TN> Dtile; Dtile.clear();   @@ -925,33 +930,35 @@ threadgroup_barrier(mem_flags::mem_threadgroup);   STEEL_PRAGMA_NO_UNROLL for (int kk1 = 0; kk1 < BK; kk1 += SK) { - NAXTile<T, TM, TK> Atile; - NAXTile<Wtype, BR, BC> Btile; + if (sg_active) { + NAXTile<T, TM, TK> Atile; + NAXTile<Wtype, BR, BC> Btile;   - volatile int compiler_barrier; + volatile int compiler_barrier;   - if constexpr (kAlignedM.value) { - Atile.load(xn + kk1, K); - } else { - Atile.load_safe(xn + kk1, K, short2(SK, sgp_sm)); - } + if constexpr (kAlignedM.value) { + Atile.load(xn + kk1, K); + } else { + Atile.load_safe(xn + kk1, K, short2(SK, sgp_sm)); + }   - if constexpr (transpose) { - Btile.template load<Wtype, BK_padded, 1>( - Ws + tn * BK_padded + kk1); - } else { - Btile.template load<Wtype, BN_padded, 1>( - Ws + tn + kk1 * BN_padded); - } + if constexpr (transpose) { + Btile.template load<Wtype, BK_padded, 1>( + Ws + tn * BK_padded + kk1); + } else { + Btile.template load<Wtype, BN_padded, 1>( + Ws + tn + kk1 * BN_padded); + }   - tile_matmad_nax( - Dtile, - Atile, - metal::bool_constant<false>{}, - Btile, - metal::bool_constant<transpose>{}); + tile_matmad_nax( + Dtile, + Atile, + metal::bool_constant<false>{}, + Btile, + metal::bool_constant<transpose>{});   - (void)compiler_barrier; + (void)compiler_barrier; + } }   xn += BK; @@ -965,37 +972,36 @@ threadgroup_barrier(mem_flags::mem_threadgroup);   STEEL_PRAGMA_NO_UNROLL for (int kk1 = 0; kk1 < BK; kk1 += SK) { - NAXTile<T, TM, TK> Atile; - NAXTile<Wtype, BR, BC> Btile; + if (sg_active) { + NAXTile<T, TM, TK> Atile; + NAXTile<Wtype, BR, BC> Btile;   - volatile int compiler_barrier; + volatile int compiler_barrier;   - const short psk = min(int(SK), max(0, (BK - kk1))); - Atile.load_safe(xn + kk1, K, short2(psk, sgp_sm)); + const short psk = min(int(SK), max(0, (BK - kk1))); + Atile.load_safe(xn + kk1, K, short2(psk, sgp_sm));   - if constexpr (transpose) { - Btile.template load<Wtype, BK_padded, 1>( - Ws + tn * BK_padded + kk1); - } else { - Btile.template load<Wtype, BN_padded, 1>( - Ws + tn + kk1 * BN_padded); - } + if constexpr (transpose) { + Btile.template load<Wtype, BK_padded, 1>( + Ws + tn * BK_padded + kk1); + } else { + Btile.template load<Wtype, BN_padded, 1>( + Ws + tn + kk1 * BN_padded); + }   - tile_matmad_nax( - Dtile, - Atile, - metal::bool_constant<false>{}, - Btile, - metal::bool_constant<transpose>{}); + tile_matmad_nax( + Dtile, + Atile, + metal::bool_constant<false>{}, + Btile, + metal::bool_constant<transpose>{});   - (void)compiler_barrier; + (void)compiler_barrier; + } } }   threadgroup_barrier(mem_flags::mem_threadgroup); - - const short m_lo_lim = min(int(sgp_sm), max(0, offset - tm)); - const short m_hi_lim = min(int(sgp_sm), max(0, offset_next - tm));   // Store results to device memory if constexpr (kAlignedN.value) {
diff --git ml-explore/mlx/mlx/backend/metal/kernels/fp_quantized_nax.metal Layr-Labs/mlx/mlx/backend/metal/kernels/fp_quantized_nax.metal index c736f1809ee85ef34c928b0a4006d14560343a7b..771b2a963a751252bd21eeb6a8c3bd9644ecdcd7 100644 --- ml-explore/mlx/mlx/backend/metal/kernels/fp_quantized_nax.metal +++ Layr-Labs/mlx/mlx/backend/metal/kernels/fp_quantized_nax.metal @@ -24,7 +24,7 @@ fp_ ## name, \ type, \ group_size, \ bits, \ - aligned) + aligned, bm, bk, bn, wm, wn)   #define instantiate_quantized_aligned_batched(mode, name, type, bm, bn, bk, wm, wn, aligned, batched, group_size, bits) \ instantiate_kernel( \ @@ -34,7 +34,7 @@ type, \ group_size, \ bits, \ aligned, \ - batched) + batched, bm, bk, bn, wm, wn)   #define instantiate_gather_qmm_rhs(func, name, type, bm, bn, bk, wm, wn, transpose, mode, group_size, bits) \ instantiate_kernel( \ @@ -57,7 +57,11 @@ instantiate_quantized_aligned(mode, gather_qmm_t_nax, type, 64, 64, 64, 2, 2, false, group_size, bits) \ instantiate_quantized_aligned_batched(mode, qmm_t_nax, type, 64, 64, 64, 2, 2, true, 1, group_size, bits) \ instantiate_quantized_aligned_batched(mode, qmm_t_nax, type, 64, 64, 64, 2, 2, true, 0, group_size, bits) \ instantiate_quantized_aligned_batched(mode, qmm_t_nax, type, 64, 64, 64, 2, 2, false, 1, group_size, bits) \ - instantiate_quantized_aligned_batched(mode, qmm_t_nax, type, 64, 64, 64, 2, 2, false, 0, group_size, bits) + instantiate_quantized_aligned_batched(mode, qmm_t_nax, type, 64, 64, 64, 2, 2, false, 0, group_size, bits) \ + instantiate_quantized_aligned_batched(mode, qmm_t_nax, type, 32, 64, 64, 2, 2, true, 1, group_size, bits) \ + instantiate_quantized_aligned_batched(mode, qmm_t_nax, type, 32, 64, 64, 2, 2, true, 0, group_size, bits) \ + instantiate_quantized_aligned_batched(mode, qmm_t_nax, type, 32, 64, 64, 2, 2, false, 1, group_size, bits) \ + instantiate_quantized_aligned_batched(mode, qmm_t_nax, type, 32, 64, 64, 2, 2, false, 0, group_size, bits)   #define instantiate_quantized_all_rhs(type, mode, group_size, bits) \
diff --git ml-explore/mlx/mlx/backend/metal/kernels/quantized_nax.h Layr-Labs/mlx/mlx/backend/metal/kernels/quantized_nax.h index 31e51a5b7e3aa55364be525e32829606937bddb4..ed32eb59a7f24b9b603b7cdedba9be4465109a02 100644 --- ml-explore/mlx/mlx/backend/metal/kernels/quantized_nax.h +++ Layr-Labs/mlx/mlx/backend/metal/kernels/quantized_nax.h @@ -491,17 +491,16 @@ bits == 2 || bits == 3 || bits == 4 || bits == 5 || bits == 6 || bits == 8, "Template undefined for bits not in {2, 3, 4, 5, 6, 8}");   + const float s = float(scale); + const float b = float(bias); + if (bits == 2) { - U s[4] = { - scale, - scale / static_cast<U>(4.0f), - scale / static_cast<U>(16.0f), - scale / static_cast<U>(64.0f)}; + float sc[4] = {s, s / 4.0f, s / 16.0f, s / 64.0f}; for (int i = 0; i < (N / 4); i++) { - w_local[4 * i] = s[0] * (w[i] & 0x03) + bias; - w_local[4 * i + 1] = s[1] * (w[i] & 0x0c) + bias; - w_local[4 * i + 2] = s[2] * (w[i] & 0x30) + bias; - w_local[4 * i + 3] = s[3] * (w[i] & 0xc0) + bias; + w_local[4 * i] = static_cast<U>(sc[0] * (w[i] & 0x03) + b); + w_local[4 * i + 1] = static_cast<U>(sc[1] * (w[i] & 0x0c) + b); + w_local[4 * i + 2] = static_cast<U>(sc[2] * (w[i] & 0x30) + b); + w_local[4 * i + 3] = static_cast<U>(sc[3] * (w[i] & 0xc0) + b); } }   @@ -510,22 +509,24 @@ for (int i = 0; i < (N / 8); i++) { w_local += 8 * i; w += 3 * i;   - w_local[0] = (w[0] & 0x7) * scale + bias; - w_local[1] = ((w[0] & 0x38) >> 3) * scale + bias; - w_local[2] = (((w[0] & 0xc0) >> 6) + ((w[1] & 0x1) << 2)) * scale + bias; - w_local[3] = ((w[1] & 0xe) >> 1) * scale + bias; - w_local[4] = ((w[1] & 0x70) >> 4) * scale + bias; - w_local[5] = (((w[1] & 0x80) >> 7) + ((w[2] & 0x3) << 1)) * scale + bias; - w_local[6] = ((w[2] & 0x1c) >> 2) * scale + bias; - w_local[7] = ((w[2] & 0xe0) >> 5) * scale + bias; + w_local[0] = static_cast<U>((w[0] & 0x7) * s + b); + w_local[1] = static_cast<U>(((w[0] & 0x38) >> 3) * s + b); + w_local[2] = + static_cast<U>((((w[0] & 0xc0) >> 6) + ((w[1] & 0x1) << 2)) * s + b); + w_local[3] = static_cast<U>(((w[1] & 0xe) >> 1) * s + b); + w_local[4] = static_cast<U>(((w[1] & 0x70) >> 4) * s + b); + w_local[5] = + static_cast<U>((((w[1] & 0x80) >> 7) + ((w[2] & 0x3) << 1)) * s + b); + w_local[6] = static_cast<U>(((w[2] & 0x1c) >> 2) * s + b); + w_local[7] = static_cast<U>(((w[2] & 0xe0) >> 5) * s + b); } }   else if (bits == 4) { - U s[2] = {scale, scale / static_cast<U>(16.0f)}; + float sc[2] = {s, s / 16.0f}; for (int i = 0; i < (N / 2); i++) { - w_local[2 * i] = s[0] * (w[i] & 0x0f) + bias; - w_local[2 * i + 1] = s[1] * (w[i] & 0xf0) + bias; + w_local[2 * i] = static_cast<U>(sc[0] * (w[i] & 0x0f) + b); + w_local[2 * i + 1] = static_cast<U>(sc[1] * (w[i] & 0xf0) + b); } }   @@ -534,14 +535,18 @@ for (int i = 0; i < (N / 8); i++) { w_local += 8 * i; w += 5 * i;   - w_local[0] = (w[0] & 0x1f) * scale + bias; - w_local[1] = (((w[0] & 0xe0) >> 5) + ((w[1] & 0x3) << 3)) * scale + bias; - w_local[2] = ((w[1] & 0x7c) >> 2) * scale + bias; - w_local[3] = (((w[1] & 0x80) >> 7) + ((w[2] & 0xf) << 1)) * scale + bias; - w_local[4] = (((w[2] & 0xf0) >> 4) + ((w[3] & 0x1) << 4)) * scale + bias; - w_local[5] = ((w[3] & 0x3e) >> 1) * scale + bias; - w_local[6] = (((w[3] & 0xc0) >> 6) + ((w[4] & 0x7) << 2)) * scale + bias; - w_local[7] = ((w[4] & 0xf8) >> 3) * scale + bias; + w_local[0] = static_cast<U>((w[0] & 0x1f) * s + b); + w_local[1] = + static_cast<U>((((w[0] & 0xe0) >> 5) + ((w[1] & 0x3) << 3)) * s + b); + w_local[2] = static_cast<U>(((w[1] & 0x7c) >> 2) * s + b); + w_local[3] = + static_cast<U>((((w[1] & 0x80) >> 7) + ((w[2] & 0xf) << 1)) * s + b); + w_local[4] = + static_cast<U>((((w[2] & 0xf0) >> 4) + ((w[3] & 0x1) << 4)) * s + b); + w_local[5] = static_cast<U>(((w[3] & 0x3e) >> 1) * s + b); + w_local[6] = + static_cast<U>((((w[3] & 0xc0) >> 6) + ((w[4] & 0x7) << 2)) * s + b); + w_local[7] = static_cast<U>(((w[4] & 0xf8) >> 3) * s + b); } }   @@ -549,16 +554,18 @@ else if (bits == 6) { for (int i = 0; i < (N / 4); i++) { w_local += 4 * i; w += 3 * i; - w_local[0] = (w[0] & 0x3f) * scale + bias; - w_local[1] = (((w[0] >> 6) & 0x03) + ((w[1] & 0x0f) << 2)) * scale + bias; - w_local[2] = (((w[1] >> 4) & 0x0f) + ((w[2] & 0x03) << 4)) * scale + bias; - w_local[3] = ((w[2] >> 2) & 0x3f) * scale + bias; + w_local[0] = static_cast<U>((w[0] & 0x3f) * s + b); + w_local[1] = + static_cast<U>((((w[0] >> 6) & 0x03) + ((w[1] & 0x0f) << 2)) * s + b); + w_local[2] = + static_cast<U>((((w[1] >> 4) & 0x0f) + ((w[2] & 0x03) << 4)) * s + b); + w_local[3] = static_cast<U>(((w[2] >> 2) & 0x3f) * s + b); } }   else if (bits == 8) { for (int i = 0; i < N; i++) { - w_local[i] = scale * w[i] + bias; + w_local[i] = static_cast<U>(s * w[i] + b); } } } @@ -1569,6 +1576,10 @@ } } threadgroup_barrier(mem_flags::mem_none);   + const short m_lo_lim = min(int(sgp_sm), max(0, offset - tm)); + const short m_hi_lim = min(int(sgp_sm), max(0, offset_next - tm)); + const bool sg_active = m_hi_lim > m_lo_lim; + NAXTile<AccumType, TM, TN> Dtile; Dtile.clear();   @@ -1599,31 +1610,33 @@ threadgroup_barrier(mem_flags::mem_threadgroup);   STEEL_PRAGMA_NO_UNROLL for (int kk1 = 0; kk1 < BK; kk1 += SK) { - NAXTile<T, TM, TK> Atile; - NAXTile<T, BR, BC> Btile; + if (sg_active) { + NAXTile<T, TM, TK> Atile; + NAXTile<T, BR, BC> Btile;   - volatile int compiler_barrier; + volatile int compiler_barrier;   - if constexpr (kAlignedM.value) { - Atile.load(xn + kk1, K); - } else { - Atile.load_safe(xn + kk1, K, short2(SK, sgp_sm)); - } + if constexpr (kAlignedM.value) { + Atile.load(xn + kk1, K); + } else { + Atile.load_safe(xn + kk1, K, short2(SK, sgp_sm)); + }   - if constexpr (transpose) { - Btile.template load<T, BK_padded, 1>(Ws + tn * BK_padded + kk1); - } else { - Btile.template load<T, BN_padded, 1>(Ws + tn + kk1 * BN_padded); - } + if constexpr (transpose) { + Btile.template load<T, BK_padded, 1>(Ws + tn * BK_padded + kk1); + } else { + Btile.template load<T, BN_padded, 1>(Ws + tn + kk1 * BN_padded); + }   - tile_matmad_nax( - Dtile, - Atile, - metal::bool_constant<false>{}, - Btile, - metal::bool_constant<transpose>{}); + tile_matmad_nax( + Dtile, + Atile, + metal::bool_constant<false>{}, + Btile, + metal::bool_constant<transpose>{});   - (void)compiler_barrier; + (void)compiler_barrier; + } }   xn += BK; @@ -1637,35 +1650,34 @@ threadgroup_barrier(mem_flags::mem_threadgroup);   STEEL_PRAGMA_NO_UNROLL for (int kk1 = 0; kk1 < BK; kk1 += SK) { - NAXTile<T, TM, TK> Atile; - NAXTile<T, BR, BC> Btile; + if (sg_active) { + NAXTile<T, TM, TK> Atile; + NAXTile<T, BR, BC> Btile;   - volatile int compiler_barrier; + volatile int compiler_barrier;   - const short psk = min(int(SK), max(0, (BK - kk1))); - Atile.load_safe(xn + kk1, K, short2(psk, sgp_sm)); + const short psk = min(int(SK), max(0, (BK - kk1))); + Atile.load_safe(xn + kk1, K, short2(psk, sgp_sm));   - if constexpr (transpose) { - Btile.template load<T, BK_padded, 1>(Ws + tn * BK_padded + kk1); - } else { - Btile.template load<T, BN_padded, 1>(Ws + tn + kk1 * BN_padded); - } + if constexpr (transpose) { + Btile.template load<T, BK_padded, 1>(Ws + tn * BK_padded + kk1); + } else { + Btile.template load<T, BN_padded, 1>(Ws + tn + kk1 * BN_padded); + }   - tile_matmad_nax( - Dtile, - Atile, - metal::bool_constant<false>{}, - Btile, - metal::bool_constant<transpose>{}); + tile_matmad_nax( + Dtile, + Atile, + metal::bool_constant<false>{}, + Btile, + metal::bool_constant<transpose>{});   - (void)compiler_barrier; + (void)compiler_barrier; + } } }   threadgroup_barrier(mem_flags::mem_threadgroup); - - const short m_lo_lim = min(int(sgp_sm), max(0, offset - tm)); - const short m_hi_lim = min(int(sgp_sm), max(0, offset_next - tm));   // Store results to device memory if constexpr (kAlignedN.value) {
diff --git ml-explore/mlx/mlx/backend/metal/kernels/quantized_nax.metal Layr-Labs/mlx/mlx/backend/metal/kernels/quantized_nax.metal index 27302ecb5fe505db20c9cc3eea5f8df13e6f62f7..9557fd838d73011bf244dacf389e25f44ecb6486 100644 --- ml-explore/mlx/mlx/backend/metal/kernels/quantized_nax.metal +++ Layr-Labs/mlx/mlx/backend/metal/kernels/quantized_nax.metal @@ -74,7 +74,11 @@ instantiate_quantized_aligned(affine_gather_qmm_t_nax, type, group_size, bits, 64, 64, 64, 2, 2, false) \ instantiate_quantized_aligned_batched(affine_qmm_t_nax, type, group_size, bits, 64, 64, 64, 2, 2, true, 1) \ instantiate_quantized_aligned_batched(affine_qmm_t_nax, type, group_size, bits, 64, 64, 64, 2, 2, true, 0) \ instantiate_quantized_aligned_batched(affine_qmm_t_nax, type, group_size, bits, 64, 64, 64, 2, 2, false, 1) \ - instantiate_quantized_aligned_batched(affine_qmm_t_nax, type, group_size, bits, 64, 64, 64, 2, 2, false, 0) + instantiate_quantized_aligned_batched(affine_qmm_t_nax, type, group_size, bits, 64, 64, 64, 2, 2, false, 0) \ + instantiate_quantized_aligned_batched(affine_qmm_t_nax, type, group_size, bits, 32, 64, 64, 2, 2, true, 1) \ + instantiate_quantized_aligned_batched(affine_qmm_t_nax, type, group_size, bits, 32, 64, 64, 2, 2, true, 0) \ + instantiate_quantized_aligned_batched(affine_qmm_t_nax, type, group_size, bits, 32, 64, 64, 2, 2, false, 1) \ + instantiate_quantized_aligned_batched(affine_qmm_t_nax, type, group_size, bits, 32, 64, 64, 2, 2, false, 0)   #define instantiate_quantized_all_rhs(type, group_size, bits) \ instantiate_gather_qmm_rhs(affine_gather_qmm_rhs_nax, affine_gather_qmm_rhs_nax_nt, type, group_size, bits, 64, 64, 64, 2, 2, true) \
diff --git ml-explore/mlx/mlx/backend/metal/kernels/reduction/reduce_all.h Layr-Labs/mlx/mlx/backend/metal/kernels/reduction/reduce_all.h index e0d08392c0b0e7efd8c85f7191efafe3f9333ed7..47ad63fbbd2d3682ff5cd43758a34c44f63da873 100644 --- ml-explore/mlx/mlx/backend/metal/kernels/reduction/reduce_all.h +++ Layr-Labs/mlx/mlx/backend/metal/kernels/reduction/reduce_all.h @@ -37,13 +37,13 @@ }   for (IdxT b = 0; b < blocks; b++) { for (int i = 0; i < N_READS; i++) { - total = op(static_cast<U>(in[i]), total); + total = op(cast_to<U>(in[i]), total); } in += lsize.x * N_READS; } if (extra > 0) { for (int i = 0; i < extra; i++) { - total = op(static_cast<U>(in[i]), total); + total = op(cast_to<U>(in[i]), total); } }
diff --git ml-explore/mlx/mlx/backend/metal/kernels/reduction/reduce_col.h Layr-Labs/mlx/mlx/backend/metal/kernels/reduction/reduce_col.h index c109faf0bce51f9f58a8218832f4726770877e3e..b1546adb55d62f269d38c0a1623f53599e194dcf 100644 --- ml-explore/mlx/mlx/backend/metal/kernels/reduction/reduce_col.h +++ Layr-Labs/mlx/mlx/backend/metal/kernels/reduction/reduce_col.h @@ -43,13 +43,13 @@ for (IdxT r = lid.y; r < total_rows; r += lsize.y) { row = in + loop.location(); if (safe) { for (int i = 0; i < n_reads; i++) { - totals[i] = op(static_cast<U>(row[i]), totals[i]); + totals[i] = op(cast_to<U>(row[i]), totals[i]); } } else { U vals[n_reads]; for (int i = 0; i < n_reads; i++) { vals[i] = - (column + i < reduction_stride) ? static_cast<U>(row[i]) : op.init; + (column + i < reduction_stride) ? cast_to<U>(row[i]) : op.init; } for (int i = 0; i < n_reads; i++) { totals[i] = op(vals[i], totals[i]); @@ -125,7 +125,7 @@ loop.next(gid.z * lsize.y + lid.y, reduce_shape, reduce_strides); for (IdxT r = gid.z * lsize.y + lid.y; r < total_rows; r += lsize.y * gsize.z) { row = in + loop.location(); - total = op(static_cast<U>(*row), total); + total = op(cast_to<U>(*row), total); loop.next(lsize.y * gsize.z, reduce_shape, reduce_strides); }   @@ -207,13 +207,13 @@ row = in + loop.location();   if (safe) { for (int i = 0; i < n_reads; i++) { - totals[i] = op(static_cast<U>(row[i]), totals[i]); + totals[i] = op(cast_to<U>(row[i]), totals[i]); } } else { U vals[n_reads]; for (int i = 0; i < n_reads; i++) { vals[i] = - (column + i < reduction_stride) ? static_cast<U>(row[i]) : op.init; + (column + i < reduction_stride) ? cast_to<U>(row[i]) : op.init; } for (int i = 0; i < n_reads; i++) { totals[i] = op(vals[i], totals[i]); @@ -352,13 +352,13 @@ row = in + loop.location();   if (safe) { for (int i = 0; i < n_reads; i++) { - totals[i] = op(static_cast<U>(row[i]), totals[i]); + totals[i] = op(cast_to<U>(row[i]), totals[i]); } } else { U vals[n_reads]; for (int i = 0; i < n_reads; i++) { vals[i] = - (column + i < reduction_stride) ? static_cast<U>(row[i]) : op.init; + (column + i < reduction_stride) ? cast_to<U>(row[i]) : op.init; } for (int i = 0; i < n_reads; i++) { totals[i] = op(vals[i], totals[i]);
diff --git ml-explore/mlx/mlx/backend/metal/kernels/reduction/reduce_row.h Layr-Labs/mlx/mlx/backend/metal/kernels/reduction/reduce_row.h index 936d75bb52629f5ed54839a75bc81795f06e52ec..09d1d8834855800677d721ae7700b0d87fda971a 100644 --- ml-explore/mlx/mlx/backend/metal/kernels/reduction/reduce_row.h +++ Layr-Labs/mlx/mlx/backend/metal/kernels/reduction/reduce_row.h @@ -34,7 +34,7 @@ // Loop over the reduction size within thread group for (int i = 0; i < blocks; i++) { for (int j = 0; j < N_WRITES; j++) { for (int i = 0; i < N_READS; i++) { - totals[j] = op(static_cast<U>(inputs[j][i]), totals[j]); + totals[j] = op(cast_to<U>(inputs[j][i]), totals[j]); }   inputs[j] += lsize_x * N_READS; @@ -46,13 +46,13 @@ int index = lid_x * N_READS; if (index + N_READS <= extra) { for (int j = 0; j < N_WRITES; j++) { for (int i = 0; i < N_READS; i++) { - totals[j] = op(static_cast<U>(inputs[j][i]), totals[j]); + totals[j] = op(cast_to<U>(inputs[j][i]), totals[j]); } } } else { for (int j = 0; j < N_WRITES; j++) { for (int i = 0; index + i < extra; i++) { - totals[j] = op(static_cast<U>(inputs[j][i]), totals[j]); + totals[j] = op(cast_to<U>(inputs[j][i]), totals[j]); } } } @@ -337,7 +337,8 @@ IdxT out_idx = gid.y + gsize.y * IdxT(gid.z);   // lid.x * N_READS breaks the per_thread_row_reduce interface a bit. Maybe it // needs a small refactor. - in += elem_to_loc<IdxT>(out_idx, shape, strides, ndim) + lid.x * N_READS; + in += + elem_to_loc<IdxT>(out_idx, shape, strides, ndim) + IdxT(lid.x) * N_READS;   LoopedElemToLoc<NDIMS, IdxT, (NDIMS > 2)> loop(reduce_ndim); const device T* row;
diff --git ml-explore/mlx/mlx/backend/metal/kernels/rms_norm.metal Layr-Labs/mlx/mlx/backend/metal/kernels/rms_norm.metal index a50d4a25c642d7c7966d1223d808220ac05502d8..eb9b5c0af1cde571570ea749c95969d963262882 100644 --- ml-explore/mlx/mlx/backend/metal/kernels/rms_norm.metal +++ Layr-Labs/mlx/mlx/backend/metal/kernels/rms_norm.metal @@ -166,21 +166,32 @@ device T* gw, constant float& eps, constant uint& axis_size, constant uint& w_stride, + constant uint& n_rows, + constant uint& rows_per_group, uint gid [[threadgroup_position_in_grid]], uint lid [[thread_position_in_threadgroup]], uint simd_lane_id [[thread_index_in_simdgroup]], uint simd_group_id [[simdgroup_index_in_threadgroup]]) { - // Advance the input pointers - x += gid * size_t(axis_size) + lid * N_READS; - g += gid * size_t(axis_size) + lid * N_READS; w += w_stride * lid * N_READS; + float thread_w[N_READS]; + if (lid * N_READS + N_READS <= axis_size) { + for (int i = 0; i < N_READS; i++) { + thread_w[i] = w[w_stride * i]; + } + } else { + for (int i = 0; i < N_READS; i++) { + thread_w[i] = + (lid * N_READS + i < axis_size) ? (float)w[w_stride * i] : 0; + } + }   // Allocate registers for the computation and accumulators float thread_x[N_READS]; - float thread_w[N_READS]; float thread_g[N_READS]; - float sumx2 = 0; - float sumgwx = 0; + float gw_acc[N_READS]; + for (int i = 0; i < N_READS; i++) { + gw_acc[i] = 0; + }   // Allocate shared memory to implement the reduction constexpr int SIMD_SIZE = 32; @@ -189,75 +200,99 @@ threadgroup float local_sumgwx[SIMD_SIZE]; threadgroup float local_normalizer[1]; threadgroup float local_meangwx[1];   - // Read and accumulate locally - if (lid * N_READS + N_READS <= axis_size) { - for (int i = 0; i < N_READS; i++) { - thread_x[i] = x[i]; - thread_w[i] = w[w_stride * i]; - thread_g[i] = g[i]; + uint row_end = gid * rows_per_group + rows_per_group; + if (row_end > n_rows) { + row_end = n_rows; + } + for (uint row = gid * rows_per_group; row < row_end; ++row) { + const device T* x_in = x + size_t(row) * axis_size + lid * N_READS; + const device T* g_in = g + size_t(row) * axis_size + lid * N_READS;   - sumx2 += thread_x[i] * thread_x[i]; - sumgwx += thread_x[i] * thread_w[i] * thread_g[i]; - } - } else { - for (int i = 0; i < N_READS; i++) { - if ((lid * N_READS + i) < axis_size) { - thread_x[i] = x[i]; - thread_w[i] = w[w_stride * i]; - thread_g[i] = g[i]; + float sumx2 = 0; + float sumgwx = 0; + + // Read and accumulate locally + if (lid * N_READS + N_READS <= axis_size) { + for (int i = 0; i < N_READS; i++) { + thread_x[i] = x_in[i]; + thread_g[i] = g_in[i];   sumx2 += thread_x[i] * thread_x[i]; sumgwx += thread_x[i] * thread_w[i] * thread_g[i]; } + } else { + for (int i = 0; i < N_READS; i++) { + if ((lid * N_READS + i) < axis_size) { + thread_x[i] = x_in[i]; + thread_g[i] = g_in[i]; + + sumx2 += thread_x[i] * thread_x[i]; + sumgwx += thread_x[i] * thread_w[i] * thread_g[i]; + } + } } - }   - // Accumulate across threads - sumx2 = simd_sum(sumx2); - sumgwx = simd_sum(sumgwx); - if (simd_group_id == 0) { - local_sumx2[simd_lane_id] = 0; - local_sumgwx[simd_lane_id] = 0; - } - threadgroup_barrier(mem_flags::mem_threadgroup); - if (simd_lane_id == 0) { - local_sumx2[simd_group_id] = sumx2; - local_sumgwx[simd_group_id] = sumgwx; - } - threadgroup_barrier(mem_flags::mem_threadgroup); - if (simd_group_id == 0) { - sumx2 = simd_sum(local_sumx2[simd_lane_id]); - sumgwx = simd_sum(local_sumgwx[simd_lane_id]); + // Accumulate across threads + sumx2 = simd_sum(sumx2); + sumgwx = simd_sum(sumgwx); + if (simd_group_id == 0) { + local_sumx2[simd_lane_id] = 0; + local_sumgwx[simd_lane_id] = 0; + } + threadgroup_barrier(mem_flags::mem_threadgroup); if (simd_lane_id == 0) { - local_meangwx[0] = sumgwx / axis_size; - local_normalizer[0] = metal::precise::rsqrt(sumx2 / axis_size + eps); + local_sumx2[simd_group_id] = sumx2; + local_sumgwx[simd_group_id] = sumgwx; } - } - threadgroup_barrier(mem_flags::mem_threadgroup); - float meangwx = local_meangwx[0]; - float normalizer = local_normalizer[0]; - float normalizer3 = normalizer * normalizer * normalizer; - - // Write the outputs - gx += gid * size_t(axis_size) + lid * N_READS; - gw += gid * size_t(axis_size) + lid * N_READS; - if (lid * N_READS + N_READS <= axis_size) { - for (int i = 0; i < N_READS; i++) { - gx[i] = static_cast<T>( - thread_g[i] * thread_w[i] * normalizer - - thread_x[i] * meangwx * normalizer3); - if (has_w) { - gw[i] = static_cast<T>(thread_g[i] * thread_x[i] * normalizer); + threadgroup_barrier(mem_flags::mem_threadgroup); + if (simd_group_id == 0) { + sumx2 = simd_sum(local_sumx2[simd_lane_id]); + sumgwx = simd_sum(local_sumgwx[simd_lane_id]); + if (simd_lane_id == 0) { + local_meangwx[0] = sumgwx / axis_size; + local_normalizer[0] = metal::precise::rsqrt(sumx2 / axis_size + eps); } } - } else { - for (int i = 0; i < N_READS; i++) { - if ((lid * N_READS + i) < axis_size) { - gx[i] = static_cast<T>( + threadgroup_barrier(mem_flags::mem_threadgroup); + float meangwx = local_meangwx[0]; + float normalizer = local_normalizer[0]; + float normalizer3 = normalizer * normalizer * normalizer; + + // Write the outputs + device T* gx_out = gx + size_t(row) * axis_size + lid * N_READS; + if (lid * N_READS + N_READS <= axis_size) { + for (int i = 0; i < N_READS; i++) { + gx_out[i] = static_cast<T>( thread_g[i] * thread_w[i] * normalizer - thread_x[i] * meangwx * normalizer3); if (has_w) { - gw[i] = static_cast<T>(thread_g[i] * thread_x[i] * normalizer); + gw_acc[i] += thread_g[i] * thread_x[i] * normalizer; + } + } + } else { + for (int i = 0; i < N_READS; i++) { + if ((lid * N_READS + i) < axis_size) { + gx_out[i] = static_cast<T>( + thread_g[i] * thread_w[i] * normalizer - + thread_x[i] * meangwx * normalizer3); + if (has_w) { + gw_acc[i] += thread_g[i] * thread_x[i] * normalizer; + } + } + } + } + } + + if (has_w) { + gw += size_t(gid) * axis_size + lid * N_READS; + if (lid * N_READS + N_READS <= axis_size) { + for (int i = 0; i < N_READS; i++) { + gw[i] = static_cast<T>(gw_acc[i]); + } + } else { + for (int i = 0; i < N_READS; i++) { + if ((lid * N_READS + i) < axis_size) { + gw[i] = static_cast<T>(gw_acc[i]); } } }
diff --git ml-explore/mlx/mlx/backend/metal/kernels/sort.h Layr-Labs/mlx/mlx/backend/metal/kernels/sort.h index 068d43d12602485dac616225c7f78d6b82efa1fd..ea2640bace0ade86f39ea7a04784657765475a44 100644 --- ml-explore/mlx/mlx/backend/metal/kernels/sort.h +++ Layr-Labs/mlx/mlx/backend/metal/kernels/sort.h @@ -388,8 +388,11 @@ KernelMergeSort<T, U, ARG_SORT, BLOCK_THREADS, N_PER_THREAD>; using ValT = typename sort_kernel::ValT; using IdxT = typename sort_kernel::IdxT;   - auto in_block_idx = elem_to_loc(tid.y, nc_shape, in_nc_strides, nc_dim); - auto out_block_idx = elem_to_loc(tid.y, nc_shape, out_nc_strides, nc_dim); + // Signed offsets: a non-sorted axis may have a negative stride. + auto in_block_idx = + elem_to_loc<int64_t>(tid.y, nc_shape, in_nc_strides, nc_dim); + auto out_block_idx = + elem_to_loc<int64_t>(tid.y, nc_shape, out_nc_strides, nc_dim); inp += in_block_idx; out += out_block_idx;   @@ -532,7 +535,8 @@ ARG_SORT, BLOCK_THREADS, N_PER_THREAD>;   - auto block_idx = elem_to_loc(tid.y, nc_shape, nc_strides, nc_dim); + // Signed offset: a non-sorted axis may have a negative stride. + auto block_idx = elem_to_loc<int64_t>(tid.y, nc_shape, nc_strides, nc_dim); inp += block_idx; out_vals += tid.y * size_sorted_axis; out_idxs += tid.y * size_sorted_axis;
diff --git ml-explore/mlx/mlx/backend/metal/kernels/steel/attn/kernels/steel_attention.h Layr-Labs/mlx/mlx/backend/metal/kernels/steel/attn/kernels/steel_attention.h index 0d9628e83456fe885999cfc7f18610cd82982338..29fa7ba3964f9c9c3821b135f804fffe21270a33 100644 --- ml-explore/mlx/mlx/backend/metal/kernels/steel/attn/kernels/steel_attention.h +++ Layr-Labs/mlx/mlx/backend/metal/kernels/steel/attn/kernels/steel_attention.h @@ -428,7 +428,7 @@ STEEL_PRAGMA_UNROLL for (short id = 0; id < TD; id++) { STEEL_PRAGMA_UNROLL for (short ik = 0; ik < TK; ik++) { - if constexpr (BD == 128) { + if constexpr (BD >= 128) { simdgroup_barrier(mem_flags::mem_none); }   @@ -438,7 +438,7 @@ Vtile.template load<T, 1, 1, LDV_tgp, 1>( &Vs[Vs_offset + kk * LDV_tgp + dd]);   - if constexpr (BD == 128) { + if constexpr (BD >= 128) { simdgroup_barrier(mem_flags::mem_none); }
diff --git ml-explore/mlx/mlx/backend/metal/kernels/steel/attn/kernels/steel_attention.metal Layr-Labs/mlx/mlx/backend/metal/kernels/steel/attn/kernels/steel_attention.metal index 7bddfcb054d1c5f8337d93b452b98543757b6cac..fbd84004f0effc1890cd27068ba046866590eb11 100644 --- ml-explore/mlx/mlx/backend/metal/kernels/steel/attn/kernels/steel_attention.metal +++ Layr-Labs/mlx/mlx/backend/metal/kernels/steel/attn/kernels/steel_attention.metal @@ -12,9 +12,12 @@ "_wm" #wm "_wn" #wn "_mask" #mname, \ attention, dtype, bq, bk, bd, wm, wn, mtype, float)   #define instantiate_attn_shapes_helper(iname, itype, mname, mtype) \ + instantiate_attn(iname, itype, 32, 16, 256, 4, 1, mname, mtype) \ + instantiate_attn(iname, itype, 32, 16, 192, 4, 1, mname, mtype) \ instantiate_attn(iname, itype, 32, 16, 128, 4, 1, mname, mtype) \ instantiate_attn(iname, itype, 32, 32, 96, 4, 1, mname, mtype) \ instantiate_attn(iname, itype, 32, 32, 80, 4, 1, mname, mtype) \ + instantiate_attn(iname, itype, 32, 32, 72, 4, 1, mname, mtype) \ instantiate_attn(iname, itype, 32, 32, 64, 4, 1, mname, mtype)   #define instantiate_attn_mask_helper(iname, itype) \
diff --git ml-explore/mlx/mlx/backend/metal/kernels/steel/attn/kernels/steel_attention_nax.h Layr-Labs/mlx/mlx/backend/metal/kernels/steel/attn/kernels/steel_attention_nax.h index b48a9a942d320c20bcdea7a2be999d3b235072d9..4a5a9716fd25a85831b9da98e153091ca41028e2 100644 --- ml-explore/mlx/mlx/backend/metal/kernels/steel/attn/kernels/steel_attention_nax.h +++ Layr-Labs/mlx/mlx/backend/metal/kernels/steel/attn/kernels/steel_attention_nax.h @@ -484,3 +484,424 @@ } else { Otile.store(O, int(params->O_strides[2])); } } + +/////////////////////////////////////////////////////////////////////////////// +// Head-dim split attention kernel +/////////////////////////////////////////////////////////////////////////////// + +// Variant of attention_nax for wide heads (bd = 256). There, the per-simdgroup +// accumulator working set of attention_nax (TD output fragments plus the S +// fragments) is what gates tensor-unit throughput, so this kernel splits the +// head dim across the WN = 2 simdgroups of the second warp dimension: each +// simdgroup of a pair owns one half of D for Q@K.T and one half of Dv for P@V, +// halving its accumulator set. The pair exchanges its partial Q@K.T sums +// through threadgroup memory, then both simdgroups run softmax redundantly on +// the full S tile (the row statistics are cheap) and each accumulates P@V for +// its own half of Dv. + +// clang-format off +template < + typename T, + int BQ, + int BK, + int BD, + int WM, + int WN, + typename MaskType = float, + typename AccumType = float> +[[kernel, max_total_threads_per_threadgroup(WM * WN * 32)]] void attention_nax_dsplit( + const device T* Q [[buffer(0)]], + const device T* K [[buffer(1)]], + const device T* V [[buffer(2)]], + device T* O [[buffer(3)]], + const constant AttnParams* params [[buffer(4)]], + 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)]], + uint simd_lane_id [[thread_index_in_simdgroup]], + uint simd_group_id [[simdgroup_index_in_threadgroup]], + uint3 tid [[threadgroup_position_in_grid]], + uint3 lid [[thread_position_in_threadgroup]]) { // clang-format on + + // Pacifying compiler + (void)lid; + + // Move to correct block + ulong3 tidl{tid.x, tid.y, tid.z}; + + Q += tidl.z * params->Q_strides[0] + // Batch + tidl.y * params->Q_strides[1] + // Head + tidl.x * BQ * params->Q_strides[2]; // Sequence + + ulong kv_head_idx = int(tid.y) / params->gqa_factor; + K += tidl.z * params->K_strides[0] + // Batch + kv_head_idx * params->K_strides[1]; // Head + + V += tidl.z * params->V_strides[0] + // Batch + kv_head_idx * params->V_strides[1]; // Head + + O += tidl.z * params->O_strides[0] + // Batch + tidl.y * params->O_strides[1] + // Head + tidl.x * BQ * params->O_strides[2]; // Sequence + + if (has_mask) { + mask += tidl.z * mask_params->M_strides[0] + // Batch + tidl.y * mask_params->M_strides[1]; // Head + } + + const metal::uniform<float> scale2 = + make_uniform(params->scale) * make_uniform(1.44269504089f); + + // Prepare MMA tiles + constexpr short kU = 16; + + // The WM simdgroups along the first warp dimension split the Q sequence; + // the WN simdgroups along the second split the head dim. The exchange + // below reduces exactly one peer, so WN is fixed at 2. + static_assert(WN == 2, "The head-dim split kernel needs WN == 2"); + constexpr int kNWarps = WM; + static_assert( + BQ >= (kNWarps * kU) && BQ % (kNWarps * kU) == 0, + "Each simdgroup must host atleast 1 simdgroup matrix along Q sequence."); + + // Q seq frags per warp + constexpr int TQ = BQ / (kNWarps * kU); + // HeadDim frags over the full head dim + constexpr int TD = BD / kU; + // KV seq frags per warp + constexpr short TK = BK / kU; + + static_assert(TQ == 1, "Check TQ"); + static_assert(TD % WN == 0, "The head dim must split evenly across WN"); + + // HeadDim frags / columns owned by each of the WN simdgroups of a row group + constexpr int TDh = TD / WN; + constexpr int BDh = BD / WN; + + static_assert(TDh % 2 == 0, "P@V accumulates output fragments in pairs"); + static_assert(TK % 2 == 0, "S fragments are exchanged pair by pair"); + + const short row_group = simd_group_id / WN; + const short d_half = simd_group_id % WN; + + using otile_t = NAXTile<AccumType, TQ, TDh>; + otile_t Otile; + Otile.clear(); + + const short tm = kU * TQ * row_group; + Q += tm * int(params->Q_strides[2]) + d_half * BDh; + K += d_half * BDh; + V += d_half * BDh; + O += tm * int(params->O_strides[2]) + d_half * BDh; + + constexpr short kRowsPT = otile_t::kRowsPerThread; + + metal::vec<AccumType, kRowsPT> max_score; + metal::vec<AccumType, kRowsPT> sum_score{0}; + + STEEL_PRAGMA_UNROLL + for (short i = 0; i < kRowsPT; ++i) { + max_score[i] = Limits<AccumType>::finite_min; + } + + if (has_sinks) { + STEEL_PRAGMA_UNROLL + for (short i = 0; i < kRowsPT; ++i) { + max_score[i] = M_LOG2E_F * static_cast<AccumType>(sinks[tidl.y]); + sum_score[i] = 1; + } + } + + int kb_lim = params->NK; + int kb_min_causal = params->NK; + + if (do_causal) { + int q_max = (tid.x + 1) * BQ + params->qL_off; + kb_lim = (q_max + BK - 1) / BK; + kb_lim = min(params->NK, kb_lim); + + int q_min = tid.x * BQ + params->qL_off; + q_min = max(0, q_min); + kb_min_causal = (q_min / BK); + } + + const bool is_last_q = int(tid.x) == (params->NQ_aligned); + const short lim_rows_q = params->qL_rem - tm; + const short lim_rows_k = params->kL_rem; + + using stile_t = NAXTile<AccumType, TQ, TK>; + constexpr short kEPF = stile_t::NAXFrag_t::kElemsPerFrag; + + // One slot per (row group, half): a fragment pair in per-lane-linear + // layout. Both halves share the fragment-to-lane mapping, so the + // exchange needs no coordinate math. + threadgroup AccumType s_xchg[WM][WN][2 * kEPF * 32]; + + // Keep the simdgroup's Q half resident in registers for the whole KV + // loop: TDh fragments of T are cheap next to the accumulators. + NAXTile<T, 1, 1> Qtiles[TDh]; + STEEL_PRAGMA_UNROLL + for (short id = 0; id < TDh; id++) { + const int Q_load_off = id * kU; + if (!align_Q && is_last_q) { + Qtiles[id].load_rows( + Q + Q_load_off, int(params->Q_strides[2]), lim_rows_q); + } else { + Qtiles[id].load(Q + Q_load_off, int(params->Q_strides[2])); + } + } + + const short2 simd_coord = otile_t::NAXFrag_t::get_coord(); + const short sm = simd_coord.y; + const short sn = simd_coord.x; + + // Loop over KV seq length + for (int kb = 0; kb < kb_lim; kb++) { + const int is_last_k = (kb == (params->NK_aligned)); + + stile_t Stile; + Stile.clear(); + + // S = Q @ K.T, this half of D only, exchanged pair by pair. + STEEL_PRAGMA_UNROLL + for (short ik = 0; ik < TK; ik += 2) { + STEEL_PRAGMA_UNROLL + for (short id = 0; id < TDh; id++) { + NAXTile<T, 2, 1> Ktile; + const int K_load_off = ik * kU * int(params->K_strides[2]) + id * kU; + + if (!align_K && is_last_k) { + Ktile.load_rows( + K + K_load_off, int(params->K_strides[2]), lim_rows_k - ik * kU); + } else { + Ktile.load(K + K_load_off, int(params->K_strides[2])); + } + + stile_t::NAXFrag_t::mma( + Stile.frag_at(0, ik), + Stile.frag_at(0, ik + 1), + Qtiles[id].frag_at(0, 0), + metal::false_type{}, + Ktile.frag_at(0, 0), + Ktile.frag_at(1, 0), + metal::true_type{}); + } + + // Exchange the partial pair and reduce. + threadgroup AccumType* slot = s_xchg[row_group][d_half]; + thread auto& s0 = Stile.frag_at(0, ik); + thread auto& s1 = Stile.frag_at(0, ik + 1); + const short base = short(simd_lane_id) * (2 * kEPF); + STEEL_PRAGMA_UNROLL + for (short i = 0; i < kEPF; i++) { + slot[base + i] = s0[i]; + slot[base + kEPF + i] = s1[i]; + } + threadgroup_barrier(mem_flags::mem_threadgroup); + const threadgroup AccumType* peer = s_xchg[row_group][1 - d_half]; + STEEL_PRAGMA_UNROLL + for (short i = 0; i < kEPF; i++) { + s0[i] += peer[base + i]; + s1[i] += peer[base + kEPF + i]; + } + threadgroup_barrier(mem_flags::mem_threadgroup); + } + + // Scale S + STEEL_PRAGMA_UNROLL + for (short ii = 0; ii < stile_t::kElemsPerTile; ii++) { + Stile.elems()[ii] *= float(scale2); + } + + // Mask out length sequence + if (!align_K && is_last_k) { + constexpr auto neg_inf = Limits<AccumType>::finite_min; + + STEEL_PRAGMA_UNROLL + for (short ik = 0; ik < TK; ik++) { + const short col_pos = ik * kU + sn; + thread auto& fg = Stile.frag_at(0, ik); + + STEEL_PRAGMA_UNROLL + for (short ii = 0; ii < stile_t::kFragThrRows; ii++) { + STEEL_PRAGMA_UNROLL + for (short jj = 0; jj < stile_t::kFragThrCols; jj++) { + const auto loc = ii * stile_t::kFragThrCols + jj; + fg[loc] = ((col_pos + jj) < params->kL_rem) ? fg[loc] : neg_inf; + } + } + } + } + + // Mask out if causal + if (do_causal && kb >= kb_min_causal) { + constexpr auto neg_inf = Limits<AccumType>::finite_min; + + const int base_row = tid.x * BQ + params->qL_off + tm; + const int base_col = kb * BK; + + STEEL_PRAGMA_UNROLL + for (short ik = 0; ik < TK; ik++) { + thread auto& fg = Stile.frag_at(0, ik); + + STEEL_PRAGMA_UNROLL + for (short ii = 0; ii < stile_t::kFragThrRows; ii++) { + STEEL_PRAGMA_UNROLL + for (short jj = 0; jj < stile_t::kFragThrCols; jj++) { + const auto r = base_row + ii * stile_t::kFragRowsJump + sm; + const auto c = base_col + ik * kU + jj + sn; + const auto loc = ii * stile_t::kFragThrCols + jj; + fg[loc] = (r < c) ? neg_inf : fg[loc]; + } + } + } + } + + // Other masking as needed + if (has_mask) { + constexpr auto neg_inf = Limits<AccumType>::finite_min; + + const int base_row = tid.x * BQ + tm; + const int base_col = kb * BK; + + constexpr bool is_bool = is_same_v<MaskType, bool>; + using melem_t = typename metal::conditional_t<is_bool, bool, AccumType>; + using mtile_t = NAXTile<melem_t, TQ, TK>; + using mfrag_t = typename mtile_t::frag_type; + + if (base_row + kU <= params->qL && base_col + BK <= params->kL) { + STEEL_PRAGMA_UNROLL + for (short ik = 0; ik < TK; ik++) { + const int row_pos = base_row; + const int col_pos = base_col + ik * kU; + + mfrag_t mfrag; + mtile_t::NAXFrag_t::load( + mfrag, + mask, + int64_t(mask_params->M_strides[2]), + Int<1>{}, + row_pos, + col_pos); + + thread auto& fg = Stile.frag_at(0, ik); + + STEEL_PRAGMA_UNROLL + for (short jj = 0; jj < mtile_t::kElemsPerFrag; jj++) { + if constexpr (is_bool) { + fg[jj] = mfrag[jj] ? fg[jj] : neg_inf; + } else { + fg[jj] += M_LOG2E_F * AccumType(mfrag[jj]); + } + } + } + } else { + STEEL_PRAGMA_UNROLL + for (short ik = 0; ik < TK; ik++) { + const int row_pos = base_row; + const int col_pos = base_col + ik * kU; + + mfrag_t mfrag; + mtile_t::NAXFrag_t::load_safe( + mfrag, + mask, + int64_t(mask_params->M_strides[2]), + Int<1>{}, + params->qL, + params->kL, + row_pos, + col_pos); + + thread auto& fg = Stile.frag_at(0, ik); + + STEEL_PRAGMA_UNROLL + for (short jj = 0; jj < mtile_t::kElemsPerFrag; jj++) { + if constexpr (is_bool) { + fg[jj] = mfrag[jj] ? fg[jj] : neg_inf; + } else { + fg[jj] += M_LOG2E_F * AccumType(mfrag[jj]); + } + } + } + } + } + + // Do softmax (redundantly per half; the row statistics are cheap) + metal::vec<AccumType, kRowsPT> new_max; + metal::vec<AccumType, kRowsPT> factor; + STEEL_PRAGMA_UNROLL + for (short i = 0; i < kRowsPT; ++i) { + new_max[i] = max_score[i]; + } + + Stile.template row_reduce<MaxOp>(new_max); + Stile.template row_bin_op<ExpSubOp>(new_max); + + STEEL_PRAGMA_UNROLL + for (short i = 0; i < kRowsPT; ++i) { + factor[i] = fast::exp2(max_score[i] - new_max[i]); + max_score[i] = new_max[i]; + } + + STEEL_PRAGMA_UNROLL + for (short i = 0; i < kRowsPT; ++i) { + sum_score[i] = sum_score[i] * factor[i]; + } + + Stile.template row_reduce<SumOp>(sum_score); + + Otile.template row_bin_op<MulOp>(factor); + + simdgroup_barrier(mem_flags::mem_none); + + // O = P @ V, this half of Dv only. + STEEL_PRAGMA_UNROLL + for (short id = 0; id < TDh; id += 2) { + STEEL_PRAGMA_UNROLL + for (short ik = 0; ik < TK; ik++) { + NAXTile<T, 1, 2> Vtile; + + const int V_load_off = ik * kU * int(params->V_strides[2]) + id * kU; + + if (!align_K && is_last_k) { + Vtile.load_rows( + V + V_load_off, int(params->V_strides[2]), lim_rows_k - ik * kU); + } else { + Vtile.load(V + V_load_off, int(params->V_strides[2])); + } + + otile_t::NAXFrag_t::mma( + Otile.frag_at(0, id), + Otile.frag_at(0, id + 1), + Stile.frag_at(0, ik), + metal::false_type{}, + Vtile.frag_at(0, 0), + Vtile.frag_at(0, 1), + metal::false_type{}); + } + } + + // Next block + K += BK * int(params->K_strides[2]); + V += BK * int(params->V_strides[2]); + } + + // Normalize output + threadgroup_barrier(mem_flags::mem_none); + + metal::vec<AccumType, kRowsPT> rcp; + STEEL_PRAGMA_UNROLL + for (short i = 0; i < kRowsPT; ++i) { + rcp[i] = 1.f / sum_score[i]; + } + + Otile.template row_bin_op<MulOp>(rcp); + + if (!align_Q && is_last_q) { + if (lim_rows_q <= 0) + return; + Otile.store_rows(O, int(params->O_strides[2]), lim_rows_q); + } else { + Otile.store(O, int(params->O_strides[2])); + } +}
diff --git ml-explore/mlx/mlx/backend/metal/kernels/steel/attn/kernels/steel_attention_nax.metal Layr-Labs/mlx/mlx/backend/metal/kernels/steel/attn/kernels/steel_attention_nax.metal index c2b60b9cf0a83167b85e4260b616ec1bcebb3460..66d55539ab4f4c468ffaa98df4309f95dcc6dea9 100644 --- ml-explore/mlx/mlx/backend/metal/kernels/steel/attn/kernels/steel_attention_nax.metal +++ Layr-Labs/mlx/mlx/backend/metal/kernels/steel/attn/kernels/steel_attention_nax.metal @@ -11,11 +11,18 @@ "steel_attention_" #tname "_bq" #bq "_bk" #bk "_bd" #bd \ "_wm" #wm "_wn" #wn "_mask" #mname, \ attention_nax, dtype, bq, bk, bd, wm, wn, mtype, float)   -#define instantiate_attn_shapes_helper(iname, itype, mname, mtype) \ - instantiate_attn(iname, itype, 64, 32, 128, 4, 1, mname, mtype) \ - instantiate_attn(iname, itype, 64, 32, 96, 4, 1, mname, mtype) \ - instantiate_attn(iname, itype, 64, 32, 64, 4, 1, mname, mtype) \ - instantiate_attn(iname, itype, 64, 64, 128, 4, 1, mname, mtype) \ +#define instantiate_attn_dsplit(tname, dtype, bq, bk, bd, wm, wn, mname, mtype) \ + instantiate_kernel( \ + "steel_attention_dsplit_" #tname "_bq" #bq "_bk" #bk "_bd" #bd \ + "_wm" #wm "_wn" #wn "_mask" #mname, \ + attention_nax_dsplit, dtype, bq, bk, bd, wm, wn, mtype, float) + +#define instantiate_attn_shapes_helper(iname, itype, mname, mtype) \ + instantiate_attn_dsplit(iname, itype, 64, 32, 256, 4, 2, mname, mtype) \ + instantiate_attn(iname, itype, 64, 32, 128, 4, 1, mname, mtype) \ + instantiate_attn(iname, itype, 64, 32, 96, 4, 1, mname, mtype) \ + instantiate_attn(iname, itype, 64, 32, 64, 4, 1, mname, mtype) \ + instantiate_attn(iname, itype, 64, 64, 128, 4, 1, mname, mtype) \ instantiate_attn(iname, itype, 64, 64, 64, 4, 1, mname, mtype)   #define instantiate_attn_mask_helper(iname, itype) \
diff --git ml-explore/mlx/mlx/backend/metal/kernels/utils.h Layr-Labs/mlx/mlx/backend/metal/kernels/utils.h index 266f27e91c91e4ed2282ff3fae80672279022eaf..f15928282f39dd95652e476eb389bbf1016e49ef 100644 --- ml-explore/mlx/mlx/backend/metal/kernels/utils.h +++ Layr-Labs/mlx/mlx/backend/metal/kernels/utils.h @@ -446,3 +446,27 @@ template <typename T, typename U> struct ConditionalType<true, T, U> { using type = T; }; + +/////////////////////////////////////////////////////////////////////////////// +// Type casting utils +/////////////////////////////////////////////////////////////////////////////// + +template <typename U, typename T> +inline U cast_to(T val) { + return static_cast<U>(val); +} + +template <> +inline bool cast_to<bool, float>(float val) { + return (as_type<uint32_t>(val) & 0x7FFFFFFF) != 0; +} + +template <> +inline bool cast_to<bool, bfloat16_t>(bfloat16_t val) { + return (as_type<uint16_t>(val) & 0x7FFF) != 0; +} + +template <> +inline bool cast_to<bool, complex64_t>(complex64_t val) { + return cast_to<bool, float>(val.real) || cast_to<bool, float>(val.imag); +}
diff --git ml-explore/mlx/mlx/backend/metal/nojit_kernels.cpp Layr-Labs/mlx/mlx/backend/metal/nojit_kernels.cpp index 5da78db8a3b2e98ec50dcac60478423a371e0b51..3795a6fb228ad409c135a2afbfce5b4a6b18b8f9 100644 --- ml-explore/mlx/mlx/backend/metal/nojit_kernels.cpp +++ Layr-Labs/mlx/mlx/backend/metal/nojit_kernels.cpp @@ -503,7 +503,8 @@ int, int, int, int, - const array&) { + const array&, + bool) { return d.get_kernel(kernel_name, hash_name, func_consts); }
diff --git ml-explore/mlx/mlx/backend/metal/normalization.cpp Layr-Labs/mlx/mlx/backend/metal/normalization.cpp index 9a222cdd6c1c8b70d811bec00a7655e87438dbe9..f3370f087d93ce45a1b04c4ed58b01df302994b0 100644 --- ml-explore/mlx/mlx/backend/metal/normalization.cpp +++ Layr-Labs/mlx/mlx/backend/metal/normalization.cpp @@ -109,11 +109,9 @@ } array x_copy = contiguous_copy_gpu(x, s); return {x_copy, true}; }; - bool donate_g = inputs[2].is_donatable(); auto [x, copied] = check_input(inputs[0]); const array& w = inputs[1]; auto [g, g_copied] = check_input(inputs[2]); - donate_g |= g_copied; array& gx = outputs[0]; array& gw = outputs[1];   @@ -137,17 +135,22 @@ auto axis_size = static_cast<uint32_t>(x.shape().back()); int n_rows = x.data_size() / axis_size;   + const int target_groups = 512; + uint32_t rows_per_group = 1; + int n_groups = n_rows; + if (axis_size <= RMS_LOOPED_LIMIT) { + rows_per_group = (n_rows + target_groups - 1) / target_groups; + n_groups = (n_rows + rows_per_group - 1) / rows_per_group; + } + // Allocate the gradient accumulator gw and a temporary to store the // gradients before they are accumulated. - array gw_temp = - (has_w) ? array({n_rows, x.shape().back()}, gw.dtype(), nullptr, {}) : w; + array gw_temp = (has_w) + ? array({n_groups, x.shape().back()}, gw.dtype(), nullptr, {}) + : w; if (has_w) { - if (!g_in_gx && donate_g) { - gw_temp.copy_shared_buffer(g); - } else { - gw_temp.set_data(allocator::malloc(gw_temp.nbytes())); - compute_encoder.add_temporary(gw_temp); - } + gw_temp.set_data(allocator::malloc(gw_temp.nbytes())); + compute_encoder.add_temporary(gw_temp); } gw.set_data(allocator::malloc(gw.nbytes()));   @@ -174,7 +177,7 @@ size_t threadgroup_needed = (axis_size + n_reads - 1) / n_reads; size_t simds_needed = (threadgroup_needed + simd_size - 1) / simd_size; size_t threadgroup_size = simd_size * simds_needed; assert(threadgroup_size <= kernel->maxTotalThreadsPerThreadgroup()); - size_t n_threads = n_rows * threadgroup_size; + size_t n_threads = n_groups * threadgroup_size; grid_dims = MTL::Size(n_threads, 1, 1); group_dims = MTL::Size(threadgroup_size, 1, 1); } else { @@ -194,12 +197,16 @@ compute_encoder.set_output_array(gw_temp, 4); compute_encoder.set_bytes(eps_, 5); compute_encoder.set_bytes(axis_size, 6); compute_encoder.set_bytes(w_stride, 7); + if (axis_size <= looped_limit) { + compute_encoder.set_bytes(static_cast<uint32_t>(n_rows), 8); + compute_encoder.set_bytes(rows_per_group, 9); + } compute_encoder.dispatch_threads(grid_dims, group_dims); }   if (has_w) { ReductionPlan plan( - ReductionOpType::ContiguousStridedReduce, {n_rows}, {axis_size}); + ReductionOpType::ContiguousStridedReduce, {n_groups}, {axis_size}); strided_reduce_general_dispatch( gw_temp, gw, "sum", plan, {0}, compute_encoder, d, s); }
diff --git ml-explore/mlx/mlx/backend/metal/primitives.cpp Layr-Labs/mlx/mlx/backend/metal/primitives.cpp index 45929e27ddf2896bfba6f89501057e2a0116da96..d1d0e781cc3c87b7550f7fd9642b28452c075bb7 100644 --- ml-explore/mlx/mlx/backend/metal/primitives.cpp +++ Layr-Labs/mlx/mlx/backend/metal/primitives.cpp @@ -13,6 +13,7 @@ #include "mlx/backend/metal/device.h" #include "mlx/backend/metal/kernels.h" #include "mlx/backend/metal/utils.h" #include "mlx/dtype_utils.h" +#include "mlx/fast_primitives.h" #include "mlx/primitives.h" #include "mlx/scheduler.h" #include "mlx/utils.h" @@ -213,5 +214,27 @@ const std::vector<array>& inputs, std::vector<array>& outputs) { throw std::runtime_error("[LUF::eval_gpu] Metal LU factorization NYI."); } + +namespace fast { + +// There is no fused Metal cross entropy kernel yet +bool CrossEntropy::use_fallback(Stream s) { + return true; +} + +void CrossEntropy::eval_gpu( + const std::vector<array>& inputs, + std::vector<array>& outputs) { + throw std::runtime_error("[CrossEntropy::eval_gpu] Metal cross entropy NYI."); +} + +void CrossEntropyVJP::eval_gpu( + const std::vector<array>& inputs, + std::vector<array>& outputs) { + throw std::runtime_error( + "[CrossEntropyVJP::eval_gpu] Metal cross entropy NYI."); +} + +} // namespace fast   } // namespace mlx::core
diff --git ml-explore/mlx/mlx/backend/no_gpu/event.cpp Layr-Labs/mlx/mlx/backend/no_gpu/event.cpp index 6dde047ab4343aeda5e2f21ff83e7dfb4fe61ad3..8966b776137094279c4dff10a2dc95b064882b07 100644 --- ml-explore/mlx/mlx/backend/no_gpu/event.cpp +++ Layr-Labs/mlx/mlx/backend/no_gpu/event.cpp @@ -12,42 +12,52 @@ struct EventCounter { uint64_t value{0}; std::mutex mtx; std::condition_variable cv; + std::atomic<Error*> error; + + void wait(uint64_t val) { + std::unique_lock<std::mutex> lk(mtx); + if (value >= val) { + return; + } + cv.wait(lk, [this, val] { return value >= val; }); + } };   Event::Event(Stream stream) : stream_(stream) { - auto dtor = [](void* ptr) { delete static_cast<EventCounter*>(ptr); }; - event_ = std::shared_ptr<void>(new EventCounter{}, dtor); + event_ = std::make_shared<EventCounter>(); }   void Event::wait() { - auto ec = static_cast<EventCounter*>(event_.get()); - std::unique_lock<std::mutex> lk(ec->mtx); - if (ec->value >= value()) { - return; - } - ec->cv.wait(lk, [value = value(), ec] { return ec->value >= value; }); + check_error(); + cast<EventCounter>().wait(value()); + check_error(); }   void Event::wait(Stream stream) { - scheduler::enqueue(stream, [*this]() mutable { wait(); }); + scheduler::wait_event(stream, *this, [value = value()](Event& self) { + self.cast<EventCounter>().wait(value); + }); }   void Event::signal(Stream stream) { - scheduler::enqueue(stream, [*this]() mutable { - auto ec = static_cast<EventCounter*>(event_.get()); + scheduler::signal_event(stream, *this, [value = value()](Event& self) { + auto& ec = self.cast<EventCounter>(); { - std::lock_guard<std::mutex> lk(ec->mtx); - ec->value = value(); + std::lock_guard lk(ec.mtx); + ec.value = value; } - ec->cv.notify_all(); + ec.cv.notify_all(); }); }   bool Event::is_signaled() const { - auto ec = static_cast<EventCounter*>(event_.get()); - { - std::lock_guard<std::mutex> lk(ec->mtx); - return (ec->value >= value()); - } + auto& ec = cast<EventCounter>(); + std::lock_guard lk(ec.mtx); + return ec.value >= value(); } + +std::atomic<Error*>& Event::error() { + return cast<EventCounter>().error; +} + } // namespace mlx::core
diff --git ml-explore/mlx/mlx/backend/no_gpu/fence.cpp Layr-Labs/mlx/mlx/backend/no_gpu/fence.cpp index cd66d23cfe398497f66d92ed5b24c1d86ce5bc33..05852c860b7611443e015aeb67149608cc9e5756 100644 --- ml-explore/mlx/mlx/backend/no_gpu/fence.cpp +++ Layr-Labs/mlx/mlx/backend/no_gpu/fence.cpp @@ -1,54 +1,30 @@ // Copyright © 2024 Apple Inc.   -#include <condition_variable> -#include <mutex> - #include "mlx/fence.h" -#include "mlx/scheduler.h" +#include "mlx/event.h"   namespace mlx::core {   struct FenceImpl { - uint32_t count{0}; - uint32_t value{0}; - std::mutex mtx; - std::condition_variable cv; + uint32_t count; + Event event; + + FenceImpl(uint32_t count, Stream s) : count(count), event(s) {} };   -Fence::Fence(Stream) { - auto dtor = [](void* ptr) { delete static_cast<FenceImpl*>(ptr); }; - fence_ = std::shared_ptr<void>(new FenceImpl{}, dtor); +Fence::Fence(Stream s) { + fence_ = std::make_shared<FenceImpl>(0, s); }   -void Fence::wait(Stream stream, const array&) { - auto& f = *static_cast<FenceImpl*>(fence_.get()); - if (stream.device == Device::cpu) { - scheduler::enqueue(stream, [count = f.count, fence_ = fence_]() mutable { - auto& f = *static_cast<FenceImpl*>(fence_.get()); - std::unique_lock<std::mutex> lk(f.mtx); - if (f.value >= count) { - return; - } - f.cv.wait(lk, [&f, count] { return f.value >= count; }); - }); - } else { - throw std::runtime_error("[Fence::wait] Invalid stream."); - } +void Fence::wait(Stream s, const array&) { + cast<FenceImpl>().event.wait(s); }   -void Fence::update(Stream stream, const array&, bool) { - auto& f = *static_cast<FenceImpl*>(fence_.get()); +void Fence::update(Stream s, const array&, bool) { + auto& f = cast<FenceImpl>(); f.count++; - if (stream.device == Device::cpu) { - scheduler::enqueue(stream, [count = f.count, fence_ = fence_]() mutable { - auto& f = *static_cast<FenceImpl*>(fence_.get()); - std::unique_lock<std::mutex> lk(f.mtx); - f.value = count; - f.cv.notify_all(); - }); - } else { - throw std::runtime_error("[Fence::update] Invalid stream."); - } + f.event.set_value(f.count); + f.event.signal(s); }   } // namespace mlx::core
diff --git ml-explore/mlx/mlx/backend/no_gpu/primitives.cpp Layr-Labs/mlx/mlx/backend/no_gpu/primitives.cpp index b7d7a19467afd1c8c58a403e8f0e15633f53e8ce..f17f12cfaf6cd20c97cad210570122de9d2b437d 100644 --- ml-explore/mlx/mlx/backend/no_gpu/primitives.cpp +++ Layr-Labs/mlx/mlx/backend/no_gpu/primitives.cpp @@ -32,18 +32,24 @@ bool has_arr_mask, bool do_causal, bool is_training, bool output_logsumexp, + bool force_fused, Stream s) { + if (force_fused) { + throw std::invalid_argument( + "[scaled_dot_product_attention] force_fused=True but no fused " + "kernel is available in CPU backend."); + } return true; }   -bool fast::ScaledDotProductAttention::supports_bool_mask() { - return false; -} - bool fast::ScaledDotProductAttentionVJP::use_fallback( const array& q, Stream s) { return true; +} + +bool fast::ScaledDotProductAttention::supports_bool_mask() { + return false; }   NO_GPU(Abs) @@ -164,6 +170,8 @@ NO_GPU(View) NO_GPU(MaskedScatter)   namespace fast { +NO_GPU_USE_FALLBACK(CrossEntropy) +NO_GPU_MULTI(CrossEntropyVJP) NO_GPU_USE_FALLBACK(LayerNorm) NO_GPU_MULTI(LayerNormVJP) NO_GPU_USE_FALLBACK(RMSNorm)
diff --git ml-explore/mlx/mlx/distributed/mpi/mpi.cpp Layr-Labs/mlx/mlx/distributed/mpi/mpi.cpp index 3b176e6e6718cbbce451a29ae28188563dca4c2f..ea3960edc48c78fe34bd4d2aeb1df7e8d1ca5f21 100644 --- ml-explore/mlx/mlx/distributed/mpi/mpi.cpp +++ Layr-Labs/mlx/mlx/distributed/mpi/mpi.cpp @@ -166,6 +166,11 @@ bool init_safe() { if (!is_available()) { return false; } + // MPI_Init is an error to call twice, and init() can run more than once + // when it returns without a group. + if (initialized_) { + return true; + } bool success = init(nullptr, nullptr) == MPI_SUCCESS;   // Initialize custom types and ops @@ -491,6 +496,19 @@ std::shared_ptr<GroupImpl> init(bool strict /* = false */) { if (!mpi().init_safe()) { if (strict) { throw std::runtime_error("Cannot initialize MPI"); + } + return nullptr; + } + + // Open MPI initializes a world of size 1 for a program that was not started + // with mpirun, which is not a distributed group. + int size = 1; + mpi().size(mpi().world(), &size); + if (size <= 1) { + if (strict) { + throw std::runtime_error( + "[mpi] The world has a single process. Launch with mpirun to " + "initialize the mpi backend."); } return nullptr; }
diff --git ml-explore/mlx/mlx/distributed/ring/ring.cpp Layr-Labs/mlx/mlx/distributed/ring/ring.cpp index 3e0c2a3221001e1557c3e1cb5388ee4c5ce490bc..9a81010e34370ebc5c342fe808c5d9ad20d32401 100644 --- ml-explore/mlx/mlx/distributed/ring/ring.cpp +++ Layr-Labs/mlx/mlx/distributed/ring/ring.cpp @@ -5,6 +5,7 @@ #include <netinet/tcp.h> #include <sys/socket.h> #include <unistd.h>   +#include <algorithm> #include <chrono> #include <fstream> #include <future> @@ -36,6 +37,10 @@ constexpr const size_t ALL_SUM_BUFFERS = 2; constexpr const int CONN_ATTEMPTS = 5; constexpr const int CONN_WAIT = 1000; constexpr const char* RING_TAG = "[ring]"; +// send(2) and recv(2) reject a length above INT_MAX with EINVAL, so a single +// transfer of 2 GiB or more fails outright rather than being carried in +// pieces. +constexpr const size_t MAX_IO_BYTES = 1024 * 1024 * 1024;   using GroupImpl = mlx::core::distributed::detail::GroupImpl; using json = nlohmann::json; @@ -174,7 +179,8 @@ }   if (!recvs_.empty()) { auto& task = recvs_.front(); - ssize_t r = ::recv(fd_, task.buffer, task.size, 0); + ssize_t r = + ::recv(fd_, task.buffer, std::min(task.size, MAX_IO_BYTES), 0); if (r > 0) { task.buffer = static_cast<char*>(task.buffer) + r; task.size -= r; @@ -191,7 +197,8 @@ } } if (!sends_.empty()) { auto& task = sends_.front(); - ssize_t r = ::send(fd_, task.buffer, task.size, 0); + ssize_t r = + ::send(fd_, task.buffer, std::min(task.size, MAX_IO_BYTES), 0); if (r > 0) { task.buffer = static_cast<char*>(task.buffer) + r; task.size -= r;
diff --git ml-explore/mlx/mlx/einsum.cpp Layr-Labs/mlx/mlx/einsum.cpp index b683733d0bb040ff1770ac6ef1de2e9f1c92bde4..dcaa8c51ea5ba43e936bacbb7c43a81ccaceade4 100644 --- ml-explore/mlx/mlx/einsum.cpp +++ Layr-Labs/mlx/mlx/einsum.cpp @@ -95,10 +95,14 @@ } std::sort(rhs.begin(), rhs.end()); } std::vector<std::string> input_list; - std::stringstream ss(lhs); - std::string token; - while (getline(ss, token, ',')) { - input_list.push_back(token); + for (size_t start = 0;;) { + auto pos = lhs.find(',', start); + if (pos == std::string::npos) { + input_list.push_back(lhs.substr(start)); + break; + } + input_list.push_back(lhs.substr(start, pos - start)); + start = pos + 1; } return {input_list, rhs}; } @@ -356,7 +360,7 @@ std::vector<int> b_contract, std::vector<int> b_batch, std::vector<int> b_concat, StreamOrDevice s) { - // Broadcast contracting dimensions + // Broadcast contracting and batch dimensions. { auto a_shape = a.shape(); auto b_shape = b.shape(); @@ -364,6 +368,11 @@ for (int i = 0; i < a_contract.size(); ++i) { auto d = std::max(a.shape(a_contract[i]), b.shape(b_contract[i])); a_shape[a_contract[i]] = d; b_shape[b_contract[i]] = d; + } + for (int i = 0; i < a_batch.size(); ++i) { + auto d = std::max(a.shape(a_batch[i]), b.shape(b_batch[i])); + a_shape[a_batch[i]] = d; + b_shape[b_batch[i]] = d; } a = broadcast_to(a, a_shape, s); b = broadcast_to(b, b_shape, s);
diff --git ml-explore/mlx/mlx/error.h Layr-Labs/mlx/mlx/error.h new file mode 100644 index 0000000000000000000000000000000000000000..ba1164f192eaac25579a60504aae25961a02dd29 --- /dev/null +++ Layr-Labs/mlx/mlx/error.h @@ -0,0 +1,49 @@ +// Copyright © 2026 Apple Inc. + +#pragma once + +#include <atomic> +#include <memory> +#include <string> + +namespace mlx::core { + +class Error { + public: + // TODO: Use std::atomic<std::shared_ptr> when it gets supported in Xcode. + using Message = std::shared_ptr<std::string>; + + void set_message(Message msg) { + std::atomic_store(&message_, std::move(msg)); + } + + bool valid() const { + auto msg = std::atomic_load(&message_); + return msg.get(); + } + + // If |ptr| is a valid event, copy and return true. + bool store_if_valid(const Error* ptr) { + if (ptr && this != ptr) { + Message msg = std::atomic_load(&ptr->message_); + if (msg) { + set_message(std::move(msg)); + return true; + } + } + return false; + } + + // If current error is valid, throw and clear. + void check() { + auto msg = std::atomic_exchange(&message_, {}); + if (msg) { + throw std::runtime_error(*msg); + } + } + + private: + Message message_; +}; + +} // namespace mlx::core
diff --git ml-explore/mlx/mlx/event.h Layr-Labs/mlx/mlx/event.h index 66a6a75df5a42ec14b8dfa000c457af2ddcdfd3a..cf2d5cc7d6f9fd7c9ee0408666593d65049f46b4 100644 --- ml-explore/mlx/mlx/event.h +++ Layr-Labs/mlx/mlx/event.h @@ -5,6 +5,7 @@ #include <cstdint> #include <memory> #include <stdexcept>   +#include "mlx/error.h" #include "mlx/stream.h"   namespace mlx::core { @@ -26,6 +27,26 @@ // Check if the event has been signaled at its current value bool is_signaled() const;   + // Associate an error to the event + void set_error(Error& err) { + error().store(&err); + } + + // Get the error associated with the event + Error* load_error() const { + if (!valid()) { + return nullptr; + } + return error().load(); + } + + // Throw and clear the associated error + void check_error() { + if (auto* p = load_error(); p) { + p->check(); + } + } + // Check if the event is valid bool valid() const { return event_ != nullptr; @@ -47,7 +68,18 @@ } return stream_; }   + template <typename T> + auto& cast() const { + return *static_cast<T*>(event_.get()); + } + private: + std::atomic<Error*>& error(); + + const std::atomic<Error*>& error() const { + return const_cast<Event*>(this)->error(); + } + // Default constructed stream should never be used // since the event is not yet valid Stream stream_{0, Device::cpu};
diff --git ml-explore/mlx/mlx/fast.cpp Layr-Labs/mlx/mlx/fast.cpp index a668fe9abd29daa07f91a14e6627367f0f1b2b0f..f45724dc9223988441c6ae2b15771a7d7fc397c2 100644 --- ml-explore/mlx/mlx/fast.cpp +++ Layr-Labs/mlx/mlx/fast.cpp @@ -187,6 +187,107 @@ const RMSNormVJP& a_other = static_cast<const RMSNormVJP&>(other); return eps_ == a_other.eps_; }   +array cross_entropy( + const array& logits, + const array& targets, + StreamOrDevice s_ /* = {} */) { + if (logits.ndim() < 1) { + throw std::invalid_argument( + "[cross_entropy] logits must have at least 1 dimension but got input " + "with 0 dimensions."); + } + auto expected = logits.shape(); + expected.pop_back(); + if (targets.shape() != expected) { + std::ostringstream msg; + msg << "[cross_entropy] targets shape " << targets.shape() + << " does not match logits shape " << logits.shape() + << " with the last axis removed."; + throw std::invalid_argument(msg.str()); + } + if (!issubdtype(logits.dtype(), floating)) { + std::ostringstream msg; + msg << "[cross_entropy] Received unsupported logits type " << logits.dtype() + << "."; + throw std::invalid_argument(msg.str()); + } + if (!issubdtype(targets.dtype(), integer)) { + std::ostringstream msg; + msg << "[cross_entropy] targets must be integer class indices but got " + << targets.dtype() << "."; + throw std::invalid_argument(msg.str()); + } + + auto s = to_stream(s_); + auto fallback = [s](const std::vector<array>& inputs) { + auto& x = inputs[0]; + auto& y = inputs[1]; + auto score = + squeeze(take_along_axis(x, expand_dims(y, -1, s), -1, s), -1, s); + auto loss = subtract(logsumexp(x, -1, /* keepdims= */ false, s), score, s); + return std::vector<array>{astype(loss, float32, s)}; + }; + + auto passed_targets = astype(targets, int32, s); + + if (!CrossEntropy::use_fallback(s)) { + return array( + expected, + float32, + std::make_shared<CrossEntropy>(s, fallback), + {logits, passed_targets}); + } + return fallback({logits, passed_targets})[0]; +} + +std::vector<array> CrossEntropy::vjp( + const std::vector<array>& primals, + const std::vector<array>& cotangents, + const std::vector<int>& argnums, + const std::vector<array>& outputs) { + assert(primals.size() == 2); + assert(outputs.size() == 1); + assert(cotangents.size() == 1); + + for (auto arg : argnums) { + if (arg != 0) { + throw std::invalid_argument( + "[cross_entropy] Cannot differentiate with respect to the targets."); + } + } + + auto s = stream(); + auto fallback = [s](const std::vector<array>& inputs) { + auto& x = inputs[0]; + auto& y = inputs[1]; + auto& loss = inputs[2]; + auto& g = inputs[3]; + + auto score = + squeeze(take_along_axis(x, expand_dims(y, -1, s), -1, s), -1, s); + auto lse = add(loss, astype(score, float32, s), s); + auto p = + exp(subtract(astype(x, float32, s), expand_dims(lse, -1, s), s), s); + Shape class_shape(x.ndim(), 1); + class_shape.back() = x.shape(-1); + auto onehot = astype( + equal( + expand_dims(y, -1, s), + reshape(arange(x.shape(-1), y.dtype(), s), class_shape, s), + s), + float32, + s); + auto gx = multiply(expand_dims(g, -1, s), subtract(p, onehot, s), s); + return std::vector<array>{astype(gx, x.dtype(), s)}; + }; + + return {array( + primals[0].shape(), + primals[0].dtype(), + std::make_shared<CrossEntropyVJP>(s, fallback), + {primals[0], primals[1], outputs[0], cotangents[0]})}; +} + array layer_norm( const array& x, const std::optional<array>& weight, @@ -618,7 +719,8 @@ const float scale, const std::string& mask_mode /* = "" */, std::optional<array> mask_arr /* = {} */, const std::optional<array>& sinks /* = {} */, - StreamOrDevice s /* = {}*/) { + bool force_fused /* = false */, + StreamOrDevice s /* = {} */) { for (const auto& tensor : {queries, keys, values}) { if (tensor.ndim() != 4) { std::ostringstream msg; @@ -834,6 +936,7 @@ has_arr_mask, do_causal, is_training, output_logsumexp, + force_fused, stream)) { if (has_bool_mask && !ScaledDotProductAttention::supports_bool_mask()) { // Convert bool mask to additive mask. @@ -846,7 +949,13 @@ 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); + stream, + fallback, + scale, + do_causal, + has_sinks, + output_logsumexp, + force_fused); if (output_logsumexp) { return array::make_arrays( {std::move(out_shape), Shape{q.shape(0), q.shape(1), q.shape(2), 1}}, @@ -912,7 +1021,8 @@ const ScaledDotProductAttention& a_other = static_cast<const ScaledDotProductAttention&>(other); return scale_ == a_other.scale_ && do_causal_ == a_other.do_causal_ && has_sinks_ == a_other.has_sinks_ && - output_logsumexp_ == a_other.output_logsumexp_; + output_logsumexp_ == a_other.output_logsumexp_ && + force_fused_ == a_other.force_fused_; }   bool ScaledDotProductAttentionVJP::is_equivalent(const Primitive& other) const {
diff --git ml-explore/mlx/mlx/fence.h Layr-Labs/mlx/mlx/fence.h index 0ececdb6d7be1d602d12782f93434015cba86e72..3fd5da333b109e30314bcc37964cf0279448bcb8 100644 --- ml-explore/mlx/mlx/fence.h +++ Layr-Labs/mlx/mlx/fence.h @@ -32,8 +32,13 @@ void update(Stream stream, const array& x, bool cross_device); void wait(Stream stream, const array& x);   + template <typename T> + auto& cast() const { + return *static_cast<T*>(fence_.get()); + } + private: - std::shared_ptr<void> fence_{nullptr}; + std::shared_ptr<void> fence_; };   } // namespace mlx::core
diff --git ml-explore/mlx/mlx/fft.cpp Layr-Labs/mlx/mlx/fft.cpp index 8ddc1aca46d281c6d01dcee9fc3d9f2a59d18933..06860a0e3ac6c0e9a7b4262b8604bf8bc203f8b8 100644 --- ml-explore/mlx/mlx/fft.cpp +++ Layr-Labs/mlx/mlx/fft.cpp @@ -240,10 +240,16 @@ StreamOrDevice s /* = {} */) { return fft_impl(a, true, true, norm, s); }   -array fftshift( +namespace { + +// Shared implementation for fftshift/ifftshift: validates axes and computes +// the per-axis roll amount, differing only in shift sign and error prefix. +array fftshift_impl( + const char* name, const array& a, const std::vector<int>& axes, - StreamOrDevice s /* = {} */) { + bool inverse, + StreamOrDevice s) { if (axes.empty()) { return a; } @@ -254,41 +260,32 @@ // Convert negative axes to positive int axis = ax < 0 ? ax + a.ndim() : ax; if (axis < 0 || axis >= a.ndim()) { std::ostringstream msg; - msg << "[fftshift] Invalid axis " << ax << " for array with " << a.ndim() - << " dimensions."; + msg << "[" << name << "] Invalid axis " << ax << " for array with " + << a.ndim() << " dimensions."; throw std::invalid_argument(msg.str()); } // Match NumPy's implementation - shifts.push_back(a.shape(axis) / 2); + int shift = a.shape(axis) / 2; + shifts.push_back(inverse ? -shift : shift); }   return roll(a, shifts, axes, s); }   +} // namespace + +array fftshift( + const array& a, + const std::vector<int>& axes, + StreamOrDevice s /* = {} */) { + return fftshift_impl("fftshift", a, axes, false, s); +} + array ifftshift( const array& a, const std::vector<int>& axes, StreamOrDevice s /* = {} */) { - if (axes.empty()) { - return a; - } - - Shape shifts; - for (int ax : axes) { - // Convert negative axes to positive - int axis = ax < 0 ? ax + a.ndim() : ax; - if (axis < 0 || axis >= a.ndim()) { - std::ostringstream msg; - msg << "[ifftshift] Invalid axis " << ax << " for array with " << a.ndim() - << " dimensions."; - throw std::invalid_argument(msg.str()); - } - // Match NumPy's implementation - int size = a.shape(axis); - shifts.push_back(-(size / 2)); - } - - return roll(a, shifts, axes, s); + return fftshift_impl("ifftshift", a, axes, true, s); }   // Default versions that operate on all axes
diff --git ml-explore/mlx/mlx/io/gguf.cpp Layr-Labs/mlx/mlx/io/gguf.cpp index 40cca573e5b0f7fbdff71910b53920eb9c375e7f..6c27c4698763088f749faeb4b083b381e33950fc 100644 --- ml-explore/mlx/mlx/io/gguf.cpp +++ Layr-Labs/mlx/mlx/io/gguf.cpp @@ -124,8 +124,7 @@ case GGUF_VALUE_TYPE_BOOL: value = array(val->boolval, bool_); break; case GGUF_VALUE_TYPE_STRING: - value = - std::string(val->string.string, static_cast<int>(val->string.len)); + value = std::string(val->string.string, val->string.len); break; case GGUF_VALUE_TYPE_FLOAT64: value = array(val->float64, float32); @@ -174,7 +173,7 @@ std::vector<std::string> strs(size); for (auto& str : strs) { auto str_val = reinterpret_cast<gguf_string*>(data); data += (str_val->len + sizeof(gguf_string)); - str = std::string(str_val->string, static_cast<int>(str_val->len)); + str = std::string(str_val->string, str_val->len); ctx->off += (str_val->len + sizeof(gguf_string)); } value = std::move(strs); @@ -200,10 +199,102 @@ ctx->off += pv->nbytes(); } }   +inline size_t gguf_value_type_size(uint32_t type) { + switch (type) { + case GGUF_VALUE_TYPE_BOOL: + case GGUF_VALUE_TYPE_UINT8: + case GGUF_VALUE_TYPE_INT8: + return 1; + case GGUF_VALUE_TYPE_UINT16: + case GGUF_VALUE_TYPE_INT16: + return 2; + case GGUF_VALUE_TYPE_UINT32: + case GGUF_VALUE_TYPE_INT32: + case GGUF_VALUE_TYPE_FLOAT32: + return 4; + case GGUF_VALUE_TYPE_UINT64: + case GGUF_VALUE_TYPE_INT64: + case GGUF_VALUE_TYPE_FLOAT64: + return 8; + default: + return 0; + } +} + +void check_metadata_value_in_file( + const gguf_ctx* ctx, + uint32_t type, + const gguf_value* val) { + auto end = ctx->data + ctx->size; + // Bytes available from a pointer up to the end of the mapping; 0 if the + // pointer lies outside [ctx->data, end]. + auto avail = [&](const uint8_t* p) -> size_t { + return (p < ctx->data || p > end) ? 0 : static_cast<size_t>(end - p); + }; + auto base = reinterpret_cast<const uint8_t*>(val); + auto fail = [](const char* what) { + std::ostringstream msg; + msg << "[load_gguf] " << what + << " Perhaps an incomplete download or corrupt file?"; + throw std::runtime_error(msg.str()); + }; + + size_t fixed = gguf_value_type_size(type); + if (fixed) { + if (fixed > avail(base)) { + fail("Metadata value extends past the end of the file."); + } + return; + } + + auto check_string = [&](const uint8_t* p) -> const uint8_t* { + uint64_t len = reinterpret_cast<const gguf_string*>(p)->len; + if (sizeof(uint64_t) + len > avail(p)) { + fail("String metadata value extends past the end of the file."); + } + return p + sizeof(uint64_t) + len; + }; + + if (type == GGUF_VALUE_TYPE_STRING) { + if (sizeof(uint64_t) > avail(base)) { + fail("String metadata value extends past the end of the file."); + } + check_string(base); + return; + } + + if (type == GGUF_VALUE_TYPE_ARRAY) { + if (gguf_array_header_size > avail(base)) { + fail("Metadata value extends past the end of the file."); + } + const uint8_t* elt = base + gguf_array_header_size; + size_t elt_size = gguf_value_type_size(val->array.type); + if (elt_size) { + if (val->array.len > avail(elt) / elt_size) { + fail("Array metadata value extends past the end of the file."); + } + return; + } + if (val->array.type == GGUF_VALUE_TYPE_STRING) { + const uint8_t* p = elt; + for (uint64_t i = 0; i < val->array.len; i++) { + if (sizeof(uint64_t) > avail(p)) { + fail("Array metadata value extends past the end of the file."); + } + p = check_string(p); + } + } + return; + } + + throw std::runtime_error("[load_gguf] Received unexpected type."); +} + std::unordered_map<std::string, GGUFMetaData> load_metadata(gguf_ctx* ctx) { std::unordered_map<std::string, GGUFMetaData> metadata; gguf_key key; while (gguf_get_key(ctx, &key)) { + check_metadata_value_in_file(ctx, key.type, key.val); std::string key_name = std::string(key.name, key.namelen); auto& val = metadata.insert({key_name, GGUFMetaData{}}).first->second; set_mx_value_from_gguf(ctx, key.type, key.val, val); @@ -211,10 +302,6 @@ } return metadata; }   -// gguflib computes weights_data as ctx->data + ctx->data_off + the tensor's -// offset field in unsigned arithmetic, without comparing the result against the -// mapping, so a crafted offset can point outside the file or -- if the addition -// wraps -- back inside it at the wrong bytes. void check_tensor_in_file(const gguf_ctx* ctx, const gguf_tensor& tensor) { auto fail = [&tensor](const std::string& what) { std::ostringstream msg;
diff --git ml-explore/mlx/mlx/ops.cpp Layr-Labs/mlx/mlx/ops.cpp index 9c8db3a26de8c917802aa4a1778918b91a6bef7a..5fe3d9be59382a47d0e17d894b4c0926cc14c036 100644 --- ml-explore/mlx/mlx/ops.cpp +++ Layr-Labs/mlx/mlx/ops.cpp @@ -1,4 +1,4 @@ -// Copyright © 2023-2024 Apple Inc. +// Copyright © 2023-2026 Apple Inc.   // Required for using M_PI in MSVC. #define _USE_MATH_DEFINES @@ -16,6 +16,7 @@ #include "mlx/ops.h" #include "mlx/primitives.h" #include "mlx/transforms.h" #include "mlx/transforms_impl.h" +#include "mlx/types/limits.h" #include "mlx/utils.h"   namespace mlx::core { @@ -271,6 +272,7 @@ array linspace( double start, double stop, int num /* = 50 */, + bool endpoint /* = true */, Dtype dtype /* = float32 */, StreamOrDevice s /* = {} */) { if (num < 0) { @@ -282,8 +284,11 @@ if (num == 1) { return astype(array({start}), dtype, s); } auto inner_type = dtype == float64 ? float64 : float32; + // Without the endpoint the samples are spaced so that `stop` would be the + // next one after the last, i.e. the step is (stop - start) / num. + auto denominator = endpoint ? num - 1 : num; array t = - divide(arange(0, num, inner_type, s), array(num - 1, inner_type), s); + divide(arange(0, num, inner_type, s), array(denominator, inner_type), s); array t_bar = subtract(array(1, inner_type), t, s); return astype( add(multiply(t_bar, array(start, inner_type), s), @@ -1138,13 +1143,7 @@ const array& a, const Shape& indices, int axis, StreamOrDevice s /* = {} */) { - auto ax = axis < 0 ? axis + a.ndim() : axis; - if (ax < 0 || ax >= a.ndim()) { - std::ostringstream msg; - msg << "Invalid axis (" << axis << ") passed to split" - << " for array with shape " << a.shape() << "."; - throw std::invalid_argument(msg.str()); - } + auto ax = normalize_axis_index(axis, a.ndim(), "[split] ");   if (indices.empty()) { return {a}; @@ -1186,20 +1185,14 @@ }   std::vector<array> split(const array& a, int num_splits, int axis, StreamOrDevice s /* = {} */) { - auto ax = axis < 0 ? axis + a.ndim() : axis; - if (ax < 0 || ax >= a.ndim()) { - std::ostringstream msg; - msg << "Invalid axis " << axis << " passed to split" - << " for array with shape " << a.shape() << "."; - throw std::invalid_argument(msg.str()); - } + auto ax = normalize_axis_index(axis, a.ndim(), "[split] "); if (num_splits <= 0) { std::ostringstream msg; msg << "[split] num_splits must be positive and non-zero but got " << num_splits << "."; throw std::invalid_argument(msg.str()); } - auto q_and_r = std::ldiv(a.shape(axis), num_splits); + auto q_and_r = std::ldiv(a.shape(ax), num_splits); if (q_and_r.rem) { std::ostringstream msg; msg << "Array split does not result in sub arrays with equal size:" @@ -1212,7 +1205,7 @@ Shape indices(num_splits - 1); for (int i = 0; i < indices.size(); ++i) { indices[i] = (i + 1) * split_size; } - return split(a, indices, axis, s); + return split(a, indices, ax, s); }   std::vector<array> @@ -1223,13 +1216,7 @@ std::vector<array> unstack(const array& a, int axis, StreamOrDevice s /* = {} */) { auto ndim = static_cast<int>(a.ndim()); - auto ax = axis < 0 ? axis + ndim : axis; - if (ax < 0 || ax >= ndim) { - std::ostringstream msg; - msg << "[unstack] Invalid axis " << axis << " for array with " << ndim - << " dimensions."; - throw std::invalid_argument(msg.str()); - } + auto ax = normalize_axis_index(axis, ndim, "[unstack] "); auto n = a.shape(ax); std::vector<array> res; res.reserve(n); @@ -1325,7 +1312,9 @@ throw std::invalid_argument(msg.str()); };   auto shape = arrays[0].shape(); - shape[ax] = 0; + // Accumulate the concatenation axis in 64 bits so a total that does not fit + // in a shape dimension is reported rather than silently wrapping. + int64_t concat_size = 0; // Make the output shape and validate that all arrays have the same shape // except for the concatenation axis. for (auto& a : arrays) { @@ -1344,8 +1333,9 @@ if (a.shape(i) != shape[i]) { throw_invalid_shapes(); } } - shape[ax] += a.shape(ax); + concat_size += a.shape(ax); } + shape[ax] = safe_cast(concat_size, "concatenate");   // Promote all the arrays to the same type auto dtype = result_type(arrays); @@ -1421,7 +1411,8 @@ out = broadcast_to(out, shape, s);   // Reshape back into a contiguous array where S_axis is now S_axis * repeats shape.erase(shape.begin() + axis + 1); - shape[axis] *= repeats; + shape[axis] = + safe_cast(static_cast<int64_t>(shape[axis]) * repeats, "repeat"); out = reshape(out, shape, s);   return out; @@ -1628,7 +1619,10 @@ throw std::invalid_argument(msg.str()); }   auto ax = axes[i] < 0 ? a.ndim() + axes[i] : axes[i]; - out_shape[ax] += low_pad_size[i] + high_pad_size[i]; + out_shape[ax] = safe_cast( + static_cast<int64_t>(out_shape[ax]) + low_pad_size[i] + + high_pad_size[i], + "pad"); }   if (mode == "constant") { @@ -2107,11 +2101,11 @@ }   auto type_to_max = [](const auto& dtype) -> float { if (dtype == float32) { - return std::numeric_limits<float>::max(); + return numeric_limits<float>::max(); } else if (dtype == bfloat16) { - return std::numeric_limits<bfloat16_t>::max(); + return numeric_limits<bfloat16_t>::max(); } else if (dtype == float16) { - return std::numeric_limits<float16_t>::max(); + return numeric_limits<float16_t>::max(); } else { std::ostringstream msg; msg << "[nan_to_num] Does not yet support given type: " << dtype << "."; @@ -2417,6 +2411,16 @@ add(median_a, astype(slice(sorted_a, start, stop, s), dtype, s), s), array(0.5, dtype), s); } + // Sorting moves NaN to the end, so the midpoint slice never selects it. + // Propagate it explicitly to stay consistent with max, min and mean. + if (issubdtype(a.dtype(), inexact)) { + median_a = where( + any(isnan(flat_a, s), -1, /* keepdims = */ true, s), + array(std::numeric_limits<float>::quiet_NaN(), dtype), + median_a, + s); + } + median_a = squeeze(median_a, -1, s); if (keepdims) { median_a = expand_dims(median_a, sorted_axes, s); @@ -2450,7 +2454,17 @@ int ddof /* = 0*/, StreamOrDevice s /* = {}*/) { auto dtype = at_least_float(a.dtype()); auto mu = mean(a, axes, /* keepdims= */ true, s); - auto v = sum(square(subtract(a, mu, s), s), axes, keepdims, s); + auto d = subtract(a, mu, s); + // The variance of complex values is the mean squared magnitude. Squaring the + // deviations directly gives a complex result which can even be negative, so + // multiply by the conjugate instead. + auto sq = issubdtype(dtype, complexfloating) + ? real(multiply(d, conjugate(d, s), s), s) + : square(d, s); + if (issubdtype(dtype, complexfloating)) { + dtype = float32; + } + auto v = sum(sq, axes, keepdims, s);   if (ddof != 0) { auto normalizer = maximum( @@ -2780,17 +2794,9 @@ }   /** Returns a sorted copy of the array along a given axis. */ array sort(const array& a, int axis, StreamOrDevice s /* = {} */) { - // Check for valid axis - if (axis + static_cast<int>(a.ndim()) < 0 || - axis >= static_cast<int>(a.ndim())) { - std::ostringstream msg; - msg << "[sort] Received invalid axis " << axis << " for array with " - << a.ndim() << " dimensions."; - throw std::invalid_argument(msg.str()); - } - + auto ax = normalize_axis_index(axis, a.ndim(), "[sort] "); return array( - a.shape(), a.dtype(), std::make_shared<Sort>(to_stream(s), axis), {a}); + a.shape(), a.dtype(), std::make_shared<Sort>(to_stream(s), ax), {a}); }   /** Returns indices that sort the flattened array. */ @@ -2801,17 +2807,9 @@ }   /** Returns indices that sort the array along a given axis. */ array argsort(const array& a, int axis, StreamOrDevice s /* = {} */) { - // Check for valid axis - if (axis + static_cast<int>(a.ndim()) < 0 || - axis >= static_cast<int>(a.ndim())) { - std::ostringstream msg; - msg << "[argsort] Received invalid axis " << axis << " for array with " - << a.ndim() << " dimensions."; - throw std::invalid_argument(msg.str()); - } - + auto ax = normalize_axis_index(axis, a.ndim(), "[argsort] "); return array( - a.shape(), uint32, std::make_shared<ArgSort>(to_stream(s), axis), {a}); + a.shape(), uint32, std::make_shared<ArgSort>(to_stream(s), ax), {a}); }   /** @@ -2833,14 +2831,7 @@ int kth, int axis, StreamOrDevice s /* = {} */) { // Check for valid axis - if (axis + static_cast<int>(a.ndim()) < 0 || - axis >= static_cast<int>(a.ndim())) { - std::ostringstream msg; - msg << "[partition] Received invalid axis " << axis << " for array with " - << a.ndim() << " dimensions."; - throw std::invalid_argument(msg.str()); - } - int axis_ = axis < 0 ? axis + a.ndim() : axis; + int axis_ = normalize_axis_index(axis, a.ndim(), "[partition] "); int kth_ = kth < 0 ? kth + a.shape(axis) : kth; if (kth_ < 0 || kth_ >= a.shape(axis_)) { std::ostringstream msg; @@ -2874,14 +2865,7 @@ int kth, int axis, StreamOrDevice s /* = {} */) { // Check for valid axis - if (axis + static_cast<int>(a.ndim()) < 0 || - axis >= static_cast<int>(a.ndim())) { - std::ostringstream msg; - msg << "[argpartition] Received invalid axis " << axis << " for array with " - << a.ndim() << " dimensions."; - throw std::invalid_argument(msg.str()); - } - int axis_ = axis < 0 ? axis + a.ndim() : axis; + int axis_ = normalize_axis_index(axis, a.ndim(), "[argpartition] "); int kth_ = kth < 0 ? kth + a.shape(axis) : kth; if (kth_ < 0 || kth_ >= a.shape(axis_)) { std::ostringstream msg; @@ -2936,13 +2920,7 @@ /** Returns topk elements of the array along a given axis. */ array topk(const array& a, int k, int axis, StreamOrDevice s /* = {}*/) { // Check for valid axis - int axis_ = axis < 0 ? axis + a.ndim() : axis; - if (axis_ < 0 || axis_ >= static_cast<int>(a.ndim())) { - std::ostringstream msg; - msg << "[topk] Received invalid axis " << axis << " for array with " - << a.ndim() << " dimensions."; - throw std::invalid_argument(msg.str()); - } + int axis_ = normalize_axis_index(axis, a.ndim(), "[topk] "); if (k < 0 || k > a.shape(axis_)) { std::ostringstream msg; msg << "[topk] Received invalid k=" << k << " along axis " << axis @@ -3178,6 +3156,9 @@ }   array remainder(const array& a, const array& b, StreamOrDevice s /* = {} */) { auto dtype = promote_types(a.dtype(), b.dtype()); + if (issubdtype(dtype, complexfloating)) { + throw std::invalid_argument("[remainder] Complex type not supported."); + } auto inputs = broadcast_arrays( {astype(a, dtype, s), astype(b, dtype, to_stream(s))}, s); auto shape = inputs[0].shape(); @@ -3268,6 +3249,9 @@ return array(a.shape(), dtype, std::make_shared<Exp>(to_stream(s)), {input}); }   array expm1(const array& a, StreamOrDevice s /* = {} */) { + if (a.dtype() == complex64) { + throw std::invalid_argument("[expm1] Not supported for complex64."); + } auto dtype = at_least_float(a.dtype()); auto input = astype(a, dtype, s); return array( @@ -3314,6 +3298,9 @@ a.shape(), dtype, std::make_shared<ArcTan>(to_stream(s)), {input}); }   array arctan2(const array& a, const array& b, StreamOrDevice s /* = {} */) { + if (a.dtype() == complex64 || b.dtype() == complex64) { + throw std::invalid_argument("[arctan2] Not supported for complex64."); + } auto dtype = at_least_float(promote_types(a.dtype(), b.dtype())); auto inputs = broadcast_arrays({astype(a, dtype, s), astype(b, dtype, s)}, s); auto shape = inputs[0].shape(); @@ -3421,6 +3408,9 @@ std::move(inputs)); }   array sigmoid(const array& a, StreamOrDevice s /* = {} */) { + if (a.dtype() == complex64) { + throw std::invalid_argument("[sigmoid] Not supported for complex64."); + } auto dtype = at_least_float(a.dtype()); auto input = astype(a, dtype, s); return array( @@ -3428,6 +3418,9 @@ a.shape(), dtype, std::make_shared<Sigmoid>(to_stream(s)), {input}); }   array erf(const array& a, StreamOrDevice s /* = {} */) { + if (a.dtype() == complex64) { + throw std::invalid_argument("[erf] Not supported for complex64."); + } auto dtype = at_least_float(a.dtype()); return array( a.shape(), @@ -3437,6 +3430,9 @@ {astype(a, dtype, s)}); }   array erfinv(const array& a, StreamOrDevice s /* = {} */) { + if (a.dtype() == complex64) { + throw std::invalid_argument("[erfinv] Not supported for complex64."); + } auto dtype = at_least_float(a.dtype()); return array( a.shape(), @@ -3648,11 +3644,13 @@ Shape out_shape(ndim, 1);   for (int i = ndim - 1, j = a.ndim() - 1; j >= 0; j--, i--) { a_shape[2 * i] = a.shape(j); - out_shape[i] *= a.shape(j); + out_shape[i] = + safe_cast(static_cast<int64_t>(out_shape[i]) * a.shape(j), "kron"); } for (int i = ndim - 1, j = b.ndim() - 1; j >= 0; j--, i--) { b_shape[2 * i + 1] = b.shape(j); - out_shape[i] *= b.shape(j); + out_shape[i] = + safe_cast(static_cast<int64_t>(out_shape[i]) * b.shape(j), "kron"); }   return reshape( @@ -5150,11 +5148,17 @@ divide(wq, from_fp8(scales, w.dtype(), s), s), scale_encode, s); } else { // convert to e8m0 auto z = array(0, scales.dtype()); - scales = where( - equal(scales, z, s), - z, - astype(round(log2(scales, s), s), int32, s), + // Round the scale up so the block maximum stays representable, + // matching the CUDA backend. + auto exponent = astype(round(log2(scales, s), s), int32, s); + auto decoded = + power(array(2.0f, float32), astype(exponent, float32, s), s); + exponent = where( + less(decoded, astype(scales, float32, s), s), + add(exponent, array(1, int32), s), + exponent, s); + scales = where(equal(scales, z, s), z, exponent, s);   wq = divide(wq, power(array(2.0f, w.dtype()), scales, s), s); scales = astype(add(scales, array(127, int32), s), uint8, s);
diff --git ml-explore/mlx/mlx/ops.h Layr-Labs/mlx/mlx/ops.h index 01e0a9928620562c84de3e69ec7d8726d8e1c1d2..f597753b1e1e1de0db730bf84570ffcb9a637746 100644 --- ml-explore/mlx/mlx/ops.h +++ Layr-Labs/mlx/mlx/ops.h @@ -38,13 +38,25 @@ MLX_API array arange(int start, int stop, int step, StreamOrDevice s = {}); MLX_API array arange(int start, int stop, StreamOrDevice s = {}); MLX_API array arange(int stop, StreamOrDevice s = {});   -/** A 1D array of `num` evenly spaced numbers in the range `[start, stop]` */ +/** + * A 1D array of `num` evenly spaced numbers in the range `[start, stop]`, or + * in the half-open range `[start, stop)` when `endpoint` is false. + */ MLX_API array linspace( double start, double stop, - int num = 50, + int num, + bool endpoint, Dtype dtype = float32, StreamOrDevice s = {}); +inline array linspace( + double start, + double stop, + int num = 50, + Dtype dtype = float32, + StreamOrDevice s = {}) { + return linspace(start, stop, num, true, dtype, s); +}   /** Convert an array to the given data type. */ MLX_API array astype(array a, Dtype dtype, StreamOrDevice s = {});
diff --git ml-explore/mlx/mlx/primitives.cpp Layr-Labs/mlx/mlx/primitives.cpp index 3bafd194071f5fe05a2d9724996a9f4aee13c19f..9a1771394efaa0e1b505ce4c37da38d68facdd72 100644 --- ml-explore/mlx/mlx/primitives.cpp +++ Layr-Labs/mlx/mlx/primitives.cpp @@ -616,7 +616,7 @@ assert(inputs.size() == 1); assert(axes.size() == 1);   int axis_left = axes[0] >= 0 && axes[0] <= axis_; - return {{argpartition(inputs[0], axis_ + axis_left, stream())}, axes}; + return {{argpartition(inputs[0], kth_, axis_ + axis_left, stream())}, axes}; }   std::vector<array> ArgPartition::vjp( @@ -1187,9 +1187,11 @@ std::vector<Shape> Concatenate::output_shapes( const std::vector<array>& inputs) { auto shape = inputs[0].shape(); + int64_t concat_size = shape[axis_]; for (int i = 1; i < inputs.size(); ++i) { - shape[axis_] += inputs[i].shape(axis_); + concat_size += inputs[i].shape(axis_); } + shape[axis_] = safe_cast(concat_size, "concatenate"); return {std::move(shape)}; }   @@ -1939,6 +1941,11 @@ auto [a, b, to_ax] = vmap_binary_op(inputs, axes, stream()); return {{equal(a, b, stream())}, {to_ax}}; }   +bool Equal::is_equivalent(const Primitive& other) const { + const Equal& e_other = static_cast<const Equal&>(other); + return equal_nan_ == e_other.equal_nan_; +} + std::vector<array> Equal::vjp( const std::vector<array>& primals, const std::vector<array>& cotangents, @@ -2251,12 +2258,13 @@ for (auto& fft_ax : fft_axes) { if (fft_ax >= ax) { fft_ax++; } - if (real_) { - auto n = out_shape[fft_ax]; - out_shape[fft_ax] = inverse_ ? 2 * (n - 1) : n / 2 + 1; - } } } + // Only the last transformed axis changes size in a real transform + if (real_) { + auto n = out_shape[fft_axes.back()]; + out_shape[fft_axes.back()] = inverse_ ? 2 * (n - 1) : n / 2 + 1; + } return { {array( out_shape, @@ -2360,14 +2368,15 @@ const std::vector<int>& argnums) { assert(primals.size() == 1); assert(argnums.size() == 1); auto& tan = tangents[0]; + std::vector<int> axes(axes_.begin(), axes_.end()); if (real_ & inverse_) { - return {fft::irfftn(tan, fft::FFTNorm::Backward, stream())}; + return {fft::irfftn(tan, axes, fft::FFTNorm::Backward, stream())}; } else if (real_) { - return {fft::rfftn(tan, fft::FFTNorm::Backward, stream())}; + return {fft::rfftn(tan, axes, fft::FFTNorm::Backward, stream())}; } else if (inverse_) { - return {fft::ifftn(tan, fft::FFTNorm::Backward, stream())}; + return {fft::ifftn(tan, axes, fft::FFTNorm::Backward, stream())}; } else { - return {fft::fftn(tan, fft::FFTNorm::Backward, stream())}; + return {fft::fftn(tan, axes, fft::FFTNorm::Backward, stream())}; } }   @@ -2789,6 +2798,11 @@ in.dtype(), std::make_shared<Log>(stream(), base_), {in})}, axes}; +} + +bool Log::is_equivalent(const Primitive& other) const { + const Log& l_other = static_cast<const Log&>(other); + return base_ == l_other.base_; }   std::vector<array> Log1p::vjp( @@ -3421,7 +3435,7 @@ assert(inputs.size() == 1); assert(axes.size() == 1);   int axis_left = axes[0] >= 0 && axes[0] <= axis_; - return {{partition(inputs[0], axis_ + axis_left, stream())}, axes}; + return {{partition(inputs[0], kth_, axis_ + axis_left, stream())}, axes}; }   bool Partition::is_equivalent(const Primitive& other) const {
diff --git ml-explore/mlx/mlx/primitives.h Layr-Labs/mlx/mlx/primitives.h index 3a3d0ba5e5aee459b11fd1e66a7c4be46d02fdd2..0cfc71bf04f17a4aded04937b2d6ccd972f502c0 100644 --- ml-explore/mlx/mlx/primitives.h +++ Layr-Labs/mlx/mlx/primitives.h @@ -975,9 +975,9 @@ void eval_gpu(const std::vector<array>& inputs, array& out) override;   DEFINE_VMAP() DEFINE_GRADS() - DEFINE_DEFAULT_IS_EQUIVALENT() DEFINE_INPUT_OUTPUT_SHAPE()   + bool is_equivalent(const Primitive& other) const override; const char* name() const override { if (equal_nan_) { return "NaNEqual"; @@ -1325,9 +1325,9 @@ void eval_gpu(const std::vector<array>& inputs, array& out) override;   DEFINE_VMAP() DEFINE_GRADS() - DEFINE_DEFAULT_IS_EQUIVALENT() DEFINE_INPUT_OUTPUT_SHAPE()   + bool is_equivalent(const Primitive& other) const override; Base state() const { return base_; };
diff --git ml-explore/mlx/mlx/scheduler.cpp Layr-Labs/mlx/mlx/scheduler.cpp index 7507917f5bbaff870fee71c33bc981883fb13d6b..572e236bbeb0298e317f3761cee1bc47bf08b9f5 100644 --- ml-explore/mlx/mlx/scheduler.cpp +++ Layr-Labs/mlx/mlx/scheduler.cpp @@ -1,8 +1,12 @@ // Copyright © 2023-2026 Apple Inc.   -#include "mlx/scheduler.h" +#include <future> +#include <thread> + #include "mlx/backend/cpu/eval.h" #include "mlx/backend/gpu/eval.h" +#include "mlx/compile_impl.h" +#include "mlx/scheduler.h" #include "mlx/utils.h"   namespace mlx::core { @@ -13,6 +17,7 @@ auto p = std::make_shared<std::promise<void>>(); std::future<void> f = p->get_future(); scheduler::enqueue(s, [p = std::move(p)]() { p->set_value(); }); f.wait(); + scheduler::check_error(s); } else { gpu::synchronize(s); } @@ -27,12 +32,65 @@ synchronize(default_stream(default_device())); }   void clear_streams() { + detail::compile_clear_cache(detail::compile_cache()); cpu::clear_streams(); gpu::clear_streams(); }   namespace scheduler {   +struct StreamThread { + std::mutex mtx; + std::queue<std::function<void()>> q; + std::condition_variable cond; + bool stop; + std::thread thread; + Error error; + + StreamThread() : stop(false), thread(&StreamThread::thread_fn, this) {} + + ~StreamThread() { + { + std::lock_guard<std::mutex> lk(mtx); + stop = true; + } + cond.notify_one(); + thread.join(); + } + + void thread_fn() { + while (true) { + std::function<void()> task; + { + std::unique_lock<std::mutex> lk(mtx); + cond.wait(lk, [this] { return !this->q.empty() || this->stop; }); + if (q.empty() && stop) { + return; + } + task = std::move(q.front()); + q.pop(); + } + + task(); + } + } + + void enqueue(std::function<void()> f) { + if (is_main_thread()) { + error.check(); + } + { + std::lock_guard<std::mutex> lk(mtx); + if (stop) { + throw std::runtime_error( + "Cannot enqueue work after stream is stopped."); + } + q.emplace(std::move(f)); + } + cond.notify_one(); + } +}; + Scheduler::Scheduler() { is_main_thread(); gpu::init(); @@ -41,23 +99,66 @@ Scheduler::~Scheduler() = default;   void Scheduler::enqueue(Stream s, std::function<void()> task) { - StreamThread* st = nullptr; + auto& st = get_thread(s); + st.enqueue([&st, task = std::move(task)]() mutable { + try { + task(); + } catch (const std::exception& error) { + // Set error to stream only when no error happended before, to preserve + // the earliest error. + if (!st.error.valid()) { + st.error.set_message(std::make_shared<std::string>(error.what())); + } + } + }); +} + +void Scheduler::wait_event( + Stream s, + Event event, + std::function<void(Event&)> task) { + assert(s.device == Device::cpu); + auto& st = get_thread(s); + st.enqueue([&st, event = std::move(event), task = std::move(task)]() mutable { + task(event); + // Poison current stream if the waited event has error. + st.error.store_if_valid(event.load_error()); + }); +} + +void Scheduler::signal_event( + Stream s, + Event event, + std::function<void(Event&)> task) { + assert(s.device == Device::cpu); + auto& st = get_thread(s); + st.enqueue([&st, event = std::move(event), task = std::move(task)]() mutable { + // Poison the signal event if current stream has error. + if (st.error.valid()) { + event.set_error(st.error); + } + task(event); + }); +} + +void Scheduler::check_error(Stream s) { + get_thread(s).error.check(); +} + +StreamThread& Scheduler::get_thread(Stream s) { { std::shared_lock lock(threads_mtx_); auto it = threads_.find(s.index); if (it != threads_.end()) { - st = it->second.get(); + return *it->second.get(); } } - if (!st) { - std::unique_lock lock(threads_mtx_); - auto it = threads_.find(s.index); - if (it == threads_.end()) { - it = threads_.emplace(s.index, std::make_unique<StreamThread>()).first; - } - st = it->second.get(); + std::unique_lock lock(threads_mtx_); + auto it = threads_.find(s.index); + if (it == threads_.end()) { + it = threads_.emplace(s.index, std::make_unique<StreamThread>()).first; } - st->enqueue(std::move(task)); + return *it->second.get(); }   // Leak the scheduler singleton on all platforms. During static destruction,
diff --git ml-explore/mlx/mlx/scheduler.h Layr-Labs/mlx/mlx/scheduler.h index c84ab62855bee0b74c1182fc1cca8d0e76ecef9f..7bce5a9fe68a0e99821cfb1a60932f6d86de64c0 100644 --- ml-explore/mlx/mlx/scheduler.h +++ Layr-Labs/mlx/mlx/scheduler.h @@ -3,66 +3,19 @@ #pragma once   #include <atomic> -#include <future> #include <queue> #include <shared_mutex> -#include <thread> #include <unordered_map>   #include "mlx/api.h" #include "mlx/backend/gpu/eval.h" #include "mlx/device.h" #include "mlx/stream.h" +#include "mlx/utils.h"   namespace mlx::core::scheduler {   -struct StreamThread { - std::mutex mtx; - std::queue<std::function<void()>> q; - std::condition_variable cond; - bool stop; - std::thread thread; - - StreamThread() : stop(false), thread(&StreamThread::thread_fn, this) {} - - ~StreamThread() { - { - std::lock_guard<std::mutex> lk(mtx); - stop = true; - } - cond.notify_one(); - thread.join(); - } - - void thread_fn() { - while (true) { - std::function<void()> task; - { - std::unique_lock<std::mutex> lk(mtx); - cond.wait(lk, [this] { return !this->q.empty() || this->stop; }); - if (q.empty() && stop) { - return; - } - task = std::move(q.front()); - q.pop(); - } - - task(); - } - } - - void enqueue(std::function<void()> f) { - { - std::lock_guard<std::mutex> lk(mtx); - if (stop) { - throw std::runtime_error( - "Cannot enqueue work after stream is stopped."); - } - q.emplace(std::move(f)); - } - cond.notify_one(); - } -}; +class StreamThread;   class MLX_API Scheduler { public: @@ -76,6 +29,9 @@ Scheduler& operator=(const Scheduler&) = delete; Scheduler& operator=(Scheduler&&) = delete;   void enqueue(Stream s, std::function<void()> task); + void wait_event(Stream s, Event event, std::function<void(Event&)> task); + void signal_event(Stream s, Event event, std::function<void(Event&)> task); + void check_error(Stream s);   void notify_new_task(const Stream& stream) { { @@ -110,6 +66,8 @@ private: friend Stream mlx::core::new_stream(Device d);   + StreamThread& get_thread(Stream s); + int n_active_tasks_{0}; std::unordered_map<int, std::unique_ptr<StreamThread>> threads_; std::shared_mutex threads_mtx_; @@ -120,8 +78,24 @@ MLX_API Scheduler& scheduler();   template <typename F> -void enqueue(const Stream& stream, F&& f) { - scheduler().enqueue(stream, std::forward<F>(f)); +inline void enqueue(Stream s, F&& f) { + scheduler().enqueue(s, std::forward<F>(f)); +} + +// Like enqueue but the task is used for processing the passed event. +template <typename F> +inline void wait_event(Stream s, Event event, F&& f) { + scheduler().wait_event(s, std::move(event), std::forward<F>(f)); +} + +template <typename F> +inline void signal_event(Stream s, Event event, F&& f) { + scheduler().signal_event(s, std::move(event), std::forward<F>(f)); +} + +// Throw and clear the error stored in the stream, if any. +inline void check_error(Stream s) { + scheduler().check_error(s); }   inline int n_active_tasks() {
diff --git ml-explore/mlx/mlx/version.h Layr-Labs/mlx/mlx/version.h index dc8b3a86303d4038bda137dc391b4638d01fafc1..9de9f2e55838f7f996808a9de0f7a7c892bef9e6 100644 --- ml-explore/mlx/mlx/version.h +++ Layr-Labs/mlx/mlx/version.h @@ -6,7 +6,7 @@ #include "mlx/api.h"   #define MLX_VERSION_MAJOR 0 #define MLX_VERSION_MINOR 32 -#define MLX_VERSION_PATCH 1 +#define MLX_VERSION_PATCH 2 #define MLX_VERSION_NUMERIC \ (100000 * MLX_VERSION_MAJOR + 1000 * MLX_VERSION_MINOR + MLX_VERSION_PATCH)
diff --git ml-explore/mlx/python/mlx/__array_api_info.py Layr-Labs/mlx/python/mlx/__array_api_info.py new file mode 100644 index 0000000000000000000000000000000000000000..847a0bcbf3e90d52b4fcd5ec842d359a23ce24ec --- /dev/null +++ Layr-Labs/mlx/python/mlx/__array_api_info.py @@ -0,0 +1,82 @@ +class ArrayNamespaceInfo: + def capabilities(self): + return { + "boolean indexing": False, + "data-dependent shapes": False, + "max dimensions": 10, + } + + def default_device(self): + import mlx.core as mx + + return mx.default_device() + + def default_dtypes(self, *, device=None): + import mlx.core as mx + + if device is not None and not isinstance(device, mx.Device): + raise TypeError("Expected a mlx Device") + return { + "real floating": mx.float32, + "complex floating": mx.complex64, + "integral": mx.int32, + "indexing": mx.int32, + } + + def devices(self): + import mlx.core as mx + + devices = [ + mx.Device(dev_type, i) + for dev_type in (mx.cpu, mx.gpu) + for i in range(mx.device_count(dev_type)) + ] + return tuple(devices) + + def dtypes(self, *, device=None, kind=None): + import mlx.core as mx + + if device is not None and not isinstance(device, mx.Device): + raise TypeError("Expected a mlx Device") + device = device if device is not None else self.default_device() + + dtypes = { + "bool": mx.bool_, + "int8": mx.int8, + "int16": mx.int16, + "int32": mx.int32, + "int64": mx.int64, + "uint8": mx.uint8, + "uint16": mx.uint16, + "uint32": mx.uint32, + "uint64": mx.uint64, + "float32": mx.float32, + "complex64": mx.complex64, + } + if device.type == mx.cpu: + dtypes["float64"] = mx.float64 + if kind is None: + return dtypes + + signed = {"int8", "int16", "int32", "int64"} + unsigned = {"uint8", "uint16", "uint32", "uint64"} + real = {"float32", "float64"} + complex_ = {"complex64"} + kinds = { + "bool": {"bool"}, + "signed integer": signed, + "unsigned integer": unsigned, + "integral": signed | unsigned, + "real floating": real, + "complex floating": complex_, + "numeric": signed | unsigned | real | complex_, + } + kind = (kind,) if isinstance(kind, str) else kind + if not isinstance(kind, tuple) or any(k not in kinds for k in kind): + raise ValueError(f"Unsupported dtype kind: {kind!r}") + names = {name for k in kind for name in kinds[k]} + return {name: dtype for name, dtype in dtypes.items() if name in names} + + +def __array_namespace_info__(): + return ArrayNamespaceInfo()
diff --git ml-explore/mlx/python/mlx/_distributed_utils/launch.py Layr-Labs/mlx/python/mlx/_distributed_utils/launch.py index 4771c1fb5bb9a8c64ceff399dcbccfa200c80b93..0e661e358ffb13fd8b110145cc73c7e46554eac8 100644 --- ml-explore/mlx/python/mlx/_distributed_utils/launch.py +++ Layr-Labs/mlx/python/mlx/_distributed_utils/launch.py @@ -376,9 +376,31 @@ parser.error("Rank 0 should have an IP reachable from all other ranks")   jaccl_ring = args.backend == "jaccl-ring" have_rdmas = all(len(h.rdma) == len(hosts) for h in hosts) + if not have_rdmas: + parser.error( + "The hostfile is malformed: number of RDMA devices does not match hosts" + ) have_nulls = all(h.rdma[i] is None for i, h in enumerate(hosts)) - if not have_rdmas or not have_nulls: - parser.error("Malformed hostfile for jaccl backend") + if not have_nulls: + parser.error("The hostfile is malformed: RDMA device of self should be null") + + # Find pairs that miss rmda in hostfile. + n = len(hosts) + missing_rdma = [ + (i, j) + for i, h in enumerate(hosts) + for j in (((i - 1) % n, (i + 1) % n) if jaccl_ring else range(n)) + if i != j and h.rdma[j] is None + ] + + if missing_rdma: + pairs = ", ".join( + f"{hosts[i].ssh_hostname} to {hosts[j].ssh_hostname}" + for i, j in missing_rdma[:3] + ) + if len(missing_rdma) > 3: + pairs += f" and {len(missing_rdma) - 3} more" + parser.error(f"The hostfile is malformed: no RDMA device is listed for {pairs}")   coordinator = hosts[0].ips[0] env = args.env
diff --git ml-explore/mlx/python/mlx/_stub_patterns.txt Layr-Labs/mlx/python/mlx/_stub_patterns.txt index 974ce0c7a5eede1957a9b9bb775f5858a36c612a..90afb55981e0cd6448a7365bae8c83f94a48d087 100644 --- ml-explore/mlx/python/mlx/_stub_patterns.txt +++ Layr-Labs/mlx/python/mlx/_stub_patterns.txt @@ -1,10 +1,10 @@ mlx.core.__prefix__: - from typing import Any, ParamSpec, Protocol, TypeAlias, TypeVar + from typing import Any, BinaryIO as file, Literal, ParamSpec, Protocol, TypeAlias, TypeVar P = ParamSpec("P") R = TypeVar("R") class DLPackCompatible(Protocol): - __dlpack__: Callable[..., Any] - __dlpack_device__: Callable[..., Any] + def __dlpack__(self, *args: Any, **kwargs: Any) -> Any: ... + def __dlpack_device__(self, *args: Any, **kwargs: Any) -> Any: ...   mlx.core.__suffix__: scalar: TypeAlias = int | float | bool | complex @@ -12,21 +12,32 @@ list_or_scalar: TypeAlias = scalar | list["list_or_scalar"] StreamOrDevice: TypeAlias = Stream | ThreadLocalStream | Device | DeviceType | None bool_: Dtype = ...   +mlx.core.matrix_norm: + matrix_norm = linalg.norm + +mlx.core.array.__(eq|ne)__: + @overload + def __\1__(self, other: bool | int | float | array | Annotated[NDArray, dict(writable=False)] | complex) -> array: ... + @overload + def __\1__(self, other: ArrayLike) -> array | bool: ... + @overload + def __\1__(self, other: object) -> Any: ... + +mlx.core._PrintOptionsContext: + class _PrintOptionsContext: + def __init__(self, arg: PrintOptions, /) -> None: ... + def __enter__(self) -> _PrintOptionsContext: ... + def __exit__(self, *args) -> None: ... + mlx.core.distributed.__prefix__: - from mlx.core import array, Dtype, StreamOrDevice, scalar - from mlx.core.distributed import Group - from collections.abc import Sequence + from mlx.core import array, Dtype, StreamOrDevice + from collections.abc import Callable, Sequence   mlx.core.fast.__prefix__: - from mlx.core import array, Dtype, StreamOrDevice, scalar + from mlx.core import array, StreamOrDevice   mlx.core.linalg.__prefix__: - from mlx.core import array, Dtype, StreamOrDevice, scalar - from collections.abc import Sequence - -mlx.core.metal.__prefix__: - from mlx.core import array, Dtype, Device, Stream, scalar - from collections.abc import Sequence + from mlx.core import array, StreamOrDevice   mlx.core.random.__prefix__: from mlx.core import array, Dtype, StreamOrDevice, scalar, float32, int32
diff --git ml-explore/mlx/python/mlx/nn/layers/normalization.py Layr-Labs/mlx/python/mlx/nn/layers/normalization.py index e79440dce3fed8690d987449d038baa3532f394b..97f6942f04d00e0b240b379f93efd75d6a28dd28 100644 --- ml-explore/mlx/python/mlx/nn/layers/normalization.py +++ Layr-Labs/mlx/python/mlx/nn/layers/normalization.py @@ -47,6 +47,8 @@ eps: float = 1e-5, affine: bool = False, ): super().__init__() + if eps <= 0.0: + raise ValueError(f"[InstanceNorm] 'eps' must be positive but got {eps}.") if affine: self.weight = mx.ones((dims,)) self.bias = mx.zeros((dims,)) @@ -62,12 +64,19 @@ raise ValueError( f"InstanceNorm expects inputs with at least 3 dimensions" f" (N, ..., C) but the input has {x.ndim} dimensions." ) - reduction_axes = tuple(range(1, x.ndim - 1)) - # Compute stats - mean = mx.mean(x, axis=reduction_axes, keepdims=True) - var = mx.var(x, axis=reduction_axes, keepdims=True) - # Normalize - x = (x - mean) * mx.rsqrt(var + self.eps) + batch_size, features = x.shape[0], x.shape[-1] + spatial_shape = x.shape[1:-1] + channels_first = mx.transpose(x, (0, x.ndim - 1, *range(1, x.ndim - 1))) + x = mx.fast.layer_norm( + channels_first.reshape(batch_size, features, -1), + None, + None, + self.eps, + ) + x = mx.transpose( + x.reshape(batch_size, features, *spatial_shape), + (0, *range(2, len(spatial_shape) + 2), 1), + ) # Scale and shift if necessary return (self.weight * x + self.bias) if "weight" in self else x   @@ -101,6 +110,8 @@ def __init__( self, dims: int, eps: float = 1e-5, affine: bool = True, bias: bool = True ): super().__init__() + if eps <= 0.0: + raise ValueError(f"[LayerNorm] 'eps' must be positive but got {eps}.") if affine: self.weight = mx.ones((dims,)) if bias: @@ -141,6 +152,8 @@ """   def __init__(self, dims: int, eps: float = 1e-5): super().__init__() + if eps <= 0.0: + raise ValueError(f"[RMSNorm] 'eps' must be positive but got {eps}.") self.weight = mx.ones((dims,)) self.eps = eps   @@ -191,6 +204,8 @@ affine: bool = True, pytorch_compatible: bool = False, ): super().__init__() + if eps <= 0.0: + raise ValueError(f"[GroupNorm] 'eps' must be positive but got {eps}.") if num_groups <= 0: raise ValueError( f"The number of groups ({num_groups}) must be a positive integer." @@ -309,6 +324,8 @@ affine: bool = True, track_running_stats: bool = True, ): super().__init__() + if eps <= 0.0: + raise ValueError(f"[BatchNorm] 'eps' must be positive but got {eps}.")   self.num_features = num_features self.eps = eps
diff --git ml-explore/mlx/python/mlx/nn/losses.py Layr-Labs/mlx/python/mlx/nn/losses.py index 184df2a2e09aa8a7a6b09d075ecf6277c03b0b3f..b98d2765d69c6f5bbc80e235f852a0889e32b023 100644 --- ml-explore/mlx/python/mlx/nn/losses.py +++ Layr-Labs/mlx/python/mlx/nn/losses.py @@ -63,6 +63,18 @@ >>> logits = mx.array([[2.0, -1.0], [-1.0, 2.0]]) >>> targets = mx.array([[0.9, 0.1], [0.1, 0.9]]) >>> nn.losses.cross_entropy(logits, targets) array([0.348587, 0.348587], dtype=float32) + >>> + >>> # Half precision logits with class indices as targets. On CUDA a + >>> # fused kernel accumulates the reduction in float32: + >>> logits = mx.array([[2.0, -1.0], [-1.0, 2.0]], mx.bfloat16) + >>> targets = mx.array([0, 1]) + >>> nn.losses.cross_entropy(logits, targets) + array([0.0485873, 0.0485873], dtype=float32) + >>> + >>> # Metal and the CPU reduce in the dtype of the logits, so upcast + >>> # them to get the same accuracy: + >>> nn.losses.cross_entropy(logits.astype(mx.float32), targets) + array([0.0485873, 0.0485873], dtype=float32) """ if label_smoothing < 0 or label_smoothing >= 1: raise ValueError(f"Label smoothing must be in [0, 1), got {label_smoothing}.") @@ -83,31 +95,38 @@ raise ValueError( f"Targets shape {targets.shape} does not match logits shape {logits.shape}." )   - # Shift by the max first. The loss only depends on differences between - # logits, but subtracting the logsumexp of large logits loses the gap to - # rounding before the subtraction happens. - logits = logits - mx.stop_gradient(mx.max(logits, axis=axis, keepdims=True)) + use_fast = ( + mx.cuda.is_available() + and mx.default_device() == mx.gpu + and not targets_as_probs + and label_smoothing == 0 + and axis in (-1, logits.ndim - 1) + and mx.issubdtype(logits.dtype, mx.floating) + and mx.issubdtype(targets.dtype, mx.integer) + )   - if targets_as_probs: - score = mx.sum(logits * targets, axis=axis) + if use_fast: + loss = mx.fast.cross_entropy(logits, targets).astype(logits.dtype) else: - score = mx.take_along_axis(logits, mx.expand_dims(targets, axis), axis).squeeze( - axis - ) + logits = logits - mx.stop_gradient(mx.max(logits, axis=axis, keepdims=True)) + + if targets_as_probs: + score = mx.sum(logits * targets, axis=axis) + else: + score = mx.take_along_axis( + logits, mx.expand_dims(targets, axis), axis + ).squeeze(axis)   - logsumexp_logits = mx.logsumexp(logits, axis=axis) - if label_smoothing > 0: - # Adjust the true class score with label smoothing - adjusted_score = (1 - label_smoothing) * score + logsumexp_logits = mx.logsumexp(logits, axis=axis) + if label_smoothing > 0: + adjusted_score = (1 - label_smoothing) * score   - # Calculate the mean logit across the classes for smoothed loss - mean_logits = logits.mean(axis=axis) - smoothed_loss = -mean_logits * label_smoothing + mean_logits = logits.mean(axis=axis) + smoothed_loss = -mean_logits * label_smoothing   - # Combine the adjusted score and smoothed loss with the logsumexp logits - loss = logsumexp_logits - adjusted_score + smoothed_loss - else: - loss = logsumexp_logits - score + loss = logsumexp_logits - adjusted_score + smoothed_loss + else: + loss = logsumexp_logits - score   # Apply weights if provided if weights is not None:
diff --git ml-explore/mlx/python/mlx/optimizers/optimizers.py Layr-Labs/mlx/python/mlx/optimizers/optimizers.py index 65efab222da926394801e38a40d76fbf96bbebea..8be344247cc928e3bf8b673b5d0bd32754edb06a 100644 --- ml-explore/mlx/python/mlx/optimizers/optimizers.py +++ Layr-Labs/mlx/python/mlx/optimizers/optimizers.py @@ -499,6 +499,16 @@ bias_correction: bool = False, ): super().__init__()   + for i, beta in enumerate(betas): + if not 0.0 <= beta < 1.0: + raise ValueError( + f"Adam beta{i + 1} should be in [0, 1), {beta} was provided " + "instead" + ) + + if not 0.0 <= eps: + raise ValueError(f"Adam epsilon should be >=0, {eps} was provided instead") + self._maybe_schedule("learning_rate", learning_rate) self.betas = betas self.eps = eps @@ -620,10 +630,6 @@ betas: List[float] = [0.9, 0.999], eps: float = 1e-8, ): super().__init__(learning_rate, betas, eps) - if not 0.0 <= eps: - raise ValueError( - f"Epsilon value should be >=0, {self.eps} was provided instead" - )   def init_single(self, parameter: mx.array, state: dict): """Initialize optimizer state""" @@ -682,6 +688,13 @@ betas: List[float] = [0.9, 0.99], weight_decay: float = 0.0, ): super().__init__() + + for i, beta in enumerate(betas): + if not 0.0 <= beta < 1.0: + raise ValueError( + f"Lion beta{i + 1} should be in [0, 1), {beta} was provided " + "instead" + )   self._maybe_schedule("learning_rate", learning_rate) self.betas = betas
diff --git ml-explore/mlx/python/src/CMakeLists.txt Layr-Labs/mlx/python/src/CMakeLists.txt index 447271500b55bc2a8ffa51dc4e406edb9001c947..0798add4109523abc6a5b35a9654fe2d382c2943 100644 --- ml-explore/mlx/python/src/CMakeLists.txt +++ Layr-Labs/mlx/python/src/CMakeLists.txt @@ -2,6 +2,7 @@ nanobind_add_module( core NB_STATIC STABLE_ABI + FREE_THREADED LTO NOMINSIZE NB_DOMAIN
diff --git ml-explore/mlx/python/src/convert.cpp Layr-Labs/mlx/python/src/convert.cpp index 9941358e8475546077f582a890718365c133450a..a3da76ef535aedc91b54398bd02bfef704c5745a 100644 --- ml-explore/mlx/python/src/convert.cpp +++ Layr-Labs/mlx/python/src/convert.cpp @@ -505,7 +505,8 @@ PyScalarT validate_shape( T list, const mx::Shape& shape, int idx, - bool& all_python_primitive_elements) { + bool& all_python_primitive_elements, + bool& has_wide_int) { if (idx >= shape.size()) { throw std::invalid_argument("Initialization encountered extra dimension."); } @@ -524,13 +525,18 @@ for (auto l : list) { PyScalarT t; if (nb::isinstance<nb::list>(l)) { t = validate_shape( - nb::cast<nb::list>(l), shape, idx + 1, all_python_primitive_elements); + nb::cast<nb::list>(l), + shape, + idx + 1, + all_python_primitive_elements, + has_wide_int); } else if (nb::isinstance<nb::tuple>(*list.begin())) { t = validate_shape( nb::cast<nb::tuple>(l), shape, idx + 1, - all_python_primitive_elements); + all_python_primitive_elements, + has_wide_int); } else if (nb::isinstance<mx::array>(l)) { all_python_primitive_elements = false; auto arr = nb::cast<mx::array>(l); @@ -549,6 +555,13 @@ if (nb::isinstance<nb::bool_>(l)) { t = pybool; } else if (nb::isinstance<nb::int_>(l)) { t = pyint; + // Match the scalar path, which widens to int64 rather than failing + // when a python int does not fit in int32. + auto val = nb::cast<int64_t>(l); + if (val > std::numeric_limits<int>::max() || + val < std::numeric_limits<int>::min()) { + has_wide_int = true; + } } else if (nb::isinstance<nb::float_>(l)) { t = pyfloat; } else if (PyComplex_Check(l.ptr())) { @@ -594,7 +607,8 @@ mx::array array_from_list_impl( T pl, const PyScalarT& inferred_type, std::optional<mx::Dtype> specified_type, - const mx::Shape& shape) { + const mx::Shape& shape, + bool has_wide_int) { // Make the array switch (inferred_type) { case pybool: { @@ -603,7 +617,8 @@ fill_vector(pl, vals); return mx::array(vals.begin(), shape, specified_type.value_or(mx::bool_)); } case pyint: { - auto dtype = specified_type.value_or(mx::int32); + auto dtype = + specified_type.value_or(has_wide_int ? mx::int64 : mx::int32); if (dtype == mx::int64) { std::vector<int64_t> vals; fill_vector(pl, vals); @@ -663,11 +678,13 @@ get_shape(pl, shape);   // Validate the shape and type bool all_python_primitive_elements = true; - auto type = validate_shape(pl, shape, 0, all_python_primitive_elements); + bool has_wide_int = false; + auto type = + validate_shape(pl, shape, 0, all_python_primitive_elements, has_wide_int);   if (all_python_primitive_elements) { // `pl` does not contain mlx arrays - return array_from_list_impl(pl, type, dtype, shape); + return array_from_list_impl(pl, type, dtype, shape, has_wide_int); }   // `pl` contains mlx arrays
diff --git ml-explore/mlx/python/src/fast.cpp Layr-Labs/mlx/python/src/fast.cpp index e59357bc337c818a12d5e566a681f1f33a1cf167..67c3442cfff4c2317e899dcc52ad3bfbe38c8fb9 100644 --- ml-explore/mlx/python/src/fast.cpp +++ Layr-Labs/mlx/python/src/fast.cpp @@ -175,6 +175,36 @@ array: The output array. )pbdoc");   m.def( + "cross_entropy", + &mx::fast::cross_entropy, + "logits"_a, + "targets"_a, + nb::kw_only(), + "stream"_a = nb::none(), + nb::sig( + "def cross_entropy(logits: array, targets: array, *, stream: StreamOrDevice = None) -> array"), + R"pbdoc( + Cross entropy loss with class indices as targets. + + Computes ``logsumexp(logits, axis=-1) - logits[..., target]`` in a + fused kernel with accumulation in float32. + + Note: Currently is implemented only on CUDA, fallback to unfused version with + manual casting on Metal and CPU. + + Args: + logits (array): The unnormalized logits. The loss is computed over + the last axis. + targets (array): Class indices. The shape should match the shape of + ``logits`` with the last axis removed. The indices must be in + ``[0, logits.shape[-1])``. + + Returns: + array: The per-element loss in float32, with the shape of + ``targets``. + )pbdoc"); + + m.def( "rope", [](const mx::array& a, int dims, @@ -234,6 +264,7 @@ const mx::array& values, const float scale, const std::variant<std::monostate, std::string, mx::array>& mask, const std::optional<mx::array>& sinks, + bool force_fused, mx::StreamOrDevice s) { bool has_mask = !std::holds_alternative<std::monostate>(mask); bool has_str_mask = @@ -250,16 +281,32 @@ << mask_str << "'. Must be 'causal', or an array."; throw std::invalid_argument(msg.str()); } return mx::fast::scaled_dot_product_attention( - queries, keys, values, scale, mask_str, std::nullopt, sinks, s); + queries, + keys, + values, + scale, + mask_str, + std::nullopt, + sinks, + force_fused, + s); } else { auto mask_arr = std::get<mx::array>(mask); return mx::fast::scaled_dot_product_attention( - queries, keys, values, scale, "", mask_arr, sinks, s); + queries, + keys, + values, + scale, + "", + mask_arr, + sinks, + force_fused, + s); }   } else { return mx::fast::scaled_dot_product_attention( - queries, keys, values, scale, "", {}, sinks, s); + queries, keys, values, scale, "", {}, sinks, force_fused, s); } }, "q"_a, @@ -269,9 +316,10 @@ nb::kw_only(), "scale"_a, "mask"_a = nb::none(), "sinks"_a = nb::none(), + "force_fused"_a = false, "stream"_a = nb::none(), nb::sig( - "def scaled_dot_product_attention(q: array, k: array, v: array, *, scale: float, mask: None | str | array = None, sinks: array | None = None, stream: StreamOrDevice = None) -> array"), + "def scaled_dot_product_attention(q: array, k: array, v: array, *, scale: float, mask: None | str | array = None, sinks: array | None = None, force_fused: bool = False, stream: StreamOrDevice = None) -> array"), R"pbdoc( A fast implementation of multi-head attention: ``O = softmax(Q @ K.T, dim=-1) @ V``.   @@ -313,6 +361,11 @@ The ``"causal"`` mask uses lower-right alignment where the last query aligns with the last key. sinks (array, optional): An optional array of attention sinks. Default: ``None``. + force_fused (bool, optional): If ``True``, use a fused kernel + regardless of the builtin heuristics and raise error when no + fused kernel is available. For certain configurations this would + result in slower kernel getting used but can reduce memory + consumption. Default: ``False``.   Returns: array: The output array.
diff --git ml-explore/mlx/python/src/indexing.cpp Layr-Labs/mlx/python/src/indexing.cpp index 3df4c96882c7753403b3ea52367d751637ce2fc1..1dce5378808fba1a828808aa73ca6307b94c53b9 100644 --- ml-explore/mlx/python/src/indexing.cpp +++ Layr-Labs/mlx/python/src/indexing.cpp @@ -778,6 +778,8 @@ } else if (is_index_scalar(obj)) { return mlx_scatter_args_int(src, obj, vals); } else if (nb::isinstance<nb::tuple>(obj)) { return mlx_scatter_args_nd(src, nb::cast<nb::tuple>(obj), vals); + } else if (nb::isinstance<nb::ellipsis>(obj)) { + return {{}, broadcast_to(vals, src.shape()), {}}; } else if (obj.is_none()) { return {{}, broadcast_to(vals, src.shape()), {}}; } else if (nb::isinstance<nb::list>(obj)) {
diff --git ml-explore/mlx/python/src/mlx.cpp Layr-Labs/mlx/python/src/mlx.cpp index cb031cf78c143477583878d2b9c558370e6102ea..243449b385f3e361976ca988e4732e1a855cfc72 100644 --- ml-explore/mlx/python/src/mlx.cpp +++ Layr-Labs/mlx/python/src/mlx.cpp @@ -31,6 +31,10 @@ auto reprlib_fix = nb::module_::import_("mlx._reprlib_fix"); nb::set_leak_warnings(false);   + auto array_namespace_info = nb::module_::import_("mlx.__array_api_info"); + m.attr("__array_namespace_info__") = + array_namespace_info.attr("__array_namespace_info__"); + init_mlx_func(m); init_device(m); init_stream(m);
diff --git ml-explore/mlx/python/src/ops.cpp Layr-Labs/mlx/python/src/ops.cpp index d8892a45d264b70182a85bf664ced7566a7ecd19..677d998af13c88de171e696955d6fe0b954283eb 100644 --- ml-explore/mlx/python/src/ops.cpp +++ Layr-Labs/mlx/python/src/ops.cpp @@ -1,5 +1,6 @@ // Copyright © 2023-2024 Apple Inc.   +#include <limits> #include <numeric> #include <ostream> #include <variant> @@ -24,11 +25,14 @@ namespace mx = mlx::core; namespace nb = nanobind; using namespace nb::literals;   -using Scalar = std::variant<bool, int, double>; +using Scalar = std::variant<bool, int64_t, double>;   mx::Dtype scalar_to_dtype(Scalar s) { - if (std::holds_alternative<int>(s)) { - return mx::int32; + if (auto pv = std::get_if<int64_t>(&s); pv) { + return (*pv > std::numeric_limits<int>::max() || + *pv < std::numeric_limits<int>::min()) + ? mx::int64 + : mx::int32; } else if (std::holds_alternative<double>(s)) { return mx::float32; } else { @@ -37,7 +41,7 @@ } }   double scalar_to_double(Scalar s) { - if (auto pv = std::get_if<int>(&s); pv) { + if (auto pv = std::get_if<int64_t>(&s); pv) { return static_cast<double>(*pv); } else if (auto pv = std::get_if<double>(&s); pv) { return *pv; @@ -1513,7 +1517,7 @@ Args: start (float or int, optional): Starting value which defaults to ``0``. stop (float or int, optional): Stopping value. step (float or int, optional): Increment which defaults to ``1``. - dtype (Dtype, optional): Specifies the data type of the output. If unspecified will default to ``float32`` if any of ``start``, ``stop``, or ``step`` are ``float``. Otherwise will default to ``int32``. + dtype (Dtype, optional): Specifies the data type of the output. If unspecified will default to ``float32`` if any of ``start``, ``stop``, or ``step`` are ``float``. Otherwise will default to ``int32``, or ``int64`` if any of ``start``, ``stop``, or ``step`` does not fit in ``int32``.   Returns: array: The range of values. @@ -1644,22 +1648,25 @@ "linspace", [](Scalar start, Scalar stop, int num, + bool endpoint, std::optional<mx::Dtype> dtype, mx::StreamOrDevice s) { return mx::linspace( scalar_to_double(start), scalar_to_double(stop), num, + endpoint, dtype.value_or(mx::float32), s); }, "start"_a, "stop"_a, "num"_a = 50, + "endpoint"_a = true, "dtype"_a.none() = mx::float32, "stream"_a = nb::none(), nb::sig( - "def linspace(start: scalar, stop: scalar, num: int | None = 50, dtype: Dtype | None = float32, stream: StreamOrDevice = None) -> array"), + "def linspace(start: scalar, stop: scalar, num: int | None = 50, endpoint: bool = True, dtype: Dtype | None = float32, stream: StreamOrDevice = None) -> array"), R"pbdoc( Generate ``num`` evenly spaced numbers over interval ``[start, stop]``.   @@ -1667,6 +1674,9 @@ Args: start (scalar): Starting value. stop (scalar): Stopping value. num (int, optional): Number of samples, defaults to ``50``. + endpoint (bool, optional): If ``True``, ``stop`` is the last + sample. Otherwise it is not included and the samples are spaced + over the half-open interval ``[start, stop)``. Default: ``True``. dtype (Dtype, optional): Specifies the data type of the output, default to ``float32``.   @@ -1758,7 +1768,7 @@ } }, nb::arg(), "indices"_a, - "axis"_a.none(), + "axis"_a = nb::none(), nb::kw_only(), "stream"_a = nb::none(), nb::sig( @@ -3386,14 +3396,14 @@ std::string indexing, mx::StreamOrDevice s) { std::vector<mx::array> arrays = nb::cast<std::vector<mx::array>>(arrays_); - return mx::meshgrid(arrays, sparse, indexing, s); + return nb::tuple(nb::cast(mx::meshgrid(arrays, sparse, indexing, s))); }, "arrays"_a, "sparse"_a = false, "indexing"_a = "xy", "stream"_a = nb::none(), nb::sig( - "def meshgrid(*arrays: array, sparse: bool | None = False, indexing: str | None = 'xy', stream: StreamOrDevice = None) -> array"), + "def meshgrid(*arrays: array, sparse: bool | None = False, indexing: str | None = 'xy', stream: StreamOrDevice = None) -> tuple[array, ...]"), R"pbdoc( Generate multidimensional coordinate grids from 1-D coordinate arrays   @@ -3406,7 +3416,7 @@ indexing (str, optional): Cartesian ('xy') or matrix ('ij') indexing of the output arrays. Defaults to ``'xy'``.   Returns: - list(array): The output arrays. + tuple(array): The output arrays. )pbdoc"); m.def( "repeat",
diff --git ml-explore/mlx/python/src/random.cpp Layr-Labs/mlx/python/src/random.cpp index 8485faea41124883207bc11f9cc35b519bd77f42..10b82b8921a9770e73fb2298e8920c528090a494 100644 --- ml-explore/mlx/python/src/random.cpp +++ Layr-Labs/mlx/python/src/random.cpp @@ -64,6 +64,10 @@ static thread_local PyKeySequence ks; return ks; }   +void reset_random_state() { + default_key().reset(); +} + // A process-global sentinel for `mx.random.state`. Since it is the same object // on every thread, capturing it (e.g. with `mx.compile`) is thread-independent; // the pytree traversal in trees.cpp resolves it to the calling thread's key.
diff --git ml-explore/mlx/python/src/random.h Layr-Labs/mlx/python/src/random.h index 2baf9d92f1f615c92dc048c4699ce1a3e11288cc..02d81d4c7440226d3b3e15da897cdb6c2f3d6d02 100644 --- ml-explore/mlx/python/src/random.h +++ Layr-Labs/mlx/python/src/random.h @@ -9,6 +9,9 @@ namespace mx = mlx::core; namespace nb = nanobind;   +// Clear the `mx.random.state` python object in current thread. +void reset_random_state(); + // The process-global `mx.random.state` sentinel. nb::object random_state_sentinel();
diff --git ml-explore/mlx/python/src/stream.cpp Layr-Labs/mlx/python/src/stream.cpp index 004301a45c0ba0a81a424637da14472356235206..467518e991d035f682e3b2b2d20988b0616d7c36 100644 --- ml-explore/mlx/python/src/stream.cpp +++ Layr-Labs/mlx/python/src/stream.cpp @@ -9,6 +9,7 @@ #include <nanobind/stl/variant.h>   #include "mlx/stream.h" #include "mlx/utils.h" +#include "python/src/random.h"   namespace mx = mlx::core; namespace nb = nanobind; @@ -137,7 +138,10 @@ "def new_thread_local_stream(device: Device | DeviceType) -> ThreadLocalStream"), R"pbdoc(Make a new stream that will be unique per thread.)pbdoc"); m.def( "clear_streams", - &mx::clear_streams, + []() { + reset_random_state(); + mx::clear_streams(); + }, R"pbdoc(Destroy all streams created in current thread.)pbdoc");   nb::class_<PyStreamContext>(m, "StreamContext", R"pbdoc(
diff --git ml-explore/mlx/python/src/transforms.cpp Layr-Labs/mlx/python/src/transforms.cpp index 1d7aa8b9b15fcee5a9e0f79c979c11f104e2e602..e6a778eca0bdbe9622cd0bdec0c98c352dfc7b71 100644 --- ml-explore/mlx/python/src/transforms.cpp +++ Layr-Labs/mlx/python/src/transforms.cpp @@ -406,29 +406,13 @@ return tree_unflatten(py_outputs, outputs); }; }   -void ensure_compile_cache_cleanup() { - // Make sure each thread using mx.compile would clear its compile cache - // before python interpreter exits. - struct ThreadCleanup { - ~ThreadCleanup() { - if (!mx::detail::compile_cache_empty()) { - nb::gil_scoped_acquire gil; - mx::detail::compile_clear_cache(); - } - } - }; - static thread_local auto clear_cache = []() { - mx::detail::compile_clear_cache(); - return ThreadCleanup{}; - }(); -} - struct PyCompiledFun { nb::callable fun; std::uintptr_t fun_id; nb::object captured_inputs; nb::object captured_outputs; bool shapeless; + mx::detail::CompileCacheWeakPtr cache;   // Data to attach to the compiled function that contains the python output // structure and the number of arrays in said structure. @@ -438,6 +422,11 @@ int num_outputs;   AttachedData(nb::object output_structure_, int num_outputs_) : output_structure(output_structure_), num_outputs(num_outputs_) {} + + ~AttachedData() { + nb::gil_scoped_acquire gil; + output_structure.reset(); + } };   PyCompiledFun( @@ -456,15 +445,16 @@ PyCompiledFun& operator=(const PyCompiledFun&) = delete; PyCompiledFun& operator=(PyCompiledFun&& other) = delete; PyCompiledFun(PyCompiledFun&& other) : fun(std::move(other.fun)), - fun_id(reinterpret_cast<std::uintptr_t>(fun.ptr())) { + fun_id(reinterpret_cast<std::uintptr_t>(fun.ptr())), + captured_inputs(std::move(other.captured_inputs)), + captured_outputs(std::move(other.captured_outputs)), + shapeless(other.shapeless), + cache(other.cache) { other.fun_id = 0; - captured_inputs = std::move(other.captured_inputs); - captured_outputs = std::move(other.captured_outputs); - shapeless = other.shapeless; };   nb::object call_impl(const nb::args& args, const nb::kwargs& kwargs) { - ensure_compile_cache_cleanup(); + cache = mx::detail::compile_cache();   // Flat array inputs std::vector<mx::array> inputs; @@ -599,7 +589,7 @@ ~PyCompiledFun() { nb::gil_scoped_acquire gil;   - mx::detail::compile_erase(fun_id); + mx::detail::compile_erase(cache, fun_id); fun.reset(); captured_inputs.reset(); captured_outputs.reset(); @@ -1554,9 +1544,10 @@ A callable that recomputes intermediate states during gradient computation. )pbdoc");   - // Ensure the main thread cleanup will happen before the interpreter goes - // away. As a result if the other threads join the main thread we should have - // a clean tear-down. + // Clean up main thread compile cache before python interpreter shuts down. auto atexit = nb::module_::import_("atexit"); - atexit.attr("register")(nb::cpp_function(&mx::detail::compile_clear_cache)); + atexit.attr("register")( + nb::cpp_function([cache = mx::detail::compile_cache()]() { + mx::detail::compile_clear_cache(cache); + })); }
diff --git ml-explore/mlx/python/tests/__main__.py Layr-Labs/mlx/python/tests/__main__.py deleted file mode 100644 index 5230bd535428bf324012cc78272394e79168a9ff..0000000000000000000000000000000000000000 --- ml-explore/mlx/python/tests/__main__.py +++ /dev/null @@ -1,5 +0,0 @@ -from . import mlx_tests - -__unittest = True - -mlx_tests.MLXTestRunner(module=None)
diff --git ml-explore/mlx/python/tests/mlx_tests.py Layr-Labs/mlx/python/tests/mlx_tests.py index ac223f095085d626fc5fbb4838f7144bfaabd479..2b60f4615435f883d58a759f4d139d909f56a64f 100644 --- ml-explore/mlx/python/tests/mlx_tests.py +++ Layr-Labs/mlx/python/tests/mlx_tests.py @@ -1,13 +1,6 @@ # Copyright © 2023 Apple Inc.   import os - -# Use regular fp32 precision for tests -os.environ["MLX_ENABLE_TF32"] = "0" - -# Do not abort on cache thrashing -os.environ["MLX_ENABLE_CACHE_THRASHING_CHECK"] = "0" - import platform import sys import unittest
diff --git ml-explore/mlx/python/tests/run.py Layr-Labs/mlx/python/tests/run.py new file mode 100644 index 0000000000000000000000000000000000000000..df96d4dd36f2a0c2a165de4492985eb748ac4cfb --- /dev/null +++ Layr-Labs/mlx/python/tests/run.py @@ -0,0 +1,18 @@ +import os +import sys + +# Use regular fp32 precision for tests +os.environ["MLX_ENABLE_TF32"] = "0" + +# Do not abort on cache thrashing +os.environ["MLX_ENABLE_CACHE_THRASHING_CHECK"] = "0" + +__unittest = True + +import mlx_tests + +if __name__ == "__main__": + # Run all tests by default. + dirname = os.path.dirname(os.path.realpath(__file__)) + argv = [sys.argv[0], "discover", dirname, *sys.argv[1:]] + mlx_tests.MLXTestRunner(argv=argv, module=None)
diff --git ml-explore/mlx/python/tests/test_array.py Layr-Labs/mlx/python/tests/test_array.py index aee2ab88721c12396df53804081d5b53d44d5af2..80459fabe5cc2ee16413d4b474c294a4d2f90c3c 100644 --- ml-explore/mlx/python/tests/test_array.py +++ Layr-Labs/mlx/python/tests/test_array.py @@ -49,6 +49,37 @@ v = ".".join(str(int(vn)) for vn in vnums[:3]) self.assertEqual(v, mx.__version__[: len(v)])   +class TestArrayNamespsceInfo(mlx_tests.MLXTestCase): + def test(self): + namespace = mx.__array_namespace_info__() + + self.assertEqual(namespace.default_device(), mx.default_device()) + self.assertEqual( + namespace.default_dtypes(), + { + "real floating": mx.float32, + "complex floating": mx.complex64, + "integral": mx.int32, + "indexing": mx.int32, + }, + ) + self.assertEqual( + namespace.dtypes(device=mx.Device(mx.cpu), kind="real floating"), + {"float32": mx.float32, "float64": mx.float64}, + ) + if mx.is_available(mx.gpu): + self.assertEqual( + namespace.dtypes(device=mx.Device(mx.gpu), kind="real floating"), + {"float32": mx.float32}, + ) + self.assertEqual( + namespace.dtypes(kind=("bool", "complex floating")), + {"bool": mx.bool_, "complex64": mx.complex64}, + ) + with self.assertRaises(ValueError): + namespace.dtypes(kind="invalid") + + class TestDtypes(mlx_tests.MLXTestCase): def test_dtypes(self): self.assertEqual(mx.bool_.size, 1) @@ -522,6 +553,33 @@ self.assertEqual(out, x)   out = mx.array([x], dtype=mx.float64).item() self.assertEqual(out, x) + + def test_construction_from_lists_wide_ints(self): + # A python int that does not fit in int32 widens to int64, the same + # rule the scalar path already uses. It used to raise std::bad_cast. + for value in (2**31, 2**40, -(2**31) - 1, -(2**40)): + for make in ( + lambda v: [v], + lambda v: (v,), + lambda v: [[v]], + lambda v: [v, 1], + ): + x = mx.array(make(value)) + self.assertEqual(x.dtype, mx.int64, msg=f"{value} {make(value)}") + self.assertEqual(x.flatten()[0].item(), value) + self.assertEqual(mx.array(value).dtype, mx.int64) + + # Values that still fit keep int32, including both boundaries. + for value in (0, 1, 2**31 - 1, -(2**31)): + x = mx.array([value]) + self.assertEqual(x.dtype, mx.int32, msg=str(value)) + self.assertEqual(x[0].item(), value) + + # An explicit dtype still wins. + self.assertEqual(mx.array([2**40], mx.int64).dtype, mx.int64) + self.assertEqual(mx.array([1, 2], mx.int64).dtype, mx.int64) + # A float in the list still makes it float, not int64. + self.assertEqual(mx.array([2**40, 1.5]).dtype, mx.float32)   def test_construction_from_lists_of_mlx_arrays(self): dtypes = [ @@ -1251,6 +1309,28 @@ self.assertEqual(a.tolist(), [2, 1, 1])   a[0:2] = 3 self.assertEqual(a.tolist(), [3, 3, 1]) + + # Assigning through a bare Ellipsis, like a[:] and a[None] + e = mx.zeros((2, 3), mx.int32) + e[...] = 5 + self.assertEqual(e.tolist(), [[5, 5, 5], [5, 5, 5]]) + + # Broadcasting an array update through Ellipsis + e[...] = mx.array([1, 2, 3]) + self.assertEqual(e.tolist(), [[1, 2, 3], [1, 2, 3]]) + + e[...] = mx.zeros((2, 3), mx.int32) + self.assertEqual(e.tolist(), [[0, 0, 0], [0, 0, 0]]) + + # Scalar array + e = mx.array(0) + e[...] = 7 + self.assertEqual(e.item(), 7) + + # Shapes that cannot broadcast are still rejected + e = mx.zeros((2, 3), mx.int32) + with self.assertRaises(ValueError): + e[...] = mx.array([1, 2])   a[0:3] = 4 self.assertEqual(a.tolist(), [4, 4, 4])
diff --git ml-explore/mlx/python/tests/test_compile.py Layr-Labs/mlx/python/tests/test_compile.py index 7a2c6b9d0dce1be8a47ee459226a5aa50fa8fc5c..663ec295b0bff12e575745e801c2e0eec5e46ffe 100644 --- ml-explore/mlx/python/tests/test_compile.py +++ Layr-Labs/mlx/python/tests/test_compile.py @@ -87,6 +87,7 @@ mx.eval(y, z) results.append((y.item(), z.item())) except Exception as e: errors.append(e) + mx.clear_streams()   for _ in range(3): thread = threading.Thread(target=worker) @@ -97,6 +98,50 @@ if errors: raise errors[0] self.assertEqual(results, [(2.0, 2.0)] * 3) + + def test_compile_release_on_another_thread(self): + # A function traced on one thread but released on another must still + # drop its cache entry, otherwise a later compile of the same id gets + # handed the dead function's tape instead of being traced again. + traces = [] + + def fun(x): + traces.append(1) + return x + 1 + + holder = {} + traced = threading.Event() + released = threading.Event() + errors = [] + + def worker(): + try: + holder["fn"] = mx.compile(fun) + mx.eval(holder["fn"](mx.array([1.0]))) + traced.set() + self.assertTrue(released.wait(10)) + # The same callable, so the same id. + fn = mx.compile(fun) + mx.eval(fn(mx.array([1.0]))) + except Exception as e: + errors.append(e) + finally: + traced.set() + mx.clear_streams() + + # The tracing thread has to outlive the release, on exit it would tear + # down its cache anyway. + thread = threading.Thread(target=worker) + thread.start() + self.assertTrue(traced.wait(10)) + holder.clear() + gc.collect() + released.set() + thread.join() + + if errors: + raise errors[0] + self.assertEqual(len(traces), 2)   def test_compile_grad(self): def loss_fn(x): @@ -449,6 +494,7 @@ state_from_thread = {}   def grab(): state_from_thread["s"] = mx.random.state + mx.clear_streams()   t = threading.Thread(target=grab) t.start() @@ -482,6 +528,7 @@ e = fun() results["seed_changes"] = not bool( mx.allclose(c, e, 1e-2, 1e-2).item() ) + mx.clear_streams()   t = threading.Thread(target=worker) t.start() @@ -1568,6 +1615,46 @@ w = mx.arange(120, dtype=mx.float32).reshape(2, 3, 4, 5) expected = w[::-1, :, ::-1, :] + 1.0 self.assertTrue(mx.array_equal(p(w[::-1, :, ::-1, :]), expected)) + + def test_compile_abs_unsigned(self): + # abs has to compile for the wider unsigned types too + fun = lambda x: mx.abs(x) + 1 + for dtype in [mx.uint8, mx.uint16, mx.uint32, mx.uint64]: + x = mx.array([1, 2, 3], dtype) + self.assertTrue(mx.array_equal(mx.compile(fun)(x), fun(x))) + + def test_compiled_subnormal_bool_cast(self): + f32_sub = mx.array(np.array([0x00000001] * 4, dtype=np.uint32)).view(mx.float32) + f16_sub = mx.array(np.array([0x0001] * 4, dtype=np.uint16)).view(mx.float16) + bf16_sub = mx.array(np.array([0x0001] * 4, dtype=np.uint16)).view(mx.bfloat16) + + # A single-op compile does not fuse; the fused path needs >= 2 ops. + fn = mx.compile(lambda x: mx.broadcast_to(x, (2, 4)).astype(mx.bool_)) + for sub in (f32_sub, f16_sub, bf16_sub): + self.assertTrue(mx.all(fn(sub)).item()) + + def test_compile_different_log_bases(self): + # The logs are intermediates, since outputs are not simplified. + def entropies(p): + nats = -mx.sum(p * mx.log(p)) + bits = -mx.sum(p * mx.log2(p)) + return mx.stack([nats, bits]) + + p = np.array([0.1, 0.2, 0.3, 0.4], dtype=np.float32) + expected = np.array( + [-(p * np.log(p)).sum(), -(p * np.log2(p)).sum()], dtype=np.float32 + ) + out = mx.compile(entropies)(mx.array(p)) + self.assertTrue(np.allclose(out, expected, atol=1e-5)) + + def test_compile_equal_nan(self): + def fun(x): + return mx.stack( + [mx.array_equal(x, x), mx.array_equal(x, x, equal_nan=True)] + ) + + x = mx.array([1.0, float("nan"), 3.0]) + self.assertTrue(mx.array_equal(mx.compile(fun)(x), mx.array([False, True])))   if __name__ == "__main__":
diff --git ml-explore/mlx/python/tests/test_conv.py Layr-Labs/mlx/python/tests/test_conv.py index c5f9a2c1b245d5eca3a1ef665484653818c74ed0..062841c336ddfb320a7b4486cd8a41ebdc92f374 100644 --- ml-explore/mlx/python/tests/test_conv.py +++ Layr-Labs/mlx/python/tests/test_conv.py @@ -14,7 +14,7 @@ import torch import torch.nn.functional as F   has_torch = True -except ImportError as e: +except ImportError: has_torch = False   @@ -309,9 +309,11 @@ in_mx, wt_mx = map( lambda x: mx.array(x).astype(mx_dtype), (in_np, wt_np) ) in_pt, wt_pt = map( - lambda x: torch.from_numpy(x.transpose(0, 3, 1, 2)) - .to("cpu") - .to(torch_dtype), + lambda x: ( + torch.from_numpy(x.transpose(0, 3, 1, 2)) + .to("cpu") + .to(torch_dtype) + ), (in_np, wt_np), )   @@ -1069,7 +1071,6 @@ self.assertTrue(mx.allclose(y1, y2))   @unittest.skipIf(not has_torch, "requires Torch") def test_torch_conv_depthwise(self): - # fmt: off shapes = ( # N, H, W, C kH, kW, O, strides, padding, groups @@ -1214,12 +1215,130 @@ y = mx.conv_transpose2d(x, w, stream=mx.cpu) y_hat = mx.conv_transpose2d(x, w) self.assertTrue(mx.allclose(y, y_hat))   + @unittest.skipIf(not mx.metal.is_available(), "requires Metal") + def test_conv2d_winograd_batch_tiling(self): + # Use envs to test tiling without allocating large buffers. + tile_key = "MLX_CONV_WINOGRAD_TILE_BATCH" + ws_key = "MLX_CONV_WINOGRAD_WORKING_SET" + prev = {k: os.environ.get(k) for k in (tile_key, ws_key)} + + # Winograd needs 3x3 stride-1, channels in multiples of 32, + # C + O >= 256 and N * iH * iW >= 4096. + cases = ( + ((8, 48, 48, 64), (192, 3, 3, 64)), + ((5, 52, 44, 128), (128, 3, 3, 128)), + ((4, 48, 48, 192), (96, 3, 3, 192)), + ) + + def run(x, w, env={}): + for k in (tile_key, ws_key): + os.environ.pop(k, None) + os.environ.update(env) + y = mx.conv2d(x, w, padding=1) + mx.eval(y) + return np.array(y) + + try: + for in_shape, wt_shape in cases: + np.random.seed(0) + x = mx.array(np.random.normal(size=in_shape).astype(np.float32)) + # Small weights keep the output near unit scale. + w = mx.array( + (np.random.normal(size=wt_shape) * 0.05).astype(np.float32) + ) + b = mx.zeros((wt_shape[0],)) + mx.eval(x, w, b) + cpu_ref = np.array(mx.conv2d(x, w, padding=1, stream=mx.cpu)) + + untiled = run(x, w) + self.assertGreater(np.abs(untiled).max(), 0) + self.assertTrue(np.allclose(untiled, cpu_ref, atol=1e-3)) + + # Tiled winograd keeps the same per-element reduction order, + # so it is bit-identical to untiled; the implicit gemm + # fallback never is. Exact equality pins each run to its path. + # 3 divides none of the batches, so it also covers a short + # final tile. + for tile in (1, 3): + with self.subTest(in_shape=in_shape, tile=tile): + tiled = run(x, w, {tile_key: str(tile)}) + self.assertTrue(np.array_equal(untiled, tiled)) + + # A consumer op checks the output is fenced across + # command encoders. + os.environ[tile_key] = str(tile) + fused = mx.conv2d(x, w, padding=1) + b + mx.eval(fused) + self.assertTrue(np.allclose(untiled, fused, atol=1e-4)) + os.environ.pop(tile_key, None) + + # Budget for ~2 batch elements so the selector itself must + # tile; mirrors the winograd_batch_step arithmetic. + n, iH, iW, C = in_shape + O = wt_shape[0] + pH = 6 * ((iH + 2 - 2 + 5) // 6) + 2 + pW = 6 * ((iW + 2 - 2 + 5) // 6) + 2 + per_n = ( + pH * pW * C * 4 + + 64 * ((iH + 5) // 6) * ((iW + 5) // 6) * (C + O) * 4 + ) + used = (n * iH * iW * (C + O) + 64 * C * O) * 4 + with self.subTest(in_shape=in_shape, budget="tiled"): + budget = str(int((used + 5 * per_n // 2) / 0.75)) + tiled = run(x, w, {ws_key: budget}) + self.assertTrue(np.array_equal(untiled, tiled)) + + # Too small for even one batch element: must fall back. + with self.subTest(in_shape=in_shape, budget="infeasible"): + fallback = run(x, w, {ws_key: "1"}) + self.assertFalse(np.array_equal(untiled, fallback)) + self.assertTrue(np.allclose(fallback, cpu_ref, atol=1e-3)) + + # A forced tile is capped by the budget, so this must still + # fall back. + with self.subTest(in_shape=in_shape, budget="forced+infeasible"): + capped = run(x, w, {ws_key: "1", tile_key: "1"}) + self.assertFalse(np.array_equal(untiled, capped)) + self.assertTrue(np.allclose(capped, cpu_ref, atol=1e-3)) + finally: + for k, v in prev.items(): + if v is None: + os.environ.pop(k, None) + else: + os.environ[k] = v + def test_conv2d_large_filter_small_channels(self): x = mx.random.normal(shape=(1, 181, 181, 1)) w = mx.random.normal(shape=(1, 182, 182, 1)) y = mx.conv2d(x, w, (1, 1), (1, 1), stream=mx.cpu) y_hat = mx.conv2d(x, w, (1, 1), (1, 1)) self.assertTrue(mx.allclose(y, y_hat, rtol=1e-3, atol=1e-3)) + + def test_conv_3D_small_kd_decomposition(self): + # Exercises the small kernel-depth 3D -> KD x 2D decomposition (#3625): + # N=1, small KD, depth stride/dilation 1, no depth padding, mod16 channels. + # Validated against the CPU reference, which uses a different code path. + for T, H, W, Cin, Cout, kd, kh, kw in [ + (5, 16, 16, 32, 32, 3, 3, 3), # canonical 3x3x3 (2D hits Winograd) + (4, 12, 10, 16, 48, 3, 3, 3), # Cout != Cin + (6, 14, 14, 32, 32, 1, 3, 3), # KD = 1 + (5, 12, 12, 16, 16, 5, 1, 1), # larger KD, 1x1 spatial + (4, 10, 10, 32, 16, 2, 3, 3), # KD = 2 + ]: + x = mx.random.normal((1, T, H, W, Cin)) + w = mx.random.normal((Cout, kd, kh, kw, Cin)) + # flip mirrors every kernel axis, including the decomposed depth + for flip in (False, True): + y_gpu = mx.conv_general(x, w, stride=(1, 1, 1), flip=flip) + y_cpu = mx.conv_general( + x, w, stride=(1, 1, 1), flip=flip, stream=mx.cpu + ) + mx.eval(y_gpu, y_cpu) + self.assertTrue( + mx.allclose(y_gpu, y_cpu, rtol=1e-4, atol=1e-4), + f"3D small-kd mismatch T{T} H{H} W{W} " + f"C{Cin}->{Cout} k{kd}{kh}{kw} flip={flip}", + )   if __name__ == "__main__":
diff --git ml-explore/mlx/python/tests/test_conv_transpose.py Layr-Labs/mlx/python/tests/test_conv_transpose.py index e6def081a7f5b90a09ea043b7b4a9f55cc97a601..7d12cb16f42828b7ef29e4abfa5c04f8d3374982 100644 --- ml-explore/mlx/python/tests/test_conv_transpose.py +++ Layr-Labs/mlx/python/tests/test_conv_transpose.py @@ -486,6 +486,8 @@ for idim, kdim, stride, padding in ( ((1, 1, 1), (1, 1, 1), (1, 1, 1), (0, 0, 0)), ((3, 3, 3), (3, 1, 1), (1, 1, 1), (0, 0, 0)), ((15, 15, 15), (3, 3, 3), (3, 3, 3), (2, 2, 2)), + # Exercises the Metal phase-aware stride-2/kernel-2 path. + ((3, 4, 5), (2, 2, 2), (2, 2, 2), (0, 0, 0)), ): run_conv_transpose3D( N, C, O, idim, kdim, stride, padding, dtype=dtype
diff --git ml-explore/mlx/python/tests/test_double.py Layr-Labs/mlx/python/tests/test_double.py index 65603cd9376f9977b369622a686bb2ebbaa77dc7..3186e7e57306d33c17f67c4e01efeb06fc583f0f 100644 --- ml-explore/mlx/python/tests/test_double.py +++ Layr-Labs/mlx/python/tests/test_double.py @@ -336,8 +336,14 @@ self.assertEqual(padded.dtype, dtype)   def test_linspace(self): with mx.stream(mx.cpu): - vals = mx.linspace(0, math.pi, 2, mx.float64) + vals = mx.linspace(0, math.pi, 2, dtype=mx.float64) self.assertEqual(vals.tolist()[1], math.pi) + + vals = mx.linspace(0, math.pi, 4, endpoint=False, dtype=mx.float64) + self.assertEqual(vals.dtype, mx.float64) + self.assertTrue( + np.allclose(vals.tolist(), np.linspace(0, math.pi, 4, endpoint=False)) + )   if __name__ == "__main__":
diff --git ml-explore/mlx/python/tests/test_einsum.py Layr-Labs/mlx/python/tests/test_einsum.py index a73ea381872f8888a8ef070c4676272de03b9dee..08884ae9812d9de8ff8f39ccd2124ac8c4e39b7a 100644 --- ml-explore/mlx/python/tests/test_einsum.py +++ Layr-Labs/mlx/python/tests/test_einsum.py @@ -65,6 +65,39 @@ inputs = [mx.array(i) for i in inputs] mx_path = mx.einsum_path(case, *inputs) self.assertEqual(np_path[0][1:], mx_path[0])   + def test_scalar_operands(self): + # An empty subscript is a scalar operand. A trailing one used to be + # dropped by the parser, so "i,->i" looked like a single input. + s1 = mx.array(2.0) + s2 = mx.array(3.0) + v = mx.random.uniform(shape=(3,)) + m = mx.random.uniform(shape=(2, 3)) + + cases = [ + ("->", (s1,)), + (",->", (s1, s2)), + (",,->", (s1, s2, s1)), + ("i,->i", (v, s1)), + (",i->i", (s1, v)), + ("ij,->ij", (m, s1)), + (",ij->ij", (s1, m)), + ("i,,->i", (v, s1, s2)), + ] + for spec, operands in cases: + mx_out = mx.einsum(spec, *operands) + np_out = np.einsum(spec, *[np.array(o) for o in operands]) + self.assertEqual(mx_out.shape, np_out.shape) + self.assertTrue(np.allclose(mx_out, np_out, rtol=1e-4, atol=1e-4)) + + # Operand count still has to match the number of subscripts + with self.assertRaises(ValueError): + mx.einsum(",->", s1) + with self.assertRaises(ValueError): + mx.einsum("i,->i", v) + # An empty subscript requires a 0-d operand + with self.assertRaises(ValueError): + mx.einsum(",->", v, s1) + def test_simple_einsum(self): a = mx.arange(4 * 4).reshape(4, 4) a_mx = mx.einsum("ii->i", a) @@ -188,7 +221,6 @@ def test_broadcasting(self): a = mx.full((5, 1), 1.0) b = mx.full((8, 2), 1.0) a_mx = mx.einsum("ab,bc->c", a, b) - return a_np = np.einsum("ab,bc->c", a, b) self.assertTrue(np.array_equal(a_mx, a_np))   @@ -357,6 +389,40 @@ for test_case in error_tests: inputs = inputs_for_case(test_case[0]) with self.assertRaises(ValueError): mx.einsum(test_case[1], *inputs) + + def test_ellipses_broadcast(self): + # Size 1 batch dimensions covered by an ellipsis have to broadcast + # against the other operands, including when the smaller operand + # comes first. + shape_pairs = [ + ((1, 3, 4), (2, 4, 5)), + ((2, 3, 4), (1, 4, 5)), + ((1, 1, 3, 4), (5, 2, 4, 5)), + ((5, 1, 3, 4), (1, 2, 4, 5)), + ] + for sa, sb in shape_pairs: + a = mx.random.uniform(shape=sa) + b = mx.random.uniform(shape=sb) + mx_out = mx.einsum("...ij,...jk->...ik", a, b) + np_out = np.einsum("...ij,...jk->...ik", np.array(a), np.array(b)) + self.assertEqual(mx_out.shape, np_out.shape) + self.assertTrue(np.allclose(mx_out, np_out, rtol=1e-4, atol=1e-4)) + + for sa, sb in [((1, 4), (5, 4)), ((5, 4), (1, 4))]: + a = mx.random.uniform(shape=sa) + b = mx.random.uniform(shape=sb) + mx_out = mx.einsum("...i,...i->...", a, b) + np_out = np.einsum("...i,...i->...", np.array(a), np.array(b)) + self.assertEqual(mx_out.shape, np_out.shape) + self.assertTrue(np.allclose(mx_out, np_out, rtol=1e-4, atol=1e-4)) + + # Same thing with explicit labels rather than an ellipsis + a = mx.random.uniform(shape=(1, 3, 4)) + b = mx.random.uniform(shape=(2, 4, 5)) + mx_out = mx.einsum("bij,bjk->bik", a, b) + np_out = np.einsum("bij,bjk->bik", np.array(a), np.array(b)) + self.assertEqual(mx_out.shape, np_out.shape) + self.assertTrue(np.allclose(mx_out, np_out, rtol=1e-4, atol=1e-4))   if __name__ == "__main__":
diff --git ml-explore/mlx/python/tests/test_eval.py Layr-Labs/mlx/python/tests/test_eval.py index da7f0ea8df91b36385ebe2bff8a3ebd28be8b465..265a090aa2d7db301fec18202ab160f4f571e100 100644 --- ml-explore/mlx/python/tests/test_eval.py +++ Layr-Labs/mlx/python/tests/test_eval.py @@ -227,6 +227,15 @@ # Fresh computations after the failure stay correct. x = mx.full((512,), 2.0) self.assertEqual((x + 1.0).sum().item(), 512.0 * 3.0)   + @unittest.skipIf( + mx.cuda.is_available(), "CUDA backend waits cpu stream synchronously" + ) + def test_async_eval_error_in_synchronize(self): + a = mx.linalg.inv(mx.array([[1.0, 2.0], [2.0, 4.0]]), stream=mx.cpu) + mx.async_eval(a) + with self.assertRaises(RuntimeError): + mx.synchronize(mx.cpu) +   if __name__ == "__main__": mlx_tests.MLXTestRunner()
diff --git ml-explore/mlx/python/tests/test_fast.py Layr-Labs/mlx/python/tests/test_fast.py index 5dacaa605c26a09ef3f939eb239c54198b870ddc..200781d3727823d141435d7c5ac365e1e154b971 100644 --- ml-explore/mlx/python/tests/test_fast.py +++ Layr-Labs/mlx/python/tests/test_fast.py @@ -525,6 +525,67 @@ gx2, gw2 = mx.grad(gf(f2), argnums=(0, 1))(x, w, y) self.assertLess(mx.abs(gx1 - gx2).max(), 1e-5) self.assertLess(mx.abs(gw1 - gw2).max() / mx.abs(gw1).mean(), 1e-5)   + def test_cross_entropy(self): + def cross_entropy_ref(logits, targets): + score = mx.take_along_axis(logits, mx.expand_dims(targets, -1), -1).squeeze( + -1 + ) + return mx.logsumexp(logits.astype(mx.float32), axis=-1) - score.astype( + mx.float32 + ) + + tolerances = {mx.float32: 1e-5, mx.float16: 3e-2, mx.bfloat16: 3e-1} + + for V in [7, 32, 128, 255, 256, 1000, 4096, 8192]: + for dtype in [mx.float32, mx.float16, mx.bfloat16]: + logits = (mx.random.normal(shape=(4, 7, V), scale=3.0) * 2).astype( + dtype + ) + targets = mx.random.randint(0, V, shape=(4, 7)) + expected = cross_entropy_ref(logits, targets) + out = mx.fast.cross_entropy(logits, targets) + self.assertEqual(out.dtype, mx.float32) + self.assertEqual(out.shape, targets.shape) + self.assertLess(mx.abs(out - expected).max().item(), tolerances[dtype]) + + def test_cross_entropy_shape_checks(self): + logits = mx.random.normal(shape=(4, 16)) + with self.assertRaises(ValueError): + mx.fast.cross_entropy(logits, mx.zeros((5,), mx.int32)) + with self.assertRaises(ValueError): + # Probability targets are not supported by the fused op. + mx.fast.cross_entropy(logits, mx.zeros((4, 16), mx.int32)) + with self.assertRaises(ValueError): + mx.fast.cross_entropy(logits, mx.zeros((4,), mx.float32)) + + def test_cross_entropy_grad(self): + def ref(logits, targets): + score = mx.take_along_axis(logits, mx.expand_dims(targets, -1), -1).squeeze( + -1 + ) + return mx.logsumexp(logits, axis=-1) - score + + f1 = lambda x, y: ref(x, y).mean() + f2 = lambda x, y: mx.fast.cross_entropy(x, y).mean() + + for V in [7, 128, 1000, 4096]: + logits = mx.random.normal(shape=(4, 7, V), scale=2.0) + targets = mx.random.randint(0, V, shape=(4, 7)) + g1 = mx.grad(f1, argnums=0)(logits, targets) + g2 = mx.grad(f2, argnums=0)(logits, targets) + self.assertEqual(g2.shape, logits.shape) + self.assertLess(mx.abs(g1 - g2).max().item(), 1e-6) + + w = mx.random.uniform(shape=(4, 7)) + f3 = lambda x, y: (ref(x, y) * w).sum() + f4 = lambda x, y: (mx.fast.cross_entropy(x, y) * w).sum() + logits = mx.random.normal(shape=(4, 7, 512), scale=2.0) + targets = mx.random.randint(0, 512, shape=(4, 7)) + g1 = mx.grad(f3, argnums=0)(logits, targets) + g2 = mx.grad(f4, argnums=0)(logits, targets) + self.assertEqual(g2.shape, logits.shape) + self.assertLess(mx.abs(g1 - g2).max().item(), 1e-6) + def test_layer_norm_dim_check(self): with self.assertRaises(ValueError): weight = mx.ones((129,)) @@ -1055,6 +1116,35 @@ out_b = call_kernel( a, "uint e = thread_position_in_grid.x; out[e] = inp[e] + 100.0f;" ) mx.eval(out_a, out_b) # one batch — the reported failure case + self.assertTrue(mx.array_equal(out_a, a * 2.0)) + self.assertTrue(mx.array_equal(out_b, a + 100.0)) + + @unittest.skipIf(not mx.cuda.is_available(), "CUDA is not available") + def test_cuda_kernel_same_name_different_source(self): + # The CUDA module cache was keyed on the kernel name alone, so the + # second kernel here silently ran the first one's code. Metal had the + # same bug, fixed in #3833. + def call_kernel(a, source): + kernel = mx.fast.cuda_kernel( + name="dup_name", + input_names=["inp"], + output_names=["out"], + source=source, + ) + return kernel( + inputs=[a], + grid=(a.size, 1, 1), + threadgroup=(a.size, 1, 1), + output_shapes=[a.shape], + output_dtypes=[a.dtype], + stream=mx.gpu, + )[0] + + a = mx.arange(32, dtype=mx.float32) + elem = "auto e = cooperative_groups::this_grid().thread_rank();" + out_a = call_kernel(a, f"{elem} out[e] = inp[e] * 2.0f;") + out_b = call_kernel(a, f"{elem} out[e] = inp[e] + 100.0f;") + mx.eval(out_a, out_b) self.assertTrue(mx.array_equal(out_a, a * 2.0)) self.assertTrue(mx.array_equal(out_b, a + 100.0))
diff --git ml-explore/mlx/python/tests/test_fft.py Layr-Labs/mlx/python/tests/test_fft.py index 9358ede794eb344209cb1a07d2c6f88a639ac141..1f96aad5668af880d76cdca6485363f010bf5628 100644 --- ml-explore/mlx/python/tests/test_fft.py +++ Layr-Labs/mlx/python/tests/test_fft.py @@ -446,6 +446,41 @@ dfdx = mx.grad(f)(x) dgdx = torch.func.grad(g)(x_torch) self.assertLess((dfdx - dgdx).abs().max() / dgdx.abs().mean(), 1e-4)   + def make_ffts(self): + mxffts = { + (True, True): mx.fft.irfftn, + (True, False): mx.fft.rfftn, + (False, True): mx.fft.ifftn, + (False, False): mx.fft.fftn, + } + shape = (3, 8, 6) + r = np.random.rand(*shape).astype(np.float32) + i = np.random.rand(*shape).astype(np.float32) + for (real, inverse), fftn in mxffts.items(): + a_np = r if real and not inverse else r + 1j * i + for axes in [(-1,), (0,), (-2, -1), (-1, -2), (0, 1)]: + yield fftn, a_np, axes + + def test_fft_vmap(self): + for fftn, a_np, axes in self.make_ffts(): + a = mx.array(a_np) + f = lambda x: fftn(x, axes=axes) + expected = mx.stack([f(a[i]) for i in range(a.shape[0])]) + out = mx.vmap(f)(a) + self.assertEqual(tuple(out.shape), tuple(expected.shape)) + np.testing.assert_allclose(out, expected, atol=1e-5, rtol=1e-5) + + def test_fft_jvp(self): + # The fft is linear so the jvp is the fft of the tangent + for fftn, a_np, axes in self.make_ffts(): + a = mx.array(a_np) + t = mx.array(np.random.rand(*a_np.shape).astype(a_np.dtype)) + f = lambda x: fftn(x, axes=axes) + expected = f(t) + out = mx.jvp(f, [a], [t])[1][0] + self.assertEqual(tuple(out.shape), tuple(expected.shape)) + np.testing.assert_allclose(out, expected, atol=1e-5, rtol=1e-5) +   if __name__ == "__main__": mlx_tests.MLXTestRunner()
diff --git ml-explore/mlx/python/tests/test_load.py Layr-Labs/mlx/python/tests/test_load.py index 1c52f333a69f0475289c49b166fa59548bb71b57..f5947de471d0dd958d0e519c6e18cdfbc05de4a2 100644 --- ml-explore/mlx/python/tests/test_load.py +++ Layr-Labs/mlx/python/tests/test_load.py @@ -88,6 +88,40 @@ np.save(save_file, c) with self.assertRaises(Exception): out = mx.load(save_file, stream=mx.cpu)   + def test_load_npy_read_error(self): + save_file = os.path.join(self.test_dir, "truncated.npy") + expected = np.arange(16, dtype=np.float32) + np.save(save_file, expected) + with open(save_file, "r+b") as f: + f.truncate(os.path.getsize(save_file) - expected.nbytes) + + out = mx.load(save_file, stream=mx.cpu) + with self.assertRaises(RuntimeError): + mx.eval(out) + + def test_async_load_npy_read_error_across_streams(self): + save_file = os.path.join(self.test_dir, "truncated_async.npy") + expected = np.arange(16, dtype=np.float32) + np.save(save_file, expected) + with open(save_file, "r+b") as f: + f.truncate(os.path.getsize(save_file) - expected.nbytes) + + producer_stream = mx.new_stream(mx.cpu) + consumer_stream = mx.new_stream(mx.cpu) + out = mx.add( + mx.load(save_file, stream=producer_stream), + 1.0, + stream=consumer_stream, + ) + with self.assertRaises(RuntimeError): + mx.eval(out) + # Depending on backend the error might be caught early before poisoning + # the producer_stream, but still sync to clear the errors. + try: + mx.synchronize(producer_stream) + except Exception: + pass + def test_save_and_load_safetensors(self): test_file = os.path.join(self.test_dir, "test.safetensors") with self.assertRaises(Exception):
diff --git ml-explore/mlx/python/tests/test_nn.py Layr-Labs/mlx/python/tests/test_nn.py index c5e6db94a72927be57c8248c62c43ef61b34858d..26b8fd11623e04db0b2799d03cab3f04cfd4a55b 100644 --- ml-explore/mlx/python/tests/test_nn.py +++ Layr-Labs/mlx/python/tests/test_nn.py @@ -410,6 +410,28 @@ layer = nn.Bilinear(input1_dims=2, input2_dims=4, output_dims=6) outputs = layer(inputs1, inputs2) self.assertEqual(outputs.shape, (10, 6))   + def test_norm_eps_validation(self): + # eps is added under a square root. A negative one makes rsqrt take the + # root of a negative number, so the layer emits NaN for whichever + # elements have a small enough variance, which is only a partial NaN and + # easy to miss. Zero is rejected too: it leaves rsqrt(0) for any input + # whose variance is zero, which is a whole NaN row. This matches the eps + # guards the optimizers already carry. + builders = ( + ("LayerNorm", lambda eps: nn.LayerNorm(16, eps=eps)), + ("RMSNorm", lambda eps: nn.RMSNorm(16, eps=eps)), + ("GroupNorm", lambda eps: nn.GroupNorm(4, 16, eps=eps)), + ("InstanceNorm", lambda eps: nn.InstanceNorm(16, eps=eps)), + ("BatchNorm", lambda eps: nn.BatchNorm(16, eps=eps)), + ) + for name, build in builders: + for eps in (-1.0, -1e-30, 0.0): + with self.assertRaisesRegex(ValueError, "must be positive"): + build(eps) + # Anything positive still constructs, including a very small eps. + for eps in (1e-30, 1e-5, 1.0): + build(eps) + def test_group_norm(self): x = mx.arange(100, dtype=mx.float32) x = x.reshape(1, 10, 10, 1) @@ -671,6 +693,21 @@ ], ] self.assertTrue(x.shape == y.shape) self.assertTrue(np.allclose(y, expected_y, atol=1e-5)) + # Reduced-precision statistics must not overflow for finite feature maps. + checkerboard = np.indices((4, 4, 4)).sum(axis=0) % 2 + x = mx.array( + np.stack( + [ + np.where(checkerboard, -512, 512), + np.where(checkerboard, -256, 256), + ], + axis=-1, + ).astype(np.float16) + )[None] + y = nn.InstanceNorm(dims=2)(x) + self.assertEqual(y.dtype, mx.float16) + self.assertTrue(mx.allclose(y.min(), mx.array(-1.0, dtype=mx.float16))) + self.assertTrue(mx.allclose(y.max(), mx.array(1.0, dtype=mx.float16))) # Test repr self.assertTrue(str(inorm) == "InstanceNorm(3, eps=1e-05, affine=False)") # Raise for inputs without spatial dimensions
diff --git ml-explore/mlx/python/tests/test_ops.py Layr-Labs/mlx/python/tests/test_ops.py index 1b46237ce5de4b0fe6c21d0a0464aa69e764d54e..0c94989e0205c7f6ce35e1191a8ecc43de482fbb 100644 --- ml-explore/mlx/python/tests/test_ops.py +++ Layr-Labs/mlx/python/tests/test_ops.py @@ -126,6 +126,27 @@ with self.assertRaises(OverflowError) as cm: mx.broadcast_to(a, [too_big, 1]) self.assertIn(str(too_big), str(cm.exception))   + # A concatenation axis that does not fit is computed rather than given, + # so it has to be reported instead of wrapping into a bogus dimension. + # These stay lazy, so nothing near this size is allocated. + big = mx.zeros(2**30) + for parts in (3, 4, 5): + with self.assertRaises(OverflowError) as cm: + mx.concatenate([big] * parts) + self.assertIn(str(2**30 * parts), str(cm.exception)) + + # repeat and kron multiply a dimension, and used to wrap into a + # negative or zero one that only surfaced later as a confusing reshape + # error naming a shape the caller never asked for. + for parts in (2, 3, 4): + with self.assertRaises(OverflowError) as cm: + mx.repeat(big, parts) + self.assertIn(str(2**30 * parts), str(cm.exception)) + + with self.assertRaises(OverflowError) as cm: + mx.kron(mx.zeros(2**16), mx.zeros(2**16)) + self.assertIn(str(2**32), str(cm.exception)) + # Negative overflow (< int32 min) is caught too. too_negative = -(2**31) - 1 with self.assertRaises(OverflowError) as cm: @@ -137,6 +158,20 @@ self.assertEqual(mx.zeros(4).shape, (4,)) self.assertEqual(mx.zeros((2, 3)).shape, (2, 3)) self.assertEqual(mx.ones([2, 3]).shape, (2, 3)) self.assertEqual(mx.full((2, 3), 1.5).tolist(), [[1.5] * 3] * 2) + + def test_integer_index_protocol(self): + a = mx.arange(4) + + index = np.int32(2) + self.assertEqual(mx.topk(a, index).shape, (2,)) + self.assertEqual(mx.reshape(a, [index, 2]).shape, (2, 2)) + + for value in (np.float32(2), "2"): + with self.subTest(value=value): + with self.assertRaises(TypeError): + mx.topk(a, value) + with self.assertRaises(TypeError): + mx.reshape(a, [value, 2])   def test_scalar_inputs(self): # Check combinations of python types @@ -344,6 +379,14 @@ self.assertEqual(z.dtype, mx.int32) self.assertEqual(z.item(), 2)   def test_remainder(self): + # Complex is not supported and has to say so rather than quietly + # computing a componentwise remainder, which no other library defines + z = mx.array([7 + 3j], mx.complex64) + with self.assertRaises(ValueError): + mx.remainder(z, z) + with self.assertRaises(ValueError): + z % z + for dt in [mx.int32, mx.float32, mx.float16, mx.bfloat16]: x = mx.array(2, dtype=dt) y = mx.array(4, dtype=dt) @@ -962,6 +1005,38 @@ out = mx.median(x, axis=(0, 1, 3), keepdims=True) out_np = np.median(x, axis=(0, 1, 3), keepdims=True) self.assertTrue(np.allclose(out, out_np))   + def test_median_nan(self): + nan = float("nan") + + # Odd and even lengths, with the NaN in a few different positions. + for vals in ([1.0, nan, 0.0], [nan, 1.0, 0.0], [1.0, 2.0, nan, 4.0]): + for dtype in (mx.float16, mx.bfloat16, mx.float32): + out = mx.median(mx.array(vals, dtype=dtype)) + self.assertTrue(mx.isnan(out).item(), msg=f"{vals} {dtype}") + + x = mx.array([[1.0, nan, 3.0], [4.0, 5.0, 6.0]]) + self.assertTrue( + np.array_equal( + np.array(mx.median(x, axis=1)), np.median(x, axis=1), equal_nan=True + ) + ) + self.assertTrue( + np.array_equal( + np.array(mx.median(x, axis=0)), np.median(x, axis=0), equal_nan=True + ) + ) + self.assertTrue(mx.isnan(mx.median(x)).item()) + self.assertEqual(mx.median(x, axis=1, keepdims=True).shape, (2, 1)) + + # Complex NaN propagates too, matching NumPy. + out = mx.median(mx.array([complex(1, 0), complex(nan, 0), complex(0, 0)])) + self.assertTrue(mx.isnan(out).item()) + + # A NaN-free array is unaffected, and integers are never NaN. + x = mx.array([[1.0, 2.0, 3.0], [4.0, 5.0, 6.0]]) + self.assertTrue(np.allclose(mx.median(x, axis=1), np.median(x, axis=1))) + self.assertEqual(mx.median(mx.array([0, 1, 2, 3, 4])).item(), 2) + def test_var(self): x = mx.array( [ @@ -985,11 +1060,21 @@ x = mx.array([1.0, 2.0]) out = mx.var(x, ddof=3) self.assertEqual(out.item(), float("inf"))   + x = mx.array([1 + 2j, -3 - 4j, 0.5 - 0.25j]) + x_np = np.array(x) + self.assertEqual(mx.var(x).dtype, mx.float32) + self.assertAlmostEqual(mx.var(x).item(), x_np.var().item(), places=5) + def test_std(self): x = mx.random.uniform(shape=(5, 5)) x_np = np.array(x) self.assertAlmostEqual(mx.std(x).item(), x_np.std().item(), places=6)   + x = mx.array([1 + 2j, -3 - 4j, 0.5 - 0.25j]) + x_np = np.array(x) + self.assertEqual(mx.std(x).dtype, mx.float32) + self.assertAlmostEqual(mx.std(x).item(), x_np.std().item(), places=5) + def test_abs(self): a = mx.array([-1.0, 1.0, -2.0, 3.0]) result = mx.abs(a) @@ -1155,11 +1240,28 @@ expected = np.expm1(a) np.seterr(over=errs["over"]) self.assertTrue(np.allclose(result, expected, rtol=1e-3, atol=1e-4))   + # Complex is not supported and has to say so rather than quietly + # computing on the real part + z = mx.array([1 + 2j], mx.complex64) + with self.assertRaises(ValueError): + mx.expm1(z) + with self.assertRaises(ValueError): + mx.sigmoid(z) + with self.assertRaises(ValueError): + mx.arctan2(z, z) + def test_erf(self): inputs = [-5, 0.0, 0.5, 1.0, 2.0, 10.0] x = mx.array(inputs) expected = np.array([math.erf(i) for i in inputs]) self.assertTrue(np.allclose(mx.erf(x), expected)) + + # Complex is not supported and has to say so rather than abort + z = mx.array([1 + 2j], mx.complex64) + with self.assertRaises(ValueError): + mx.erf(z) + with self.assertRaises(ValueError): + mx.erfinv(z)   def test_erfinv(self): inputs = [-5.0, -1.0, 0.5, 0.0, 0.5, 1.0, 5.0] @@ -1299,6 +1401,21 @@ self.assertEqual(mx.any(a, axis=[1]).tolist(), [True, False]) self.assertEqual(mx.any(a, axis=0).tolist(), [True, False]) self.assertEqual(mx.any(a, axis=1).tolist(), [True, False])   + def test_subnormal_bool_cast(self): + f32_sub = mx.array(np.array([0x00000001], dtype=np.uint32)).view(mx.float32) + f16_sub = mx.array(np.array([0x0001], dtype=np.uint16)).view(mx.float16) + bf16_sub = mx.array(np.array([0x0001], dtype=np.uint16)).view(mx.bfloat16) + + self.assertTrue(f32_sub.astype(mx.bool_).item()) + self.assertTrue(f16_sub.astype(mx.bool_).item()) + self.assertTrue(bf16_sub.astype(mx.bool_).item()) + self.assertTrue(mx.any(f32_sub).item()) + self.assertTrue(mx.any(f16_sub).item()) + self.assertTrue(mx.any(bf16_sub).item()) + self.assertTrue(mx.all(f32_sub).item()) + self.assertTrue(mx.all(f16_sub).item()) + self.assertTrue(mx.all(bf16_sub).item()) + def test_stop_gradient(self): def func(x): return mx.sum(2 * x + mx.stop_gradient(3 * x)) @@ -1606,9 +1723,11 @@ with self.assertRaises(ValueError): a = mx.arange(float("inf"), 1, float("inf")) with self.assertRaises(ValueError): a = mx.arange(float("inf"), 1, 5) - with self.assertRaises(TypeError): + with self.assertRaises(ValueError): INT_MAX = 2147483647 a = mx.arange(0, INT_MAX + 1, 1) + with self.assertRaises(ValueError): + a = mx.arange(0, 2**40)   a = mx.arange(5) expected = [0, 1, 2, 3, 4] @@ -1672,6 +1791,57 @@ a = mx.arange(1.0, 3.0, 0.2, dtype=mx.int32) self.assertEqual(a.dtype, mx.int32)   + # Integers that do not fit in int32 widen the inferred dtype to int64, + # matching the scalar inference of mx.array and numpy. + a = mx.arange(2**40, 2**40 + 3) + self.assertEqual(a.dtype, mx.int64) + self.assertListEqual(a.tolist(), [2**40, 2**40 + 1, 2**40 + 2]) + + a = mx.arange(-(2**40), -(2**40) + 3) + self.assertEqual(a.dtype, mx.int64) + self.assertListEqual(a.tolist(), [-(2**40), -(2**40) + 1, -(2**40) + 2]) + + # int32 boundaries themselves still infer int32. + a = mx.arange(2**31 - 3, 2**31 - 1) + self.assertEqual(a.dtype, mx.int32) + self.assertListEqual(a.tolist(), [2**31 - 3, 2**31 - 2]) + + a = mx.arange(-(2**31), -(2**31) + 2) + self.assertEqual(a.dtype, mx.int32) + self.assertListEqual(a.tolist(), [-(2**31), -(2**31) + 1]) + + # The first values that no longer fit widen as well. + a = mx.arange(-(2**31) - 1, -(2**31) + 1) + self.assertEqual(a.dtype, mx.int64) + self.assertListEqual(a.tolist(), [-(2**31) - 1, -(2**31)]) + + # A large step also widens the inferred dtype. + a = mx.arange(2**40, 2**40 + 3, 2**40) + self.assertEqual(a.dtype, mx.int64) + self.assertListEqual(a.tolist(), [2**40]) + + a = mx.arange(stop=2, step=2**40) + self.assertEqual(a.dtype, mx.int64) + self.assertListEqual(a.tolist(), [0]) + + # A negative step with widened values. + a = mx.arange(2**40 + 3, 2**40, -1) + self.assertEqual(a.dtype, mx.int64) + self.assertListEqual(a.tolist(), [2**40 + 3, 2**40 + 2, 2**40 + 1]) + + # The stop-only overload widens too, even for an empty result. + a = mx.arange(stop=2**40, step=-1) + self.assertEqual(a.dtype, mx.int64) + self.assertEqual(a.shape, (0,)) + + # An explicit dtype takes precedence over the widened inference. + a = mx.arange(2**40, 2**40 + 3, dtype=mx.int32) + self.assertEqual(a.dtype, mx.int32) + + # A float in the mix still infers float32. + a = mx.arange(0.5, 2**40, 2**39) + self.assertEqual(a.dtype, mx.float32) + def test_arange_corner_cases_cast(self): a = mx.arange(0, 3, 0.2, dtype=mx.int32) expected = [0] * 15 @@ -1725,6 +1895,20 @@ a = mx.arange(0, -10, float("-inf")) expected = [0] self.assertListEqual(a.tolist(), expected) + + # The range crossing the int32 limit widens the dtype to int64 instead + # of saturating or wrapping. + n = mx.iinfo(mx.int32).max + result = mx.arange(n - 1, n + 3) + self.assertEqual(result.shape, (4,)) + self.assertEqual(result.dtype, mx.int64) + self.assertEqual(result.tolist(), [n - 1, n, n + 1, n + 2]) + + # An explicit dtype keeps the previous wrapping behaviour. + result = mx.arange(n - 1, n + 3, dtype=mx.int32) + self.assertEqual(result.shape, (4,)) + self.assertEqual(result.dtype, mx.int32) + self.assertEqual(result.tolist(), [n - 1, n, -2147483648, -2147483647])   def test_hanning_general(self): a = mx.hanning(10) @@ -2132,6 +2316,11 @@ def test_meshgrid(self): x = mx.array([1, 2, 3], dtype=mx.int32) y = np.array([1, 2, 3], dtype=np.int32)   + # Test return type is a tuple + self.assertIsInstance(mx.meshgrid(x), tuple) + self.assertIsInstance(mx.meshgrid(x, x), tuple) + self.assertIsInstance(mx.meshgrid(x, x, x, sparse=True), tuple) + # Test single input a_mlx = mx.meshgrid(x) a_np = np.meshgrid(y) @@ -2262,7 +2451,7 @@ out_np = np.nan_to_num(a) self.assertTrue(np.allclose(out_mx, out_np))   for t in [mx.float32, mx.float16]: - a = mx.array([float("inf"), 6.9, float("nan"), float("-inf")]) + a = mx.array([float("inf"), 6.9, float("nan"), float("-inf")]).astype(t) out_mx = mx.nan_to_num(a) out_np = np.nan_to_num(a) self.assertTrue(np.allclose(out_mx, out_np)) @@ -2271,6 +2460,16 @@ a = mx.array([float("inf"), 6.9, float("nan"), float("-inf")]).astype(t) out_np = np.nan_to_num(a, nan=0.0, posinf=1000, neginf=-1000) out_mx = mx.nan_to_num(a, nan=0.0, posinf=1000, neginf=-1000) self.assertTrue(np.allclose(out_mx, out_np)) + + # bfloat16 has no numpy analogue; infinities should clamp to the + # dtype's largest finite value, not 0 + a = mx.array([float("inf"), 6.9, float("nan"), float("-inf")]).astype( + mx.bfloat16 + ) + out_mx = mx.nan_to_num(a) + bf_max = mx.finfo(mx.bfloat16).max + expected = mx.array([bf_max, 6.9, 0.0, -bf_max]).astype(mx.bfloat16) + self.assertTrue(mx.array_equal(out_mx, expected))   def test_pad_reflect_symmetric(self): # mx.pad reflect/symmetric must match numpy.pad exactly. Covers @@ -2478,6 +2677,36 @@ mx.synchronize() mem4 = mx.get_peak_memory() self.assertEqual(mem2, mem4)   + def test_scan_size_one_axis(self): + # A size one axis can carry any stride and still be row contiguous, so + # the scan must not take its row count from that stride. + for op in ["cumsum", "cumprod", "cummax", "cummin"]: + for start in (1, 2, 3): + with self.subTest(op=op, start=start): + base = mx.arange(1, 11, dtype=mx.float32).reshape(1, 10) + a = base[:, start:] + mx.eval(a) + # The axis has size one, so an inclusive scan is the identity + expected = np.array(a).copy() + out = getattr(mx, op)(a, axis=0) + self.assertTrue(np.array_equal(np.array(out), expected)) + + def test_scans_complex_exclusive(self): + a = mx.array([-3 + 1j, -1 + 2j, -4 + 0j, 0 + 5j, 2 - 1j]) + for op in ("cummax", "cummin", "logcumsumexp"): + mxop = getattr(mx, op) + for reverse in (False, True): + inclusive = mxop(a, axis=0, inclusive=True, reverse=reverse) + exclusive = mxop(a, axis=0, inclusive=False, reverse=reverse) + if reverse: + got, want = exclusive[:-1], inclusive[1:] + else: + got, want = exclusive[1:], inclusive[:-1] + self.assertTrue( + mx.allclose(got, want), + msg=f"{op} reverse={reverse}", + ) + def test_cummax_cummin_nan(self): nan = float("nan") cases = [ @@ -2686,6 +2915,27 @@ y_mx = mx.sort(a, axis=-1) y_np = np.sort(np.array(a), axis=-1) self.assertTrue(np.array_equal(y_np, y_mx))   + # Negative stride on an axis that is not sorted, single and multi block + np.random.seed(0) + for dtype in ("int32", "float32"): + for size in (4, 32769): + with self.subTest(dtype=dtype, size=size): + a_np = np.random.uniform(0, 100, size=(3, size)) + a_np = a_np.astype(getattr(np, dtype)) + a_mx = mx.array(a_np)[::-1, :] + a_np = a_np[::-1, :] + + b_np = np.sort(a_np, axis=-1) + self.assertTrue(np.array_equal(b_np, mx.sort(a_mx, axis=-1))) + + idx = mx.argsort(a_mx, axis=-1) + self.assertTrue( + np.array_equal(b_np, mx.take_along_axis(a_mx, idx, axis=-1)) + ) + + b_mx = mx.partition(a_mx, 1, axis=-1) + self.assertTrue(np.array_equal(b_np[:, 1], np.array(b_mx)[:, 1])) + def test_partition(self): shape = (3, 4, 5) for dtype in ("int32", "float32"): @@ -2945,7 +3195,7 @@ expected = mx.array(np.linspace(0, 1)) self.assertEqualArray(a, expected)   # Test int64 dtype - b = mx.linspace(0, 10, 5, mx.int64) + b = mx.linspace(0, 10, 5, dtype=mx.int64) expected = mx.array(np.linspace(0, 10, 5, dtype=int)) self.assertEqualArray(b, expected)   @@ -2971,6 +3221,57 @@ for (a, b), n in zip(ranges, nums): d = mx.linspace(a, b, n).tolist() self.assertEqual(d[0], a) self.assertEqual(d[-1], b) + + def test_linspace_endpoint(self): + # endpoint=True is the default and matches the old behaviour + a = mx.linspace(0, 1, 5, endpoint=True) + self.assertEqualArray(a, mx.array(np.linspace(0, 1, 5, endpoint=True))) + self.assertEqualArray(a, mx.linspace(0, 1, 5)) + + # endpoint=False drops the stop value and uses a step of + # (stop - start) / num instead of (stop - start) / (num - 1) + for num in [0, 1, 2, 5, 50]: + b = mx.linspace(0, 10, num, endpoint=False) + expected = mx.array(np.linspace(0, 10, num, endpoint=False)) + self.assertEqualArray(b, expected) + + c = mx.linspace(-2.7, -0.7, 7, endpoint=False) + self.assertEqualArray(c, mx.array(np.linspace(-2.7, -0.7, 7, endpoint=False))) + + # endpoint is the fourth positional argument, before dtype, as in numpy + self.assertEqualArray( + mx.linspace(0, 10, 5, False), mx.array(np.linspace(0, 10, 5, False)) + ) + + # dtype still applies + d = mx.linspace(0, 10, 5, False, mx.int64) + self.assertEqual(d.dtype, mx.int64) + self.assertEqualArray( + d, mx.array(np.linspace(0, 10, 5, endpoint=False, dtype=int)) + ) + + # the start is kept and the stop is excluded + e = mx.linspace(3.0, 4.0, 4, endpoint=False).tolist() + self.assertEqual(e[0], 3.0) + self.assertNotIn(4.0, e) + + # decreasing ranges drop the stop value too + f = mx.linspace(10, 0, 5, endpoint=False) + self.assertEqualArray(f, mx.array(np.linspace(10, 0, 5, endpoint=False))) + + # start == stop keeps every sample at that value + g = mx.linspace(5, 5, 4, endpoint=False) + self.assertEqualArray(g, mx.array(np.linspace(5, 5, 4, endpoint=False))) + + # integer dtype truncates fractional steps, as in numpy + h = mx.linspace(0, 10, 3, endpoint=False, dtype=mx.int32) + self.assertEqualArray( + h, mx.array(np.linspace(0, 10, 3, endpoint=False, dtype=np.int32)) + ) + + # num must still be non-negative + with self.assertRaises(ValueError): + mx.linspace(0, 1, -1, endpoint=False)   def test_repeat(self): # Setup data for the tests @@ -3096,6 +3397,22 @@ mx_out = mx.divmod(mx.array(a_np), mx.array(b_np)) self.assertTrue( np.allclose(np_out[0], mx_out[0]), msg=f"Shapes {s1} {s2}, Type {t}" ) + + # Mixed signs floor, matching python's divmod and numpy, so + # q * b + r == a holds + av = [-7, 7, -7, 7, -1, 1, -5, 5, 6, -6] + bv = [2, 2, -2, -2, 3, -3, 3, -3, 3, 3] + a, b = mx.array(av), mx.array(bv) + q, r = mx.divmod(a, b) + self.assertEqual(q.tolist(), [x // y for x, y in zip(av, bv)]) + self.assertEqual(r.tolist(), [x % y for x, y in zip(av, bv)]) + self.assertTrue(mx.array_equal(q * b + r, a)) + + af = mx.array([-7.0, 7.0, -7.5, 7.5]) + bf = mx.array([2.0, -2.0, 2.0, -2.0]) + q, r = mx.divmod(af, bf) + self.assertTrue(mx.array_equal(q, mx.array([-4.0, -4.0, -4.0, -4.0]))) + self.assertTrue(mx.array_equal(q * bf + r, af))   def test_tile(self): self.assertCmpNumpy([(2,), [2]], mx.tile, np.tile)
diff --git ml-explore/mlx/python/tests/test_quantized.py Layr-Labs/mlx/python/tests/test_quantized.py index 15bc892bd84ceeef5d78eb424c98032a2c9cdad0..461175f013e10620c14b58b47a255cf690b44a35 100644 --- ml-explore/mlx/python/tests/test_quantized.py +++ Layr-Labs/mlx/python/tests/test_quantized.py @@ -1,5 +1,6 @@ # Copyright © 2023-2026 Apple Inc.   +import math import os import platform import subprocess @@ -37,6 +38,18 @@ for b in [2, 3, 4, 5, 6, 8]: w_q, scales, biases = mx.quantize(a, gs, b) a_hat = mx.dequantize(w_q, scales, biases, gs, b) self.assertTrue(mx.all(a_hat == 0)) + + # slices + if mx.default_device() == mx.gpu: + w = mx.random.normal(shape=(2, 256, 32)) + quant = {"group_size": 32, "bits": 4} + wq, scales, biases = mx.quantize(w, **quant) + wq_s = wq[:, :16, :] + scales_s = scales[:, :16, :] + biases_s = biases[:, :16, :] + dq_cpu = mx.dequantize(wq_s, scales_s, biases_s, **quant, stream=mx.cpu) + dq_gpu = mx.dequantize(wq_s, scales_s, biases_s, **quant, stream=mx.gpu) + self.assertTrue(mx.abs(dq_cpu - dq_gpu).max().item() < 1e-6)   def test_mxfp4_quantize_dequantize(self): lut = mx.array( @@ -120,6 +133,28 @@ a = mx.zeros((256, 512)) w_q, scales = mx.quantize(a, mode="mxfp8") w_hat = mx.dequantize(w_q, scales, mode="mxfp8") self.assertTrue(mx.all(w_hat == 0)) + + def test_mxfp8_block_scale_does_not_saturate(self): + # E4M3 has three mantissa bits, so an in-range element loses at most + # half a step, 6.25%. More than that means the block scale rounded + # below amax/448 and the block maximum saturated. + mx.random.seed(0) + group_size = 32 + n_blocks = 512 + + # Sweep the block magnitude across one binade so both scale rounding + # directions are covered. + w = mx.random.normal(shape=(n_blocks, group_size)) + w = w * mx.exp(mx.arange(n_blocks) / n_blocks * math.log(2.0)).reshape(-1, 1) + + w_q, scales = mx.quantize(w, group_size=group_size, mode="mxfp8") + w_hat = mx.dequantize(w_q, scales, group_size=group_size, mode="mxfp8") + + # Quantization is monotone in |w|, so a block's largest output is the + # reconstruction of its largest input. + amax = mx.max(mx.abs(w), axis=1) + rel = mx.abs(amax - mx.max(mx.abs(w_hat), axis=1)) / amax + self.assertLess(mx.max(rel).item(), 0.0626)   def test_nvfp4_quantize_dequantize(self): lut = mx.array( @@ -335,6 +370,7 @@ group_size, bits = 64, 4 K = 128 tests = [ (16, 32840), # unaligned N > 2**15, M < 32: partial M-tile + (32, 32840), # M at the small-block dispatch boundary (33, 32840), # unaligned N > 2**15, M % 32 != 0 (33000, 64), # M > 2**15: row distance overflows (aligned N) ] @@ -439,6 +475,42 @@ check_affine(M, K, 128, 32, bits, dtype) for mode in modes: with self.subTest(M=M, K=K, mode=mode, dtype=dtype): check_fp(M, K, 128, mode, dtype) + + def test_qmm_small_m_block(self): + # The batched and fp-mode variants of the small-M block, which the + # test_qmm_large_dims shapes cannot reach. + if mx.default_device() == mx.cpu: + self.skipTest("Covers GPU kernels only") + key = mx.random.key(0) + k1, k2 = mx.random.split(key) + K = 1024 + tests = [ + # mode, group_size, bits, M, N, batch + ("affine", 64, 4, 14, 8256, (2,)), # batched w + ("mxfp4", None, None, 14, 8256, ()), + ] + for mode, group_size, bits, M, N, batch in tests: + dtype = mx.float16 if mode == "affine" else mx.bfloat16 + with self.subTest( + mode=mode, group_size=group_size, bits=bits, M=M, N=N, batch=batch + ): + x = (mx.random.normal(batch + (M, K), key=k1) / K**0.5).astype(dtype) + w = (mx.random.normal(batch + (N, K), key=k2) / K**0.5).astype(dtype) + if mode == "affine": + wq = mx.quantize(w, group_size=group_size, bits=bits) + else: + wq = mx.quantize(w, mode=mode) + w_hat = mx.dequantize(*wq, group_size=group_size, bits=bits, mode=mode) + y_ref = x @ w_hat.swapaxes(-1, -2) + y = mx.quantized_matmul( + x, + *wq, + transpose=True, + group_size=group_size, + bits=bits, + mode=mode, + ) + self.assertLess((y_ref - y).abs().max(), 1e-3)   def test_qmm_vjp(self): key = mx.random.key(0)
diff --git ml-explore/mlx/python/tests/test_reduce.py Layr-Labs/mlx/python/tests/test_reduce.py index 6ac8fc1504ea0ae4e50a8b097a1670282a5fbeee..164e2dd803d39995fdf77e3786cab67b6dfb27ce 100644 --- ml-explore/mlx/python/tests/test_reduce.py +++ Layr-Labs/mlx/python/tests/test_reduce.py @@ -1,6 +1,5 @@ # Copyright © 2023 Apple Inc.   -import unittest from itertools import combinations, permutations   import mlx.core as mx @@ -46,6 +45,16 @@ mx.eval(z_mlx) self.assertTrue( np.allclose(z_npy, np.array(z_mlx), atol=1e-4) ) + + def test_row_reduce_negative_stride(self): + x_npy = np.arange(1, 131).reshape(2, 65)[::-1] + x_mlx = mx.arange(1, 131).reshape(2, 65)[::-1] + + for op in ["sum", "max", "min", "mean", "var"]: + with self.subTest(op=op): + expected = getattr(np, op)(x_npy, axis=-1) + actual = getattr(mx, op)(x_mlx, axis=-1) + self.assertTrue(np.allclose(expected, actual))   def test_dtypes(self): int_dtypes = [
diff --git ml-explore/mlx/python/tests/test_vmap.py Layr-Labs/mlx/python/tests/test_vmap.py index 99c30a2dc23b53cbd3ac3db2dcce7dd5ac69f708..b050a803840d92d95b8111889cbb17c43f235ab8 100644 --- ml-explore/mlx/python/tests/test_vmap.py +++ Layr-Labs/mlx/python/tests/test_vmap.py @@ -252,6 +252,85 @@ out = mx.vmap(lambda x: mx.argmax(x))(a) expected = mx.array([2, 1]) self.assertTrue(mx.array_equal(out, expected))   + def _unstack(self, x, axis): + return [s.squeeze(axis) for s in mx.split(x, x.shape[axis], axis=axis)] + + def test_vmap_partition(self): + # Distinct values so each lane has a single valid kth element + a = mx.random.permutation(2 * 3 * 4).reshape(2, 3, 4).astype(mx.float32) + + for in_axis in (0, 1, 2): + slices = self._unstack(a, in_axis) + # Axis of the batched output that the inner axis maps onto + out_axes_map = [d for d in range(a.ndim) if d != in_axis] + for axis in (0, 1, -1): + oaxis = out_axes_map[axis if axis >= 0 else axis + 2] + for kth in range(slices[0].shape[axis]): + expected = mx.stack( + [mx.partition(x, kth, axis=axis) for x in slices], + axis=in_axis, + ) + pivot = mx.take(expected, mx.array([kth]), axis=oaxis) + + out = mx.vmap( + lambda x: mx.partition(x, kth, axis=axis), + in_axes=in_axis, + out_axes=in_axis, + )(a) + self.assertEqual(out.shape, expected.shape) + # partition only pins the kth element; the two sides are + # an arbitrary permutation, so compare against the sorted + # input rather than element-wise. + self.assertTrue( + mx.array_equal(mx.sort(out, axis=oaxis), mx.sort(a, axis=oaxis)) + ) + self.assertTrue( + mx.array_equal(mx.take(out, mx.array([kth]), axis=oaxis), pivot) + ) + + idx = mx.vmap( + lambda x: mx.argpartition(x, kth, axis=axis), + in_axes=in_axis, + out_axes=in_axis, + )(a) + self.assertEqual(idx.shape, expected.shape) + gathered = mx.take_along_axis(a, idx, axis=oaxis) + self.assertTrue( + mx.array_equal( + mx.sort(gathered, axis=oaxis), mx.sort(a, axis=oaxis) + ) + ) + self.assertTrue( + mx.array_equal( + mx.take(gathered, mx.array([kth]), axis=oaxis), pivot + ) + ) + + def test_vmap_topk(self): + a = mx.random.permutation(2 * 3 * 4).reshape(2, 3, 4).astype(mx.float32) + + for in_axis in (0, 1, 2): + slices = self._unstack(a, in_axis) + out_axes_map = [d for d in range(a.ndim) if d != in_axis] + for axis in (0, 1, -1): + oaxis = out_axes_map[axis if axis >= 0 else axis + 2] + for k in range(1, slices[0].shape[axis] + 1): + out = mx.vmap( + lambda x: mx.topk(x, k, axis=axis), + in_axes=in_axis, + out_axes=in_axis, + )(a) + expected = mx.stack( + [mx.topk(x, k, axis=axis) for x in slices], axis=in_axis + ) + self.assertEqual(out.shape, expected.shape) + # topk does not promise an order within the k elements + self.assertTrue( + mx.array_equal( + mx.sort(out, axis=oaxis), mx.sort(expected, axis=oaxis) + ) + ) + def test_vmap_mean(self): a = mx.arange(8).reshape(2, 4) out = mx.vmap(mx.mean)(a) @@ -943,6 +1022,20 @@ expected = vmap_fn(z, w) out = cvmap_fn(z, w) self.assertTrue(mx.array_equal(expected, out)) self.assertEqual(6, counter[0]) + + def test_vmap_sort(self): + a = mx.random.uniform(shape=(3, 5)) + expected = mx.stack([mx.sort(a[:, i]) for i in range(a.shape[1])], axis=1) + for axis in (0, -1): + out = mx.vmap(lambda x: mx.sort(x, axis=axis), in_axes=1, out_axes=1)(a) + self.assertTrue(mx.array_equal(out, expected)) + + def test_vmap_argsort(self): + a = mx.random.uniform(shape=(3, 5)) + expected = mx.stack([mx.argsort(a[:, i]) for i in range(a.shape[1])], axis=1) + for axis in (0, -1): + out = mx.vmap(lambda x: mx.argsort(x, axis=axis), in_axes=1, out_axes=1)(a) + self.assertTrue(mx.array_equal(out, expected))   if __name__ == "__main__":
diff --git ml-explore/mlx/setup.py Layr-Labs/mlx/setup.py index 3c2f1380487d372bfd29efad27d0cea2c38f405d..52cc4f7492124ca6d2385d519ea3f9ec820a39bd 100644 --- ml-explore/mlx/setup.py +++ Layr-Labs/mlx/setup.py @@ -53,7 +53,22 @@ return version   -build_stage = int(os.environ.get("MLX_BUILD_STAGE", 0)) +# Release builds for PyPi are separated into 2 packages: +# +# Frontend package: +# - Triggered with `MLX_BUILD_FRONTEND_PACKAGE=1` +# - Include everything except backend-specific binaries (e.g. libmlx.so, mlx.metallib, etc) +# - Wheel has Python ABI and platform tags +# - Wheel should be built for the cross-product of python version and platforms +# - Package name is "mlx" and it depends on backend packages (e.g. mlx-metal, mlx-cuda) +# Backend package: +# - Triggered with `MLX_BUILD_BACKEND_PACKAGE=1` +# - Include headers and backend binaries. +# - Wheel has only platform tags +# - Wheel should be built only for different platforms +# - Package name is back-end specific, e.g mlx-metal, mlx-cuda +build_frontend = int(os.environ.get("MLX_BUILD_FRONTEND_PACKAGE", 0)) +build_backend = int(os.environ.get("MLX_BUILD_BACKEND_PACKAGE", 0)) build_macos = platform.system() == "Darwin" build_cuda = "MLX_BUILD_CUDA=ON" in os.environ.get("CMAKE_ARGS", "")   @@ -77,9 +92,7 @@ if platform.system() == "Windows": self.build_temp = os.path.dirname(self.build_temp)   def build_extension(self, ext: CMakeExtension) -> None: - # Must be in this form due to bug in .resolve() only fixed in Python 3.10+ - ext_fullpath = Path.cwd() / self.get_ext_fullpath(ext.name) # type: ignore[no-untyped-call] - extdir = ext_fullpath.parent.resolve() + extdir = self._get_ext_dir(ext)   debug = int(os.environ.get("DEBUG", 0)) if self.debug is None else self.debug cfg = "Debug" if debug else "Release" @@ -88,17 +101,9 @@ build_temp = Path(self.build_temp) / ext.name if not build_temp.exists(): build_temp.mkdir(parents=True)   - install_prefix = extdir - pybind_out_dir = extdir - if build_stage == 1: - # Don't include MLX libraries in the wheel - install_prefix = build_temp - elif build_stage == 2: - # Don't include Python bindings in the wheel - pybind_out_dir = build_temp cmake_args = [ - f"-DCMAKE_INSTALL_PREFIX={install_prefix}", - f"-DMLX_PYTHON_BINDINGS_OUTPUT_DIRECTORY={pybind_out_dir}", + f"-DCMAKE_INSTALL_PREFIX={extdir}", + f"-DMLX_PYTHON_BINDINGS_OUTPUT_DIRECTORY={extdir}", f"-DCMAKE_BUILD_TYPE={cfg}", f"-DPython_EXECUTABLE={sys.executable}", "-DMLX_BUILD_PYTHON_BINDINGS=ON", @@ -113,8 +118,7 @@ # (needed e.g. to build for ARM OSx on conda-forge) if "CMAKE_ARGS" in os.environ: cmake_args += [item for item in os.environ["CMAKE_ARGS"].split(" ") if item]   - # For release wheel force building for all supported arches. - if build_stage == 2 and build_cuda: + if build_backend and build_cuda: # Last arch is always real and virtual for forward-compatibility cuda_archs = [ "75-real", @@ -186,16 +190,50 @@ subprocess.run( ["cmake", "--install", build_temp, "--component", "core_stub"], check=True, ) + # Copy the type stubs to extdir so they are included in wheels. + stubs_dir = Path("python/mlx/core") + if stubs_dir.exists(): + extdir = self._get_ext_dir(ext) + self.copy_tree(stubs_dir, extdir / "core") + + def _get_ext_dir(self, ext): + # Must be in this form due to bug in .resolve() only fixed in Python 3.10+ + ext_fullpath = Path.cwd() / self.get_ext_fullpath(ext.name) # type: ignore[no-untyped-call] + return ext_fullpath.parent.resolve()   class MLXBdistWheel(bdist_wheel): def get_tag(self) -> tuple[str, str, str]: impl, abi, plat_name = super().get_tag() - if build_stage == 2: + if build_backend: impl = self.python_tag abi = "none" return (impl, abi, plat_name)   + def write_wheelfile(self, *args, **kwargs) -> None: + super().write_wheelfile(*args, **kwargs) + + mlx_dir = Path(self.bdist_dir, "mlx") + + def is_backend_file(file): + if file.is_relative_to(Path(mlx_dir, "lib")): + return True + if file.is_relative_to(Path(mlx_dir, "include")): + return True + if file.is_relative_to(Path(mlx_dir, "share")): + return True + if file.suffix == ".dll": + return True + return False + + if build_frontend or build_backend: + for file in Path(self.bdist_dir).rglob("*"): + if not file.is_relative_to(mlx_dir) or not file.is_file(): + continue + bf = is_backend_file(file) + if (build_frontend and bf) or (build_backend and not bf): + file.unlink() +   # Read the content of README.md with open(Path(__file__).parent / "README.md", encoding="utf-8") as f: @@ -260,24 +298,8 @@ ] } install_requires = []   - # Release builds for PyPi are in two stages. - # Each stage should be run from a clean build: - # python setup.py clean --all - # - # Stage 1: - # - Triggered with `MLX_BUILD_STAGE=1` - # - Include everything except backend-specific binaries (e.g. libmlx.so, mlx.metallib, etc) - # - Wheel has Python ABI and platform tags - # - Wheel should be built for the cross-product of python version and platforms - # - Package name is mlx and it depends on subpackage in stage 2 (e.g. mlx-metal) - # Stage 2: - # - Triggered with `MLX_BUILD_STAGE=2` - # - Includes only backend-specific binaries (e.g. libmlx.so, mlx.metallib, etc) - # - Wheel has only platform tags - # - Wheel should be built only for different platforms - # - Package name is back-end specific, e.g mlx-metal - if build_stage != 2: - if build_stage == 1: + if not build_backend: + if build_frontend: install_requires.append( f'mlx-metal=={version}; platform_system == "Darwin"' ) @@ -320,9 +342,10 @@ "nvidia-cuda-nvrtc-cu12==12.9.*", ] elif toolkit == 13: install_requires += [ - "nvidia-cublas", - "nvidia-cufft", - "nvidia-cuda-nvrtc", + "nvidia-cublas==13.*", + "nvidia-cufft==12.*", + "nvidia-cuda-nvrtc==13.*", + "nvidia-cuda-runtime==13.*", ] else: raise ValueError(f"Unknown toolkit {toolkit}")
diff --git ml-explore/mlx/tests/load_tests.cpp Layr-Labs/mlx/tests/load_tests.cpp index 89749194760373f7a688390c1fecf458bbe18de4..6ef7bc276e1d3093b2107a77292d67087caa1bdd 100644 --- ml-explore/mlx/tests/load_tests.cpp +++ Layr-Labs/mlx/tests/load_tests.cpp @@ -257,6 +257,120 @@ CHECK_THROWS_AS(load_gguf(file_path), std::runtime_error); } }   +// Writes a metadata-only GGUF (no tensors) whose metadata KV section is +// `kv_section` verbatim, so a caller can encode values whose lengths exceed the +// file to exercise check_metadata_value_in_file(). `kv_count` must match the +// number of KV pairs encoded in `kv_section`. +void write_raw_gguf_metadata( + const std::string& path, + uint64_t kv_count, + const std::vector<char>& kv_section) { + std::ofstream out(path, std::ios::binary); + auto u32 = [&out](uint32_t v) { + out.write(reinterpret_cast<const char*>(&v), 4); + }; + auto u64 = [&out](uint64_t v) { + out.write(reinterpret_cast<const char*>(&v), 8); + }; + out.write("GGUF", 4); + u32(3); // version + u64(0); // tensor_count + u64(kv_count); // metadata_kv_count + out.write(kv_section.data(), kv_section.size()); +} + +TEST_CASE("test gguf metadata value validation") { + // A STRING/ARRAY metadata value claiming a length larger than the file must + // be rejected rather than read past the end of the mapping. See PR #4212. + + auto append_string_kv = [](std::vector<char>& b, + const std::string& key, + uint64_t claimed_len, + bool write_payload) { + auto put = [&](const void* p, size_t n) { + b.insert( + b.end(), + static_cast<const char*>(p), + static_cast<const char*>(p) + n); + }; + uint64_t klen = key.size(); + put(&klen, 8); + put(key.data(), key.size()); + uint32_t vt = 8; // GGUF_VALUE_TYPE_STRING + put(&vt, 4); + put(&claimed_len, 8); + if (write_payload) { + b.insert(b.end(), claimed_len, '\0'); + } + }; + + auto append_array_kv = [](std::vector<char>& b, + const std::string& key, + uint32_t elt_type, + uint64_t claimed_len) { + auto put = [&](const void* p, size_t n) { + b.insert( + b.end(), + static_cast<const char*>(p), + static_cast<const char*>(p) + n); + }; + uint64_t klen = key.size(); + put(&klen, 8); + put(key.data(), key.size()); + uint32_t vt = 9; // GGUF_VALUE_TYPE_ARRAY + put(&vt, 4); + put(&elt_type, 4); + put(&claimed_len, 8); + }; + + SUBCASE("valid empty and small strings load") { + std::vector<char> kv; + append_string_kv(kv, "empty", 0, false); + append_string_kv(kv, "small", 5, true); + std::string file_path = get_temp_file("test_gguf_meta_ok.gguf"); + write_raw_gguf_metadata(file_path, 2, kv); + auto [weights, metadata] = load_gguf(file_path); + CHECK(weights.empty()); + CHECK(std::get<std::string>(metadata.at("empty")) == ""); + CHECK(std::get<std::string>(metadata.at("small")) == std::string(5, '\0')); + } + + SUBCASE("string length extends past the end of the file") { + // Claims 100 bytes of payload, none of which are present. + std::vector<char> kv; + append_string_kv(kv, "s", 100, false); + std::string file_path = get_temp_file("test_gguf_meta_str_past.gguf"); + write_raw_gguf_metadata(file_path, 1, kv); + CHECK_THROWS_AS(load_gguf(file_path), std::runtime_error); + } + + SUBCASE("string length far past the end of the file") { + std::vector<char> kv; + append_string_kv(kv, "s", 1ull << 40, false); + std::string file_path = get_temp_file("test_gguf_meta_str_far.gguf"); + write_raw_gguf_metadata(file_path, 1, kv); + CHECK_THROWS_AS(load_gguf(file_path), std::runtime_error); + } + + SUBCASE("fixed-size array length extends past the end of the file") { + // GGUF_VALUE_TYPE_UINT8 = 0; claims 2^40 elements, none present. + std::vector<char> kv; + append_array_kv(kv, "a", 0, 1ull << 40); + std::string file_path = get_temp_file("test_gguf_meta_arr_past.gguf"); + write_raw_gguf_metadata(file_path, 1, kv); + CHECK_THROWS_AS(load_gguf(file_path), std::runtime_error); + } + + SUBCASE("string array element length extends past the end of the file") { + // GGUF_VALUE_TYPE_STRING = 8; two elements, neither present. + std::vector<char> kv; + append_array_kv(kv, "a", 8, 2); + std::string file_path = get_temp_file("test_gguf_meta_strarr_past.gguf"); + write_raw_gguf_metadata(file_path, 1, kv); + CHECK_THROWS_AS(load_gguf(file_path), std::runtime_error); + } +} + TEST_CASE("test gguf metadata") { std::string file_path = get_temp_file("test_arr.gguf"); using dict = std::unordered_map<std::string, array>;
diff --git ml-explore/mlx/tests/ops_tests.cpp Layr-Labs/mlx/tests/ops_tests.cpp index f7a2b8ab921268bcfadd581d791f07bd79a8982a..09236f1da1d93510d8842061f61df8f434efddd7 100644 --- ml-explore/mlx/tests/ops_tests.cpp +++ Layr-Labs/mlx/tests/ops_tests.cpp @@ -3348,11 +3348,27 @@ auto x = linspace(0, 10, 5); auto expected = array({0.0f, 2.5f, 5.0f, 7.5f, 10.0f}, {5}); CHECK(array_equal(x, expected).item<bool>());   - x = linspace(0, 10, 5, int32); + x = linspace(0, 10, 5, true, int32); expected = array({0, 2, 5, 7, 10}, {5}); CHECK(array_equal(x, expected).item<bool>());   x = linspace(0, 1, 0); + expected = array(std::initializer_list<float>{}, {0}); + CHECK(array_equal(x, expected).item<bool>()); + + x = linspace(0, 10, 5, false); + expected = array({0.0f, 2.0f, 4.0f, 6.0f, 8.0f}, {5}); + CHECK(array_equal(x, expected).item<bool>()); + + x = linspace(0, 10, 5, false, int32); + expected = array({0, 2, 4, 6, 8}, {5}); + CHECK(array_equal(x, expected).item<bool>()); + + x = linspace(1, 10, 1, false); + expected = array({1.0f}, {1}); + CHECK(array_equal(x, expected).item<bool>()); + + x = linspace(0, 1, 0, false); expected = array(std::initializer_list<float>{}, {0}); CHECK(array_equal(x, expected).item<bool>()); } @@ -4473,6 +4489,14 @@ Shape{1, 6, 6, 1}); CHECK_EQ( conv_transpose2d(in_t, wt, {2, 2}, {1, 1}, {1, 1}, {1, 1}).shape(), Shape{1, 8, 8, 1}); +} + +TEST_CASE("test pad shape overflow") { + // A padding sum that overflows int32 is rejected, not wrapped. + // https://github.com/ml-explore/mlx/issues/3611 + const int imax = 2147483647; + CHECK_THROWS_AS( + pad(zeros({8}), {0}, Shape{imax}, Shape{imax}), std::overflow_error); }   TEST_CASE("test fp8 conversion") {

The fork deletes the upstream workflows that cannot run here: the PyPI release build and the documentation deploy. build_and_test.yml keeps the lint, sanitizer and Fedora jobs. The fork deletes the jobs that run only in the upstream repository, and the composite actions that only the deleted workflows and jobs used.

The gate that keeps this page honest: check_forkdiff.py fails CI when base.hash is not the merge-base with upstream, when a file the fork changes is not described by a section above, or when a section names files the fork no longer changes.

fork-tests.yml builds the fork with Metal and runs its C++ and Python tests on a GitHub-hosted macOS runner. The setup action pins its actions to full commit SHAs, because the organization blocks actions that are not pinned. The organization also blocks hendrikmuhs/ccache-action, so setup installs ccache with Homebrew and keeps the ccache files with actions/cache.

diff --git ml-explore/mlx/.github/actions/build-docs/action.yml Layr-Labs/mlx/.github/actions/build-docs/action.yml deleted file mode 100644 index 4d4286e3c86292b903eef16d1dc22c2689e300d1..0000000000000000000000000000000000000000 --- ml-explore/mlx/.github/actions/build-docs/action.yml +++ /dev/null @@ -1,42 +0,0 @@ -name: 'Build Documentation' -description: 'Build documentation' - -runs: - using: "composite" - steps: - - name: Setup machine - id: setup - uses: ./.github/actions/setup - with: - ccache-key: 'release' - ccache-save: false - - - name: Install dependencies - shell: bash - env: - CMAKE_ARGS: ${{ steps.setup.outputs.cmake-args }} - run: | - sudo apt-get install -y doxygen - uv pip install -r docs/requirements.txt - uv pip install . -v - - - name: Build documentation - shell: bash - run: | - cd docs - doxygen - make html O=-W - - - name: Create artifact tar - shell: bash - run: tar -cf artifact.tar -C docs --dereference build/html index.html - - # Do it manually because upload-pages-artifact requires gtar - - name: Upload artifact - id: upload-artifact - uses: actions/upload-artifact@v7 - with: - name: github-pages - path: artifact.tar - retention-days: 1 - if-no-files-found: error
diff --git ml-explore/mlx/.github/actions/build/action.yml Layr-Labs/mlx/.github/actions/build/action.yml deleted file mode 100644 index 5756f6e1461543c3e9a3f935cfff7b54a87d9ac1..0000000000000000000000000000000000000000 --- ml-explore/mlx/.github/actions/build/action.yml +++ /dev/null @@ -1,44 +0,0 @@ -name: 'Build' -description: 'Build C++ and Python binaries for testing on Linux and Windows' - -inputs: - cmake-args: - description: 'The args for generating CMake project' - required: true - debug: - description: 'Do debug build' - required: true - -runs: - using: 'composite' - steps: - - name: Install Python package - shell: bash - env: - DEBUG: ${{ inputs.debug == 'true' && 1 || 0 }} - CMAKE_ARGS: ${{ inputs.cmake-args }} - run: | - echo "::group::Install Python package" - # Install cpu-only torch to save space. - uv pip install torch --torch-backend=cpu - uv pip install --no-build-isolation -e ".[dev]" -v - echo "::endgroup::" - - - name: Build CPP only - shell: bash - env: - # The cpp build is using some extra settings to reuse the compile cache - # generated by the python install: - # 1. Use the same ccache options with setup.py. - # 2. Build dynamic library. - # 3. Put the build dir in the same depth with python build dir. - CCACHE_BASEDIR: ${{ github.workspace }}/build/cpp/mlx - CCACHE_NOHASHDIR: true - run: | - echo "::group::Build CPP only" - cmake . -B build/cpp/mlx ${{ inputs.cmake-args }} \ - -DBUILD_SHARED_LIBS=ON \ - -DCMAKE_BUILD_TYPE=${{ inputs.debug == 'true' && 'Debug' || 'Release' }} - cmake --build build/cpp/mlx \ - -j ${{ runner.os == 'Windows' && '$NUMBER_OF_PROCESSORS' || '$(nproc)' }} - echo "::endgroup::"
diff --git ml-explore/mlx/.github/actions/test-wheel/action.yml Layr-Labs/mlx/.github/actions/test-wheel/action.yml deleted file mode 100644 index db68579bbb1964ca20ad15cfd82407659fdf0609..0000000000000000000000000000000000000000 --- ml-explore/mlx/.github/actions/test-wheel/action.yml +++ /dev/null @@ -1,51 +0,0 @@ -name: 'Test wheel' -description: 'Run tests with the built wheels' - -inputs: - toolkit: - description: 'Which toolkit to test' - required: false - default: 'cpu' - -runs: - using: 'composite' - steps: - - name: Get Python version - id: python - shell: bash - run: python -c "import sys; print(f'version={sys.version_info.major}.{sys.version_info.minor}')" >> $GITHUB_OUTPUT - - - name: Download frontend packages - uses: actions/download-artifact@v8 - with: - pattern: frontend-${{ runner.os }}-${{ runner.arch }}-py${{ steps.python.outputs.version }} - path: wheelhouse - - - name: Download backend packages - uses: actions/download-artifact@v8 - with: - pattern: backend-${{ inputs.toolkit}}-${{ runner.os }}-${{ runner.arch }} - path: wheelhouse - - - name: Test local packages - shell: bash - run: | - echo "::group::Test local packages" - # Fall back to PyPI if PyTorch index has issues. - uv pip install torch --torch-backend=cpu - uv pip install numpy - if ${{ inputs.toolkit == 'cpu' }} ; then - uv pip install wheelhouse/mlx_cpu*.whl - uv pip install wheelhouse/mlx-*.whl - elif ${{ startsWith(inputs.toolkit, 'cuda') }} ; then - uv pip install wheelhouse/mlx_cuda*.whl - uv pip install wheelhouse/mlx-*.whl - elif ${{ inputs.toolkit == 'metal' }} ; then - uv pip install wheelhouse/mlx_metal-*-macosx_26_0_arm64.whl - uv pip install wheelhouse/mlx-*-macosx_26_0_arm64.whl - else - echo "No matching backend wheel to install" - exit 1 - fi - python -m unittest discover -v python/tests - echo "::endgroup::"
diff --git ml-explore/mlx/.github/actions/test-windows/action.yml Layr-Labs/mlx/.github/actions/test-windows/action.yml deleted file mode 100644 index ac812d8af9760388eb37caded69665f5ae81a6d5..0000000000000000000000000000000000000000 --- ml-explore/mlx/.github/actions/test-windows/action.yml +++ /dev/null @@ -1,23 +0,0 @@ -name: 'Run tests' -description: 'Run Python and C++ tests on Windows' - -runs: - using: 'composite' - steps: - - name: Run Python tests - CPU - shell: bash - env: - DEVICE: cpu - run: | - echo "::group::Python tests - CPU" - python -m unittest discover python/tests -v - echo "::endgroup::" - - - name: Run CPP tests - CPU - shell: bash - env: - DEVICE: cpu - run: | - echo "::group::CPP tests - CPU" - ./build/cpp/mlx/tests.exe -tce="*gguf*" - echo "::endgroup::"
diff --git ml-explore/mlx/.github/workflows/build_and_test.yml Layr-Labs/mlx/.github/workflows/build_and_test.yml index 6de4218b2c0afa7f7b63a1bfa0c6089226196145..0bcffcd79ff0a5ff64090a9f4f5603bcb01e85c4 100644 --- ml-explore/mlx/.github/workflows/build_and_test.yml +++ Layr-Labs/mlx/.github/workflows/build_and_test.yml @@ -1,6 +1,7 @@ name: Build and Test   on: + workflow_dispatch: pull_request: push: branches: @@ -20,124 +21,17 @@ check_lint: name: Check Lint runs-on: ubuntu-22.04 steps: - - uses: actions/checkout@v7 - - uses: pre-commit/action@v3.0.1 - - build_and_test: - name: ${{ matrix.os }} (${{ matrix.toolkit }}, ${{ matrix.arch }}) - if: github.repository == 'ml-explore/mlx' - needs: check_lint - strategy: - fail-fast: false - matrix: - os: ['Linux', 'Windows'] - arch: ['x86_64', 'aarch64'] - toolkit: ['cpu', 'cuda-12.6', 'cuda-12.9', 'cuda-13.0'] - exclude: - # CUDA does not support Windows on arm. - - os: 'Windows' - arch: 'aarch64' - toolkit: 'cuda-12.6' - - os: 'Windows' - arch: 'aarch64' - toolkit: 'cuda-12.9' - - os: 'Windows' - arch: 'aarch64' - toolkit: 'cuda-13.0' - # CUDA 12.6 does not compile with CUTLASS on Windows. - - os: 'Windows' - arch: 'x86_64' - toolkit: 'cuda-12.6' - runs-on: |- - ${{ case(matrix.os == 'Windows', - case(matrix.arch == 'aarch64', 'windows-11-arm', - startsWith(matrix.toolkit, 'cuda'), 'windows-2022-large', - 'windows-2022'), - case(matrix.arch == 'x86_64' && startsWith(matrix.toolkit, 'cuda'), 'gpu-t4-4-core', - matrix.arch == 'aarch64', 'ubuntu-22.04-arm', - 'ubuntu-22.04')) - }} - steps: - - uses: actions/checkout@v7 - - uses: ./.github/actions/setup - id: setup - with: - toolkit: ${{ matrix.toolkit }} - - uses: ./.github/actions/build - with: - cmake-args: ${{ steps.setup.outputs.cmake-args }} - # For MSVC, Ninja/Release is the only config supported by ccache. - debug: ${{ matrix.os != 'Windows' }} - - uses: ./.github/actions/test-linux - if: matrix.os == 'Linux' && (matrix.toolkit == 'cpu' || matrix.arch == 'x86_64') - - uses: ./.github/actions/test-windows - if: matrix.os == 'Windows' && matrix.toolkit == 'cpu' - - mac_build: - name: macOS (${{ matrix.macos-target }}, ${{ matrix.toolkit }}) - if: github.repository == 'ml-explore/mlx' - strategy: - matrix: - macos-target: ['14.0', '15.0', '26.2'] - toolkit: ['cpu', 'metal', 'jit'] - runs-on: 'macos-26' - needs: check_lint - steps: - - uses: actions/checkout@v7 - - uses: ./.github/actions/setup - id: setup - with: - toolkit: ${{ matrix.toolkit }} - ccache-key: 'test-${{ matrix.macos-target }}' - # Use the compile cache of metal build for all builds. - ccache-save: ${{ matrix.toolkit == 'metal' }} - ccache-toolkit: 'metal' - - uses: ./.github/actions/build-macos - with: - cmake-args: ${{ steps.setup.outputs.cmake-args }} - macos-target: ${{ matrix.macos-target }} - - uses: ./.github/actions/test-macos - if: matrix.toolkit == 'cpu' - with: - toolkit: 'cpu' - - uses: actions/upload-artifact@v7 - if: matrix.toolkit == 'metal' - with: - name: mlx-${{ matrix.toolkit }}-macos${{ matrix.macos-target }} - path: | - dist/mlx-*.whl - build/mlx/backend/metal/kernels/mlx.metallib - build/tests/tests - if-no-files-found: error - - mac_test: - name: Test macOS - if: github.repository == 'ml-explore/mlx' - runs-on: [self-hosted, macos] - needs: mac_build - steps: - - uses: actions/checkout@v7 - - uses: ./.github/actions/setup - with: - toolkit: 'metal' - use-ccache: false - - uses: actions/download-artifact@v8 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7 + - name: Install pre-commit + run: | + python -m pip install pre-commit + python -m pip freeze --local + - uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4 with: - path: artifact - pattern: mlx-metal-* - - run: ls -lhR artifact - - uses: ./.github/actions/test-macos - with: - toolkit: 'metal' - - build_documentation: - name: Build Documentation - if: github.repository == 'ml-explore/mlx' - runs-on: ubuntu-22.04 - needs: check_lint - steps: - - uses: actions/checkout@v7 - - uses: ./.github/actions/build-docs + path: ~/.cache/pre-commit + key: pre-commit-3|${{ env.pythonLocation }}|${{ hashFiles('.pre-commit-config.yaml') }} + - name: Run pre-commit + run: pre-commit run --show-diff-on-failure --color=always --all-files   linux_sanitizer_build_and_test: name: Linux Sanitizer Tests (${{ matrix.sanitizer }}) @@ -151,7 +45,7 @@ # sanitizer: [ASAN, UBSAN, TSAN] runs-on: ubuntu-22.04-arm steps: - name: Checkout code - uses: actions/checkout@v7 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7   - name: Install Dependencies run: | @@ -189,7 +83,7 @@ container: image: fedora:42 steps: - name: Checkout code - uses: actions/checkout@v7 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7   - name: CPP Build Test - No Release run: |
diff --git ml-explore/mlx/.github/workflows/documentation.yml Layr-Labs/mlx/.github/workflows/documentation.yml deleted file mode 100644 index fad1647999807332410edc20de0a6254dd617b5b..0000000000000000000000000000000000000000 --- ml-explore/mlx/.github/workflows/documentation.yml +++ /dev/null @@ -1,28 +0,0 @@ -name: Documentation - -on: - workflow_dispatch: - -permissions: - contents: read - -jobs: - build: - runs-on: ubuntu-22.04 - steps: - - uses: actions/checkout@v7 - - uses: ./.github/actions/build-docs - - deploy: - needs: build - permissions: - pages: write - id-token: write - runs-on: ubuntu-latest - environment: - name: github-pages - url: ${{ steps.deployment.outputs.page_url }} - steps: - - name: Deploy to GitHub Pages - id: deployment - uses: actions/deploy-pages@v5
diff --git ml-explore/mlx/.github/workflows/fork-tests.yml Layr-Labs/mlx/.github/workflows/fork-tests.yml new file mode 100644 index 0000000000000000000000000000000000000000..f28404b459d1d4a160a211215d6169d47da04d14 --- /dev/null +++ Layr-Labs/mlx/.github/workflows/fork-tests.yml @@ -0,0 +1,43 @@ +name: Fork Tests + +# Build the fork with Metal and run its C++ and Python tests on a +# GitHub-hosted macOS runner. The steps use the actions in .github/actions. + +on: + pull_request: + push: + branches: + - main + workflow_dispatch: + +jobs: + metal-build-and-test: + runs-on: macos-26 + timeout-minutes: 60 + permissions: + contents: read + concurrency: + group: ${{ github.workflow }}-${{ github.ref }} + # Cancel an older run of the same pull request. Let runs on main finish. + cancel-in-progress: ${{ github.event_name == 'pull_request' }} + steps: + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 + + - uses: ./.github/actions/setup + id: setup + with: + toolkit: metal + # This cache key is not an upstream key. The setup action saves + # the cache on pushes to main only, never on pull requests. + ccache-key: fork-tests + + - uses: ./.github/actions/build-macos + with: + cmake-args: ${{ steps.setup.outputs.cmake-args }} + # The runner has macOS 26. Version 26.2 is the lowest target that + # builds all Metal kernels. + macos-target: '26.2' + + - uses: ./.github/actions/test-macos + with: + toolkit: metal
diff --git ml-explore/mlx/.github/workflows/release.yml Layr-Labs/mlx/.github/workflows/release.yml deleted file mode 100644 index bd94af697b8ea85df6a188b40a62533d1f614fb7..0000000000000000000000000000000000000000 --- ml-explore/mlx/.github/workflows/release.yml +++ /dev/null @@ -1,317 +0,0 @@ -name: 'Release build' -description: 'Build Python wheels for nightly or offical releases' - -on: - push: - tags: - - 'v*' - workflow_dispatch: - inputs: - publish: - description: 'Publish to PyPI' - required: false - type: boolean - dev-release: - description: 'Development release (DEV_RELEASE=1)' - required: false - type: boolean - schedule: - - cron: 33 6 * * 1-5 - -# In jobs we must use |*publish| instead of |inputs.publish| because we can not -# set default value for workflow_dispatch inputs reliably. -env: - publish: &publish ${{ inputs.publish || github.event_name == 'push' }} - pypi-env: &pypi-env ${{ (inputs.publish || github.event_name == 'push') && 'pypi' || 'dry-run' }} - -permissions: - contents: read - -jobs: - build_documentation: - name: Build documentation - if: github.repository == 'ml-explore/mlx' - runs-on: ubuntu-22.04 - steps: - - uses: actions/checkout@v7 - - uses: ./.github/actions/build-docs - - deploy_documentation: - name: Deploy documentation - if: *publish - needs: build_documentation - permissions: - pages: write - id-token: write - runs-on: ubuntu-latest - environment: - name: github-pages - url: ${{ steps.deployment.outputs.page_url }} - steps: - - name: Deploy to GitHub Pages - id: deployment - uses: actions/deploy-pages@v5 - - build_frontend: - name: ${{ matrix.os }} (python-${{ matrix.python-version }}, ${{ matrix.arch }}) - strategy: - matrix: - os: ['Linux', 'Windows'] - arch: ['x86_64', 'aarch64'] - python-version: ['3.10', '3.11', '3.12', '3.13', '3.14'] - # There is no cp310 binary for Windows on arm. - exclude: - - os: 'Windows' - arch: 'aarch64' - python-version: '3.10' - runs-on: |- - ${{ case(matrix.os == 'Windows', - case(matrix.arch == 'aarch64', 'windows-11-arm', - 'windows-2022'), - case(matrix.arch == 'aarch64', 'ubuntu-22.04-arm', - 'ubuntu-22.04')) - }} - env: &build-env - PYPI_RELEASE: 1 - DEV_RELEASE: ${{ inputs.dev-release && 1 || 0 }} - steps: - - uses: actions/checkout@v7 - - uses: ./.github/actions/setup - id: setup - with: - python-version: ${{ matrix.python-version }} - ccache-key: 'release' - ccache-save: false - - uses: ./.github/actions/build-wheel - with: - cmake-args: ${{ steps.setup.outputs.cmake-args }} - build-backend: false - - uses: actions/upload-artifact@v7 - with: - name: frontend-${{ runner.os }}-${{ runner.arch }}-py${{ matrix.python-version }} - path: wheelhouse/mlx-*.whl - if-no-files-found: error - - build_backend: - name: ${{ matrix.os }} (${{ matrix.toolkit }}, ${{ matrix.arch }}) - if: github.repository == 'ml-explore/mlx' - strategy: - matrix: - os: ['Linux', 'Windows'] - arch: ['x86_64', 'aarch64'] - toolkit: ['cpu', 'cuda-12.9', 'cuda-13.0'] - exclude: - # CUDA does not support Windows on arm. - - os: 'Windows' - arch: 'aarch64' - toolkit: 'cuda-12.9' - - os: 'Windows' - arch: 'aarch64' - toolkit: 'cuda-13.0' - runs-on: |- - ${{ case(matrix.os == 'Windows', - case(matrix.arch == 'aarch64', 'windows-11-arm', - 'windows-2022-large'), - case(matrix.arch == 'aarch64', 'ubuntu-22-large-arm', - 'ubuntu-22-large')) - }} - env: *build-env - steps: - - uses: actions/checkout@v7 - - uses: ./.github/actions/setup - id: setup - with: - toolkit: ${{ matrix.toolkit }} - ccache-key: 'release' - - uses: ./.github/actions/build-wheel - with: - cmake-args: ${{ steps.setup.outputs.cmake-args }} - build-frontend: false - - uses: actions/upload-artifact@v7 - with: - name: backend-${{ matrix.toolkit }}-${{ runner.os }}-${{ runner.arch }} - path: wheelhouse/*.whl - if-no-files-found: ignore - - build_mac_wheels: - name: macOS (python-${{ matrix.python-version }}) - if: github.repository == 'ml-explore/mlx' - strategy: - matrix: - python-version: ['3.10', '3.11', '3.12', '3.13', '3.14'] - runs-on: 'macos-26' - env: *build-env - steps: - - uses: actions/checkout@v7 - - uses: ./.github/actions/setup - id: setup - with: - toolkit: 'metal' - python-version: ${{ matrix.python-version }} - ccache-key: 'release' - ccache-save: ${{ matrix.python-version == '3.10' }} - - name: Build macOS 14 package - uses: ./.github/actions/build-wheel - with: - macos-target: '14.0' - cmake-args: ${{ steps.setup.outputs.cmake-args }} - build-backend: ${{ matrix.python-version == '3.10' }} - - name: Build macOS 15 package - uses: ./.github/actions/build-wheel - with: - macos-target: '15.0' - cmake-args: ${{ steps.setup.outputs.cmake-args }} - build-backend: ${{ matrix.python-version == '3.10' }} - - name: Build macOS 26 package - uses: ./.github/actions/build-wheel - with: - macos-target: '26.2' - cmake-args: ${{ steps.setup.outputs.cmake-args }} - build-backend: ${{ matrix.python-version == '3.10' }} - - name: Upload frontend packages - uses: actions/upload-artifact@v7 - with: - name: frontend-${{ runner.os }}-${{ runner.arch }}-py${{ matrix.python-version }} - path: wheelhouse/mlx-*.whl - if-no-files-found: error - - name: Upload backend packages - if: matrix.python-version == '3.10' - uses: actions/upload-artifact@v7 - with: - name: backend-metal-${{ runner.os }}-${{ runner.arch }} - path: wheelhouse/mlx_metal-*.whl - if-no-files-found: error - - test_wheel: - name: Test (${{ matrix.os }}, ${{ matrix.toolkit }}, ${{ matrix.arch }}) - if: github.repository == 'ml-explore/mlx' - needs: [build_frontend, build_backend, build_mac_wheels] - strategy: - matrix: - os: ['Linux', 'Windows'] - arch: ['aarch64'] - toolkit: ['cpu'] - include: - - os: 'Linux' - arch: 'x86_64' - toolkit: 'cpu' - - os: 'Linux' - arch: 'x86_64' - toolkit: 'cuda-12.9' - - os: 'Linux' - arch: 'x86_64' - toolkit: 'cuda-13.0' - - os: 'Windows' - arch: 'x86_64' - toolkit: 'cpu' - - os: 'macOS' - arch: 'aarch64' - toolkit: 'metal' - runs-on: |- - ${{ case(matrix.os == 'Windows', case(matrix.arch == 'aarch64', 'windows-11-arm', - 'windows-2022'), - matrix.os == 'macOS', fromJson('["self-hosted","macos"]'), - case(matrix.arch == 'x86_64' && startsWith(matrix.toolkit, 'cuda'), 'gpu-t4-4-core', - matrix.arch == 'aarch64', 'ubuntu-22.04-arm', - 'ubuntu-22.04')) - }} - steps: - - uses: actions/checkout@v7 - - uses: ./.github/actions/setup - with: - toolkit: ${{ matrix.toolkit }} - use-ccache: false - - uses: ./.github/actions/test-wheel - with: - toolkit: ${{ matrix.toolkit }} - - pypi-publish-frontend: - name: Publish mlx to PyPI - runs-on: ubuntu-latest - needs: [test_wheel] - permissions: - id-token: write - environment: - name: *pypi-env - url: https://pypi.org/p/mlx - steps: - - uses: actions/download-artifact@v8 - with: - pattern: frontend-* - merge-multiple: true - path: dist - - name: Display structure of downloaded files - run: du -ah dist - - name: Publish package distributions to PyPI - if: *publish - uses: pypa/gh-action-pypi-publish@release/v1 - with: - repository-url: https://upload.pypi.org/legacy/ - - pypi-publish-cuda: - name: Publish mlx-cuda to PyPI - runs-on: ubuntu-latest - needs: [test_wheel] - permissions: - id-token: write - environment: - name: *pypi-env - url: https://pypi.org/p/mlx-cuda - steps: - - uses: actions/download-artifact@v8 - with: - pattern: backend-cuda* - merge-multiple: true - path: dist - - name: Display structure of downloaded files - run: du -ah dist - - name: Publish package distributions to PyPI - if: *publish - uses: pypa/gh-action-pypi-publish@release/v1 - with: - repository-url: https://upload.pypi.org/legacy/ - - pypi-publish-cpu: - name: Publish mlx-cpu to PyPI - runs-on: ubuntu-latest - needs: [test_wheel] - permissions: - id-token: write - environment: - name: *pypi-env - url: https://pypi.org/p/mlx-cpu - steps: - - uses: actions/download-artifact@v8 - with: - pattern: backend-cpu-* - merge-multiple: true - path: dist - - name: Display structure of downloaded files - run: du -ah dist - - name: Publish package distributions to PyPI - if: *publish - uses: pypa/gh-action-pypi-publish@release/v1 - with: - repository-url: https://upload.pypi.org/legacy/ - - pypi-publish-metal: - name: Publish mlx-metal to PyPI - runs-on: ubuntu-latest - needs: [test_wheel] - permissions: - id-token: write - environment: - name: *pypi-env - url: https://pypi.org/p/mlx-metal - steps: - - uses: actions/download-artifact@v8 - with: - pattern: backend-metal-* - path: dist - - name: Display structure of downloaded files - run: du -ah dist - - name: Publish package distributions to PyPI - if: *publish - uses: pypa/gh-action-pypi-publish@release/v1 - with: - repository-url: https://upload.pypi.org/legacy/
diff --git ml-explore/mlx/FORKDIFF.md Layr-Labs/mlx/FORKDIFF.md new file mode 100644 index 0000000000000000000000000000000000000000..5047d49e804e1f0460e05af88abc62adbafda1c5 --- /dev/null +++ Layr-Labs/mlx/FORKDIFF.md @@ -0,0 +1,70 @@ +# Fork diff: what this fork changes, and keeping that page honest + +This repository is a fork of [`ml-explore/mlx`](https://github.com/ml-explore/mlx). +Everything it changes relative to upstream is published as a browsable page: + +**https://layr-labs.github.io/mlx/** + +The page is rendered by [`protolambda/forkdiff`](https://github.com/protolambda/forkdiff) +from [`fork.yaml`](fork.yaml) at the repo root, in the style of +[op-geth's go-ethereum fork diff](https://op-geth.optimism.io/). `fork.yaml` +groups the changed files into sections with a paragraph each, names the exact +upstream commit the fork is based on (`base.hash`), and lists files that are +not code (`ignore`). Every push to `main` re-renders and redeploys it +(`.github/workflows/forkdiff-pages.yml`). + +**One-time setup.** GitHub Pages has to be switched on by a repo admin: +Settings → Pages → Source: **GitHub Actions**. The workflow token cannot do +this itself. Until it's done the deploy workflow still builds the page (kept +as a run artifact, `forkdiff-page`) and exits with a notice instead of failing. + +## The gate + +A fork-diff page is only useful while it is true, and two things make it go +stale silently. `scripts/check_forkdiff.py` runs on every PR and on every push +to `main` (`.github/workflows/forkdiff-check.yml`) and fails on both: + +| Drift | Check | Why it matters | +|-------|-------|----------------| +| **Rebase onto newer upstream** | `base.hash` must equal `git merge-base HEAD upstream/main` and be an ancestor of upstream `main` | After a rebase the old base makes upstream's own commits look like fork changes: thousands of files, none of them ours. | +| **New fork change nobody described** | every path in `git diff --name-only base.hash HEAD` must match a section glob or a global `ignore` | An undescribed file is a change the page can't explain. | +| **Section describing code we no longer carry** | every glob must match at least one changed path | Stale sections are as misleading as missing ones. | + +The check also renders the page, so a `fork.yaml` that forkdiff itself rejects +cannot merge. The deploy workflow runs the same check before publishing, so a +stale analysis is never served. + +## Day to day + +**Adding or changing fork files in a PR.** If the check lists uncovered files, +add each to the section that explains it in `fork.yaml` (or start a new +section with a short description). Files that are not code go under the +top-level `ignore`. Keep globs specific: a `mlx/**` catch-all would swallow +upstream's changes after a bad rebase and defeat the gate. + +**Upstream sync policy (2026-09-25).** This fork does not merge or rebase onto +upstream. `base.hash` stays fixed. An upstream fix the fork needs is ported as +its own commit and described in its own `fork.yaml` section. The check keeps +comparing the fork against the fixed base: + +``` +python3 -m pip install pyyaml +python3 scripts/check_forkdiff.py --upstream-ref refs/remotes/upstream/main +``` + +**Previewing the page locally** (Go 1.21+): + +``` +go run github.com/protolambda/forkdiff@v0.1.1 -repo . -fork fork.yaml -out tmp/index.html +open tmp/index.html +``` + +## Design notes + +- `base.hash` is a full 40-hex commit id, never a branch name: a symbolic base + would move underneath the page and the gate alike. +- Section `ignore` lists count as coverage (forkdiff still lists those files, + grayed out); the top-level `ignore` is for things that aren't code at all. +- The glob semantics are forkdiff's: `*` and `?` stop at `/`, `**` spans + directories (and may match none), `[!x]` negates a class. +- The check is pure git + PyYAML so the same command runs locally and in CI.
diff --git ml-explore/mlx/README.md Layr-Labs/mlx/README.md index fb30a24f2ce87d6f81822fc51e01c9fb6d38de28..b316d5a61bb3f321fc5befe81bc75ae2969875ad 100644 --- ml-explore/mlx/README.md +++ Layr-Labs/mlx/README.md @@ -6,6 +6,8 @@ [**Examples**](#examples)   [![CircleCI](https://circleci.com/gh/ml-explore/mlx.svg?style=svg)](https://circleci.com/gh/ml-explore/mlx)   +> **This is a fork.** `Layr-Labs/mlx` tracks [`ml-explore/mlx`](https://github.com/ml-explore/mlx) and carries the kernel, allocator and Metal-runtime work behind Layr-Labs' Apple-silicon inference stack (`mlx` → `mlx-c` → `mlx-swift` → `mlx-swift-lm`). Everything changed relative to upstream is published as a fork diff at **https://layr-labs.github.io/mlx/**, described in [`fork.yaml`](fork.yaml) and kept honest by CI — see [FORKDIFF.md](FORKDIFF.md). + MLX is an array framework for machine learning on Apple silicon, brought to you by Apple machine learning research.
diff --git ml-explore/mlx/scripts/check_forkdiff.py Layr-Labs/mlx/scripts/check_forkdiff.py new file mode 100644 index 0000000000000000000000000000000000000000..91ae6abe470548758dab3d8db15dd29d76a95bc1 --- /dev/null +++ Layr-Labs/mlx/scripts/check_forkdiff.py @@ -0,0 +1,265 @@ +#!/usr/bin/env python3 +"""Gate: ``fork.yaml`` must describe this fork as it is *now*. + +The fork-diff page (rendered with protolambda/forkdiff and published on GitHub +Pages) is only useful while it is true. Two things make it go stale silently: + +1. **A rebase onto newer upstream.** ``base.hash`` still points at the old + upstream commit, so the page shows upstream's own changes as if the fork + made them. This script requires ``base.hash`` to equal + ``git merge-base HEAD <upstream>``; a rebase moves the merge-base, and the + gate stays red until the hash is bumped. +2. **A new fork change nobody described.** Every path in + ``git diff --name-only base.hash HEAD`` must match a glob in some section + (or a global ``ignore``), and every glob must still match something — a + section describing code the fork no longer carries is as misleading as a + missing one. + +Both are pure git + YAML, so the same check runs locally:: + + python3 scripts/check_forkdiff.py # coverage only + python3 scripts/check_forkdiff.py --upstream-ref upstream/main + +Exit status is non-zero on any violation. ``--review-status-out`` writes the +JSON the CI comment synthesizer consumes (same shape as history-check). +""" + +from __future__ import annotations + +import argparse +import json +import re +import subprocess +import sys +from pathlib import Path +from typing import Iterable + +import yaml + +GLOB_CLASS_RE = re.compile(r"\[([^\]]*)\]") + + +def glob_to_regex(glob: str) -> re.Pattern[str]: + """Translate a forkdiff glob into a regex over the repo-relative path. + + Semantics follow the doublestar rules forkdiff uses: ``*`` and ``?`` never + cross a ``/``; ``**`` matches any number of directories (``a/**/b`` also + matches ``a/b``); ``[...]`` character classes pass through, with a leading + ``!`` meaning negation. + """ + out: list[str] = [] + i = 0 + while i < len(glob): + c = glob[i] + if c == "*": + if glob.startswith("**", i): + if glob.startswith("**/", i): + out.append("(?:.*/)?") + i += 3 + continue + out.append(".*") + i += 2 + continue + out.append("[^/]*") + elif c == "?": + out.append("[^/]") + elif c == "[": + end = glob.find("]", i + 1) + if end == -1: + out.append(re.escape(c)) + else: + cls = glob[i + 1 : end] + if cls.startswith("!"): + cls = "^" + cls[1:] + out.append("[" + cls + "]") + i = end + 1 + continue + else: + out.append(re.escape(c)) + i += 1 + return re.compile("^" + "".join(out) + "$") + + +def collect_globs(node: dict, path: str = "def") -> list[tuple[str, str]]: + """Every glob in the section tree as ``(section path, glob)``. + + A section's ``ignore`` list counts as coverage too: forkdiff still lists + those files under the section (grayed out), so they are described. + """ + found: list[tuple[str, str]] = [] + title = node.get("title") or "(untitled)" + here = f"{path} › {title}" if path != "def" else title + for g in node.get("globs") or []: + found.append((here, str(g))) + for g in node.get("ignore") or []: + found.append((here, str(g))) + for child in node.get("sub") or []: + found.extend(collect_globs(child, here)) + return found + + +def git(*args: str, cwd: Path) -> str: + result = subprocess.run( + ["git", *args], cwd=cwd, capture_output=True, text=True, check=False + ) + if result.returncode != 0: + raise RuntimeError(f"git {' '.join(args)} failed: {result.stderr.strip()}") + return result.stdout.strip() + + +def changed_paths(repo: Path, base: str, head: str) -> list[str]: + out = git("diff", "--name-only", f"{base}..{head}", cwd=repo) + return [line for line in out.splitlines() if line] + + +def check_coverage( + paths: Iterable[str], globs: list[tuple[str, str]] +) -> tuple[list[str], list[tuple[str, str]], dict[str, int]]: + """Return ``(uncovered paths, stale globs, matches per glob)``.""" + compiled = [(section, g, glob_to_regex(g)) for section, g in globs] + hits: dict[str, int] = {g: 0 for _, g in globs} + uncovered: list[str] = [] + for p in paths: + matched = False + for _, g, rx in compiled: + if rx.match(p): + hits[g] += 1 + matched = True + if not matched: + uncovered.append(p) + stale = [(section, g) for section, g in globs if hits[g] == 0] + return uncovered, stale, hits + + +def check_base(repo: Path, base: str, head: str, upstream_ref: str | None) -> list[str]: + problems: list[str] = [] + if not re.fullmatch(r"[0-9a-f]{40}", base): + problems.append( + f"base.hash must be a full 40-hex commit id, got {base!r} — a short or " + "symbolic ref would silently move under the page." + ) + return problems + try: + git("cat-file", "-e", f"{base}^{{commit}}", cwd=repo) + except RuntimeError: + problems.append( + f"base.hash {base[:12]} is not present in this repository. The fork must " + "sit on top of it (fetch upstream if the clone is shallow)." + ) + return problems + if upstream_ref is None: + return problems + try: + git("merge-base", "--is-ancestor", base, upstream_ref, cwd=repo) + except RuntimeError: + problems.append( + f"base.hash {base[:12]} is not an ancestor of {upstream_ref} — it must name " + "a commit on upstream main, not a fork commit." + ) + merge_base = git("merge-base", head, upstream_ref, cwd=repo) + if merge_base != base: + problems.append( + f"base.hash {base[:12]} != merge-base({head}, {upstream_ref}) = " + f"{merge_base[:12]}. The fork was rebased onto newer upstream; update " + "fork.yaml's base.hash and re-describe the sections (see docs/forkdiff.md)." + ) + return problems + + +def review_status(problems: list[str], detail: str) -> list[dict]: + if not problems: + return [] + return [ + { + "source": "fork diff analysis", + "results": [ + { + "kind": "action_required", + "title": "fork.yaml no longer describes the fork", + "summary": ( + problems[0] + if len(problems) == 1 + else f"{len(problems)} issues: {problems[0]}" + ), + "detail": detail, + "how_to_fix": ( + "See docs/forkdiff.md. After a rebase: set base.hash to " + "`git merge-base HEAD upstream/main`. For new files: add them to " + "the section that explains them (or to a global `ignore` if they " + "are not code). Then run `python3 scripts/check_forkdiff.py " + "--upstream-ref upstream/main` locally." + ), + } + ], + } + ] + + +def main(argv: list[str] | None = None) -> int: + ap = argparse.ArgumentParser(description=__doc__.split("\n\n")[0]) + ap.add_argument("--repo", default=".", help="path to the fork checkout") + ap.add_argument("--fork", default="fork.yaml", help="fork page definition") + ap.add_argument("--head", default="HEAD", help="the fork revision to describe") + ap.add_argument( + "--upstream-ref", + default=None, + help="a ref holding upstream main (e.g. refs/remotes/upstream/main); " + "enables the merge-base check", + ) + ap.add_argument( + "--review-status-out", default=None, help="write review-status JSON here" + ) + args = ap.parse_args(argv) + + repo = Path(args.repo).resolve() + fork_path = repo / args.fork + spec = yaml.safe_load(fork_path.read_text(encoding="utf-8")) or {} + base = str((spec.get("base") or {}).get("hash") or "").strip() + problems: list[str] = [] + lines: list[str] = [] + + problems += check_base(repo, base, args.head, args.upstream_ref) + + globs = collect_globs(spec.get("def") or {}) + globs += [("(global ignore)", str(g)) for g in spec.get("ignore") or []] + + paths: list[str] = [] + if re.fullmatch(r"[0-9a-f]{40}", base): + try: + paths = changed_paths(repo, base, args.head) + except RuntimeError as exc: + problems.append(str(exc)) + + uncovered, stale, hits = check_coverage(paths, globs) + if uncovered: + problems.append( + f"{len(uncovered)} changed file(s) are not described by any fork.yaml section" + ) + lines.append("Uncovered files (add each to the section that explains it):") + lines += [f" - {p}" for p in uncovered] + if stale: + problems.append(f"{len(stale)} glob(s) match nothing the fork changes") + lines.append("Stale globs (the fork no longer changes anything they name):") + lines += [f" - {g} [{section}]" for section, g in stale] + + covered = len(paths) - len(uncovered) + print( + f"fork.yaml: base {base[:12]} head {args.head} changed files {len(paths)} " + f"described {covered} sections+ignores {len(globs)}" + ) + if args.upstream_ref is None: + print(" (no --upstream-ref: merge-base check skipped)") + for line in lines: + print(line) + for p in problems: + print(f"::error::{p}") + + if args.review_status_out: + Path(args.review_status_out).write_text( + json.dumps(review_status(problems, "\n".join(lines))), encoding="utf-8" + ) + return 1 if problems else 0 + + +if __name__ == "__main__": + sys.exit(main())
diff --git ml-explore/mlx/.github/workflows/forkdiff-check.yml Layr-Labs/mlx/.github/workflows/forkdiff-check.yml new file mode 100644 index 0000000000000000000000000000000000000000..67f75cc69e06c7b86f10f7cfbffc6b99fdafc56c --- /dev/null +++ Layr-Labs/mlx/.github/workflows/forkdiff-check.yml @@ -0,0 +1,76 @@ +name: Fork Diff Check + +# Fails a PR whose `fork.yaml` no longer describes this fork. +# +# The fork-diff page (https://layr-labs.github.io/mlx/, rendered by +# protolambda/forkdiff from `fork.yaml`, deployed by forkdiff-pages.yml) is +# only useful while it is true, and two things make it go stale silently: +# +# 1. A rebase onto newer upstream. `base.hash` keeps pointing at the old +# upstream commit, so the page shows upstream's own changes as the fork's. +# `scripts/check_forkdiff.py` requires base.hash == merge-base(HEAD, +# upstream/main); a rebase moves the merge-base and the gate stays red +# until the hash — and the sections — are brought up to date. +# 2. A fork change nobody described. Every file in the base..HEAD diff must +# match a section glob (or a global ignore), and every glob must still +# match something. +# +# The page is also rendered here so a fork.yaml that forkdiff itself rejects +# cannot merge. See FORKDIFF.md. + +on: + pull_request: + push: + branches: [main] + workflow_dispatch: + +permissions: + contents: read + +jobs: + analysis-up-to-date: + name: fork.yaml describes the fork + runs-on: ubuntu-latest + timeout-minutes: 20 + steps: + - uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2 + with: + fetch-depth: 0 # full history: merge-base with upstream, and forkdiff diffs against base.hash + + - name: Fetch upstream main + run: | + git fetch --no-tags --quiet https://github.com/ml-explore/mlx.git \ + main:refs/remotes/upstream/main + echo "upstream/main = $(git rev-parse --short refs/remotes/upstream/main)" + echo "merge-base = $(git merge-base HEAD refs/remotes/upstream/main | cut -c1-12)" + + - uses: actions/setup-python@a309ff8b426b58ec0e2a45f0f869d46889d02405 # v6.2.0 + with: + python-version: "3.12" + + - name: Verify fork.yaml against the real diff + run: | + python3 -m pip install --quiet pyyaml + python3 scripts/check_forkdiff.py --upstream-ref refs/remotes/upstream/main + + - uses: actions/setup-go@40f1582b2485089dde7abd97c1529aa768e1baff # v5.6.0 + with: + go-version: "1.24" + cache: false + + - name: Render the fork-diff page + # fork.yaml names refs/heads/main; on a PR the checkout is a detached + # merge commit, so point the local branch at what we are checking. + run: | + git update-ref refs/heads/main HEAD + mkdir -p tmp/pages + go run github.com/protolambda/forkdiff@v0.1.1 \ + -repo . -fork fork.yaml -out tmp/pages/index.html + echo "rendered $(wc -c < tmp/pages/index.html) bytes" + + - name: Upload rendered page + uses: actions/upload-artifact@ea165f8d65b6e75b540449e92b4886f43607fa02 # v4.6.2 + with: + name: forkdiff-page + path: tmp/pages/index.html + retention-days: 7
diff --git ml-explore/mlx/.github/workflows/forkdiff-pages.yml Layr-Labs/mlx/.github/workflows/forkdiff-pages.yml new file mode 100644 index 0000000000000000000000000000000000000000..9a01dacc5cae4a9578c1badef29dc1ba7de9a83e --- /dev/null +++ Layr-Labs/mlx/.github/workflows/forkdiff-pages.yml @@ -0,0 +1,99 @@ +name: Deploy Fork Diff + +# Renders `fork.yaml` with protolambda/forkdiff and publishes the result as the +# repository's GitHub Pages site — the same setup ethereum-optimism/op-geth +# uses for its go-ethereum fork diff. Runs the analysis gate first, so a page +# that lies about the fork is never published. + +on: + push: + branches: [main] + workflow_dispatch: + +permissions: + contents: read + pages: write + id-token: write + +concurrency: + group: "pages" + cancel-in-progress: true + +jobs: + deploy: + name: Render and deploy + environment: + name: github-pages + url: ${{ steps.deployment.outputs.page_url }} + runs-on: ubuntu-latest + timeout-minutes: 20 + steps: + - uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2 + with: + fetch-depth: 0 # forkdiff diffs against base.hash, deep in history + + - name: Fetch upstream main + run: | + git fetch --no-tags --quiet https://github.com/ml-explore/mlx.git \ + main:refs/remotes/upstream/main + + - uses: actions/setup-python@a309ff8b426b58ec0e2a45f0f869d46889d02405 # v6.2.0 + with: + python-version: "3.12" + + - name: Refuse to publish a stale analysis + run: | + python3 -m pip install --quiet pyyaml + python3 scripts/check_forkdiff.py --upstream-ref refs/remotes/upstream/main + + - uses: actions/setup-go@40f1582b2485089dde7abd97c1529aa768e1baff # v5.6.0 + with: + go-version: "1.24" + cache: false + + - name: Build forkdiff + run: | + mkdir -p tmp/pages + go run github.com/protolambda/forkdiff@v0.1.1 \ + -repo . -fork fork.yaml -out tmp/pages/index.html + touch tmp/pages/.nojekyll + + - name: Is GitHub Pages enabled? + # Pages must be switched on once by a repo admin (Settings → Pages → + # Source: "GitHub Actions"). The workflow token cannot do that itself + # (`configure-pages` enablement needs an admin PAT), so until then the + # page is built and kept as an artifact but not deployed — a notice, + # not a red run. + id: pages + env: + GH_TOKEN: ${{ github.token }} + run: | + if gh api "repos/${GITHUB_REPOSITORY}/pages" --silent 2>/dev/null; then + echo "enabled=true" >> "$GITHUB_OUTPUT" + else + echo "enabled=false" >> "$GITHUB_OUTPUT" + echo "::notice::GitHub Pages is not enabled for ${GITHUB_REPOSITORY}; built the page but skipped the deploy. Enable it once under Settings → Pages (Source: GitHub Actions) and re-run this workflow." + fi + + - name: Keep the rendered page as an artifact + if: steps.pages.outputs.enabled != 'true' + uses: actions/upload-artifact@ea165f8d65b6e75b540449e92b4886f43607fa02 # v4.6.2 + with: + name: forkdiff-page + path: tmp/pages/index.html + retention-days: 30 + + - name: Setup Pages + if: steps.pages.outputs.enabled == 'true' + uses: actions/configure-pages@983d7736d9b0ae728b81ab479565c72886d7745b # v5.0.0 + + - name: Upload artifact + if: steps.pages.outputs.enabled == 'true' + uses: actions/upload-pages-artifact@7b1f4a764d45c48632c6b24a0339c27f5614fb0b # v4.0.0 + with: + path: tmp/pages + + - name: Deploy to GitHub Pages + if: steps.pages.outputs.enabled == 'true' + id: deployment + uses: actions/deploy-pages@d6db90164ac5ed86f2b6aed7e0febac5b3c0c03e # v4.0.5
diff --git ml-explore/mlx/fork.yaml Layr-Labs/mlx/fork.yaml new file mode 100644 index 0000000000000000000000000000000000000000..faf3e49a3bdc9a82707b4f18fa7f0c567179eba3 --- /dev/null +++ Layr-Labs/mlx/fork.yaml @@ -0,0 +1,296 @@ +title: "Layr-Labs/mlx - MLX fork diff overview" +footer: | + Fork-diff overview of [`Layr-Labs/mlx`](https://github.com/Layr-Labs/mlx), a fork of + [`ml-explore/mlx`](https://github.com/ml-explore/mlx) &middot; created with + [Forkdiff](https://github.com/protolambda/forkdiff) +base: + name: ml-explore/mlx + url: https://github.com/ml-explore/mlx + # The upstream commit this fork is based on. CI (`scripts/check_forkdiff.py`) + # requires this to equal `git merge-base HEAD upstream/main`, so a rebase onto + # newer upstream fails the gate until the hash — and the sections below — are + # brought up to date. Upstream main as of 2026-08-14. + hash: 77a0c1e8ae16eb781a79e2b7134ba84b2ee7b597 +fork: + name: Layr-Labs/mlx + url: https://github.com/Layr-Labs/mlx + ref: refs/heads/main +def: + title: "Layr-Labs/mlx" + description: | + This is an overview of the changes in [`Layr-Labs/mlx`](https://github.com/Layr-Labs/mlx), + a fork of [`ml-explore/mlx`](https://github.com/ml-explore/mlx). + + The fork is the core of Layr-Labs' Apple-silicon inference stack — `mlx` → `mlx-c` → + `mlx-swift` → `mlx-swift-lm` — and carries the kernel, allocator and Metal-runtime work + that serving production MoE models (Gemma 4, Qwen 3.5/3.6, GPT-OSS) needed before or + beyond what upstream ships: sorted expert-tile quantized matmul routes, a masked-tail + MXFP4 decode kernel, FP32-safe SDPA partials, coherent allocator accounting for admission + control, declared-mutable custom-kernel inputs for paged KV, and a use-after-free fix in + Metal eval. Every change is meant to be additive and rebasable; each is gated by shape, + device or environment checks so the upstream path stays the default elsewhere. + + One section is not fork work at all: the fork carries upstream's own v0.32.2 changes + (applied in PR #10 rather than merged), and those files are listed as such so the page + stays honest about what Layr-Labs actually changed. + sub: + - title: "Metal eval: use-after-free with mid-eval synchronization" + description: | + `gpu::eval` captured the `MTL::CommandBuffer*` before `eval_gpu` ran and attached the + input-buffer liveness handler to it afterwards. Any primitive that calls + `CommandEncoder::synchronize()` mid-eval — the expert-tile route's sortedness-retract + check does — commits and replaces the encoder's buffer, so the handler landed on a + dangling pointer (latent on Gemma E=128, a deterministic SIGSEGV on Qwen E=256). The + buffer is now fetched from the encoder *after* `eval_gpu`, so inputs outlive exactly the + tail of the primitive's work. (PR #5) + globs: + - "mlx/backend/metal/eval.cpp" + - title: "Sorted expert-tile quantized matmul (Gemma 4 E=128, Qwen 3.5/3.6 E=256)" + description: | + An opt-in route (`MLX_GATHER_QMM_EXPERT_SLICES=1|trust`) for gathered, affine-quantized + MoE expert matmuls: a descriptor kernel builds sorted expert tiles on device (one + instantiation per expert count) and a single expert-count-agnostic tile kernel consumes + them. `classify_gemma4_expert_qmm` is a pure route table admitting exactly the shapes + the kernel is built for; availability is all-or-nothing on a source-matched metallib so + an older library fails the route closed; `trust` skips the descriptor-retract readback + when the caller guarantees sorted indices. Diagnostics counters + (`mlx_metal_gemma4_expert_qmm_diagnostics_*`) are exposed through a C ABI so mlx-c / + mlx-swift can arm, snapshot and reset them. Also in these kernels: every affine + quantized-vector bias-sum operand is promoted to the accumulator type before summation + (PR #15). + globs: + - "mlx/backend/common/gemma4_expert_qmm.h" + - "mlx/backend/metal/device.h" + - "mlx/backend/metal/device.cpp" + - "mlx/backend/metal/kernels/quantized.h" + - "mlx/backend/metal/kernels/quantized.metal" + - "tests/gpu_tests.cpp" + - title: "GPT-OSS MXFP4 decode: masked-tail gathered qmv" + description: | + GPT-OSS 20B gathered MXFP4 decode has K=2880, which misses the fast vector kernel's + 512-element alignment and fell to the general path. `fp_gather_qmv_fast_tail` runs five + full blocks plus a 320-value tail read only by the participating lanes, preserving FP32 + accumulation and the gathered indices/strides. It is enabled automatically only on the + physical M4 Max architecture (`applegpu_g16s`) for the exact MXFP4/group-32/4-bit, + E=32, K=2880, N=2880|5760, FP32/BF16 shape; `MLX_GPTOSS_MXFP4_DECODE_FAST_TAIL=0` + restores the original route and `MLX_GPTOSS_MXFP4_PREFILL_TILE=m32n32k32` opts into a + separately measured prefill tile (`gptoss_mxfp4_policy.h`). (PR #14) + `fp_quantized.metal` also builds the two m32n32k32 prefill-tile kernels (float and + bfloat16) ahead of time, so CMake builds with the default `MLX_METAL_JIT=OFF` can load + them. (PR #24) + globs: + - "mlx/backend/metal/gptoss_mxfp4_policy.h" + - "mlx/backend/metal/kernels/fp_quantized.h" + - "mlx/backend/metal/kernels/fp_quantized.metal" + - "mlx/backend/metal/quantized.cpp" + - "tests/gptoss_mxfp4_tests.cpp" + - "tests/CMakeLists.txt" + - title: "Coherent allocator accounting for admission control" + description: | + Admission accounting used to combine independent reads of the allocator counters with + logical buffer sizes, so concurrent allocator activity could produce an incoherent + snapshot and alignment / cache reuse could make the real backing larger than the + request. `get_memory_snapshot()` reads active / cache / peak under one lock; + `get_allocation_size_upper_bound()` and the detached, lock-free + `AllocationFootprintPolicy` bound a buffer including alignment and the (now inclusive, + overflow-checked) cache-reuse limit shared with `BufferCache`. Metal, CPU and CUDA + allocators implement the same interface; the allocator's own behaviour stays the + authority. These describe allocator accounting, not process RSS. (PR #15) + globs: + - "mlx/memory.h" + - "mlx/backend/common/allocation_footprint.h" + - "mlx/backend/common/buffer_cache.h" + - "mlx/backend/metal/allocator.h" + - "mlx/backend/metal/allocator.cpp" + - "mlx/backend/cuda/allocator.h" + - "mlx/backend/cuda/allocator.cpp" + - "mlx/backend/no_gpu/allocator.cpp" + - "tests/allocator_tests.cpp" + - title: "FP32 partial sums in two-pass vector SDPA" + description: | + The two-pass vector SDPA cast unnormalized partial sums to the input dtype before + combining them: BF16 lost cancellation residuals and FP16 could overflow even when the + final output was representable. Partials now stay FP32 through the second pass and only + the final result is cast; kernel entry names change so a stale shader cannot silently + satisfy the new buffer contract. Python regressions cover cancellation and overflow + across D64/D128/D256, 32/128 partitions, masked and unmasked. (PR #16) + globs: + - "mlx/backend/metal/kernels/sdpa_vector.h" + - "mlx/backend/metal/kernels/scaled_dot_product_attention.metal" + - "mlx/backend/metal/scaled_dot_product_attention.cpp" + - "python/tests/test_fast_sdpa.py" + - title: "Declared-mutable inputs for Metal custom kernels" + description: | + `metal_kernel_with_mutable_inputs` lets a caller declare which custom-kernel inputs will + be written, so paged-KV consumers keep writes on the original allocation and Metal sees + the read/write hazard. Mutable names must identify exactly one input; mutable `device` + pointers are generated (including for small/scalar storage) and take part in the + compiled-kernel identity; an implicit contiguous copy of declared-mutable storage is + rejected; allocations are registered for writes as well as reads. A separate + `MutableInputCustomKernel` primitive serializes the declaration while the existing + `CustomKernel` export schema is untouched. (PR #18) + globs: + - "mlx/fast.h" + - "mlx/fast_primitives.h" + - "mlx/backend/common/metal_kernel.cpp" + - "mlx/backend/metal/custom_kernel.cpp" + - "mlx/export.cpp" + - title: "Cross-thread compile-cache cleanup" + description: | + Compile caches became thread-local upstream; bindings built before that (the mlx-c / + mlx-swift consumers) still call the process-wide `compile_erase` / `compile_clear_cache`. + A process-lifetime registry records each thread's cache once at construction so those + entrypoints can erase or clear every live cache, with no per-call hot-path work. (PR #10) + globs: + - "mlx/compile.cpp" + - "mlx/compile_impl.h" + - "tests/compile_tests.cpp" + - title: "Upstream v0.32.2 sync carried in PR #10" + description: | + **Not fork work.** PR #10 moved the fork from an 0.32.0-era upstream main to official + v0.32.2 by applying upstream's changes rather than merging them, so `git` sees them as + fork commits even though every file below is byte-identical to upstream `v0.32.2` + (CUDA backend, CPU SIMD, kernels, Python bindings and tests, repository CI and + templates, `AGENTS.md` / `CLAUDE.md`, docs). Listed here so the page does not credit + Layr-Labs with upstream's changes; the section disappears once the fork rebases onto + an upstream commit at or past v0.32.2. + globs: + - ".pre-commit-config.yaml" + - "AGENTS.md" + - "CLAUDE.md" + - "CMakeLists.txt" + - "CONTRIBUTING.md" + - "setup.py" + - ".github/pull_request_template.md" + - ".github/ISSUE_TEMPLATE/*" + - ".github/actions/build-macos/action.yml" + - ".github/actions/setup/action.yml" + - ".github/actions/build-wheel/action.yml" + - ".github/actions/test-linux/action.yml" + - ".github/actions/test-macos/action.yml" + - "benchmarks/python/sdpa_bench.py" + - "docs/src/install.rst" + - "docs/src/python/fast.rst" + - "examples/extensions/*" + - "mlx/array.h" + - "mlx/einsum.cpp" + - "mlx/error.h" + - "mlx/event.h" + - "mlx/fast.cpp" + - "mlx/fence.h" + - "mlx/fft.cpp" + - "mlx/ops.cpp" + - "mlx/ops.h" + - "mlx/primitives.cpp" + - "mlx/primitives.h" + - "mlx/scheduler.cpp" + - "mlx/scheduler.h" + - "mlx/version.h" + - "mlx/backend/common/load.cpp" + - "mlx/backend/cpu/*" + - "mlx/backend/cpu/simd/base_simd.h" + - "mlx/backend/cuda/CMakeLists.txt" + - "mlx/backend/cuda/cross_entropy.cu" + - "mlx/backend/cuda/cuda_utils.h" + - "mlx/backend/cuda/custom_kernel.cpp" + - "mlx/backend/cuda/delayload.cpp" + - "mlx/backend/cuda/device.cpp" + - "mlx/backend/cuda/device.h" + - "mlx/backend/cuda/dirs.cpp" + - "mlx/backend/cuda/event.cu" + - "mlx/backend/cuda/event.h" + - "mlx/backend/cuda/fence.cpp" + - "mlx/backend/cuda/jit_module.cpp" + - "mlx/backend/cuda/scaled_dot_product_attention.cpp" + - "mlx/backend/cuda/device/binary_ops.cuh" + - "mlx/backend/metal/compiled.cpp" + - "mlx/backend/metal/conv.cpp" + - "mlx/backend/metal/event.cpp" + - "mlx/backend/metal/event.h" + - "mlx/backend/metal/fence.cpp" + - "mlx/backend/metal/jit_kernels.cpp" + - "mlx/backend/metal/kernels.h" + - "mlx/backend/metal/nojit_kernels.cpp" + - "mlx/backend/metal/normalization.cpp" + - "mlx/backend/metal/primitives.cpp" + - "mlx/backend/metal/kernels/binary_ops.h" + - "mlx/backend/metal/kernels/copy.h" + - "mlx/backend/metal/kernels/fp_quantized_nax.h" + - "mlx/backend/metal/kernels/fp_quantized_nax.metal" + - "mlx/backend/metal/kernels/fp8.h" + - "mlx/backend/metal/kernels/quantized_nax.h" + - "mlx/backend/metal/kernels/quantized_nax.metal" + - "mlx/backend/metal/kernels/rms_norm.metal" + - "mlx/backend/metal/kernels/sort.h" + - "mlx/backend/metal/kernels/utils.h" + - "mlx/backend/metal/kernels/reduction/*" + - "mlx/backend/metal/kernels/steel/attn/kernels/*" + - "mlx/backend/no_gpu/event.cpp" + - "mlx/backend/no_gpu/fence.cpp" + - "mlx/backend/no_gpu/primitives.cpp" + - "mlx/distributed/mpi/mpi.cpp" + - "mlx/distributed/ring/ring.cpp" + - "mlx/io/gguf.cpp" + - "python/mlx/*" + - "python/mlx/_distributed_utils/launch.py" + - "python/mlx/nn/losses.py" + - "python/mlx/nn/layers/normalization.py" + - "python/mlx/optimizers/optimizers.py" + - "python/src/*" + - "python/tests/__main__.py" + - "python/tests/mlx_tests.py" + - "python/tests/run.py" + - "python/tests/test_array.py" + - "python/tests/test_compile.py" + - "python/tests/test_conv_transpose.py" + - "python/tests/test_conv.py" + - "python/tests/test_double.py" + - "python/tests/test_einsum.py" + - "python/tests/test_eval.py" + - "python/tests/test_fast.py" + - "python/tests/test_fft.py" + - "python/tests/test_load.py" + - "python/tests/test_nn.py" + - "python/tests/test_ops.py" + - "python/tests/test_quantized.py" + - "python/tests/test_reduce.py" + - "python/tests/test_vmap.py" + - "tests/load_tests.cpp" + - "tests/ops_tests.cpp" + - title: "Fork tooling" + description: | + The fork deletes the upstream workflows that cannot run here: the PyPI release build and + the documentation deploy. `build_and_test.yml` keeps the lint, sanitizer and Fedora jobs. + The fork deletes the jobs that run only in the upstream repository, and the composite + actions that only the deleted workflows and jobs used. + + The gate that keeps this page honest: `check_forkdiff.py` fails CI when `base.hash` is + not the merge-base with upstream, when a file the fork changes is not described by a + section above, or when a section names files the fork no longer changes. + + `fork-tests.yml` builds the fork with Metal and runs its C++ and Python tests on a + GitHub-hosted macOS runner. The `setup` action pins its actions to full commit SHAs, + because the organization blocks actions that are not pinned. The organization also + blocks `hendrikmuhs/ccache-action`, so `setup` installs ccache with Homebrew and keeps + the ccache files with `actions/cache`. + globs: + - ".github/actions/build/action.yml" + - ".github/actions/build-docs/action.yml" + - ".github/actions/build-wheel/action.yml" + - ".github/actions/test-linux/action.yml" + - ".github/actions/test-wheel/action.yml" + - ".github/actions/test-windows/action.yml" + - ".github/workflows/build_and_test.yml" + - ".github/workflows/documentation.yml" + - ".github/workflows/release.yml" + - "scripts/check_forkdiff.py" + - "FORKDIFF.md" + - "README.md" + - ".github/actions/setup/action.yml" + - ".github/workflows/fork-tests.yml" + +# ignored globally, does not count towards line count +ignore: + - "fork.yaml" + - ".github/workflows/forkdiff-check.yml" + - ".github/workflows/forkdiff-pages.yml"