From 85986547c393d4a12f4d6ceee44ef667ba2b8726 Mon Sep 17 00:00:00 2001 From: Kaan Kesgin Date: Wed, 27 May 2026 15:55:09 +0200 Subject: [PATCH 01/12] Rebase convolution suite onto merged FFT support Rebuild the FFT/MPS convolution work on top of the merged AbstractFFTs FFT support (#713), which reorganized the FFT module after this work was first written. - Port convolution.jl, its tests, and perf benchmarks - Add concatTensors and 2D-convolution MPSGraph wrappers to operations.jl - Repoint to the merged MPSGraphFFTDescriptor constructor (replaces the removed create_fft_descriptor helper) - Fix :valid-mode crop offset in the N-D FFT extraction path - Re-export the core conv API from Metal; keep internals in MPSGraphs --- lib/mpsgraphs/MPSGraphs.jl | 1 + lib/mpsgraphs/convolution.jl | 1916 +++++++++++++++++++++++++ lib/mpsgraphs/operations.jl | 92 ++ perf/Project.toml | 3 + perf/benchmark_fused_comprehensive.jl | 324 +++++ perf/benchmark_fused_conv.jl | 96 ++ perf/benchmark_inline_pad.jl | 200 +++ perf/benchmark_throughput.jl | 129 ++ perf/bottleneck_analysis.jl | 495 +++++++ perf/padding_alternatives.jl | 265 ++++ perf/simple_bottleneck.jl | 362 +++++ src/Metal.jl | 4 + test/mpsgraphs/convolution.jl | 401 ++++++ 13 files changed, 4288 insertions(+) create mode 100644 lib/mpsgraphs/convolution.jl create mode 100644 perf/benchmark_fused_comprehensive.jl create mode 100644 perf/benchmark_fused_conv.jl create mode 100644 perf/benchmark_inline_pad.jl create mode 100644 perf/benchmark_throughput.jl create mode 100644 perf/bottleneck_analysis.jl create mode 100644 perf/padding_alternatives.jl create mode 100644 perf/simple_bottleneck.jl create mode 100644 test/mpsgraphs/convolution.jl diff --git a/lib/mpsgraphs/MPSGraphs.jl b/lib/mpsgraphs/MPSGraphs.jl index 19f7cfce9..deb5305b8 100644 --- a/lib/mpsgraphs/MPSGraphs.jl +++ b/lib/mpsgraphs/MPSGraphs.jl @@ -49,5 +49,6 @@ include("random.jl") include("matmul.jl") include("fft.jl") +include("convolution.jl") end diff --git a/lib/mpsgraphs/convolution.jl b/lib/mpsgraphs/convolution.jl new file mode 100644 index 000000000..a406a2a91 --- /dev/null +++ b/lib/mpsgraphs/convolution.jl @@ -0,0 +1,1916 @@ +# FFT-based convolution for MtlArrays +# +# Implements linear convolution using the FFT convolution theorem: +# conv(u, v) = ifft(fft(u) .* fft(v)) +# +# Supported types: +# - Float32, Float16 (real inputs - uses rfft/irfft for efficiency) +# - ComplexF32, ComplexF16 (complex inputs - uses fft/ifft) +# +# For double precision, convert to Float32 first or use CPU DSP.jl. + +using AbstractFFTs + +export conv, conv_fft, conv_fft!, conv_fft_fused, xcorr +export plan_conv_fft, ConvFFTPlan +export get_cached_conv_plan, clear_conv_plan_cache!, clear_fused_conv_cache! + +# ============================================================================ +# Helper Functions +# ============================================================================ + +""" + nextfastfft(n::Integer) + +Return the smallest integer >= n that has only factors of 2, 3, 5, and 7. +These sizes are efficient for FFT computation. +""" +function nextfastfft(n::Integer) + n <= 0 && return 1 + while true + m = n + for p in (2, 3, 5, 7) + while m % p == 0 + m ÷= p + end + end + m == 1 && return n + n += 1 + end +end + +""" + _conv_output_size(signal_size, kernel_size, mode) + +Compute the output size for convolution based on mode. +""" +function _conv_output_size(signal_size::Int, kernel_size::Int, mode::Symbol) + full_size = signal_size + kernel_size - 1 + if mode == :full + return full_size + elseif mode == :same + return signal_size + elseif mode == :valid + return max(signal_size - kernel_size + 1, 0) + else + throw(ArgumentError("Unknown convolution mode: $mode. Use :full, :same, or :valid")) + end +end + +""" + _extract_conv_result(result, output_size, full_size, mode) + +Extract the appropriate portion of the convolution result based on mode. +""" +function _extract_conv_result( + result::MtlArray{T, 1}, output_size::Int, full_size::Int, mode::Symbol + ) where {T} + if mode == :full + return result[1:output_size] + elseif mode == :same + # Center the output around the same size as input + offset = (full_size - output_size) ÷ 2 + return result[(offset + 1):(offset + output_size)] + elseif mode == :valid + # Only fully overlapping region + kernel_size = full_size - output_size + 1 - 1 # Solve for K from N + K - 1 = full, valid = N - K + 1 + offset = full_size - output_size + return result[(offset ÷ 2 + 1):(offset ÷ 2 + output_size)] + end +end + +# Multi-dimensional version +function _extract_conv_result( + result::MtlArray{T, N}, output_sizes::NTuple{N, Int}, full_sizes::NTuple{N, Int}, + mode::Symbol, dims::Union{Int, Tuple} + ) where {T, N} + dims_tuple = dims isa Int ? (dims,) : Tuple(dims) + + # Build index ranges for each dimension + ranges = ntuple(N) do i + if i in dims_tuple + full_size = full_sizes[i] + output_size = output_sizes[i] + if mode == :full + 1:output_size + elseif mode == :same + offset = (full_size - output_size) ÷ 2 + (offset + 1):(offset + output_size) + else # :valid + offset = full_size - output_size + (offset ÷ 2 + 1):(offset ÷ 2 + output_size) + end + else + 1:size(result, i) + end + end + + return result[ranges...] +end + +# ============================================================================ +# Convolution Plan (for repeated convolutions with same sizes) +# ============================================================================ + +""" + ConvFFTPlan{T, N} + +Pre-computed FFT convolution plan for efficient repeated convolutions. + +When you need to convolve many signals with the same kernel, or perform +multiple convolutions with arrays of the same shape, creating a plan +avoids redundant allocations and FFT setup. + +# Fields (internal) +- Pre-allocated padded signal/kernel buffers +- Pre-computed kernel FFT (when kernel is provided at plan time) +- Cached FFT size and output parameters + +# Example +```julia +# Create plan for 1D convolution +signal_size = 1000 +kernel_size = 100 +plan = plan_conv_fft(signal_size, kernel_size, Float32) + +# Use plan for multiple convolutions +for signal in signals + result = plan * signal # Uses pre-allocated buffers +end + +# Or with pre-computed kernel FFT +kernel = MtlVector(randn(Float32, 100)) +plan_with_kernel = plan_conv_fft(1000, kernel) +for signal in signals + result = plan_with_kernel * signal # Even faster - kernel FFT cached +end +``` +""" +struct ConvFFTPlan{T, N, IsReal} + signal_size::NTuple{N, Int} + kernel_size::NTuple{N, Int} + output_size::NTuple{N, Int} + full_size::NTuple{N, Int} + fft_size::NTuple{N, Int} + dims::Tuple{Vararg{Int}} + mode::Symbol + # Pre-allocated buffers + signal_padded::MtlArray{T, N} + kernel_padded::MtlArray{T, N} + # Pre-computed kernel FFT (if kernel was provided) + kernel_fft::Union{Nothing, MtlArray{<:Complex, N}} +end + +# Plan cache for automatic reuse +const _CONV_PLAN_CACHE = Dict{UInt64, ConvFFTPlan}() +const _CONV_PLAN_CACHE_LOCK = ReentrantLock() +const _CONV_PLAN_CACHE_MAX_SIZE = 32 + +""" + _conv_plan_cache_key(signal_size, kernel_size, T, dims, mode) + +Generate a unique key for caching convolution plans. +""" +function _conv_plan_cache_key( + signal_size::NTuple{N, Int}, kernel_size::NTuple{N, Int}, + ::Type{T}, dims::Tuple, mode::Symbol + ) where {N, T} + return hash((signal_size, kernel_size, T, dims, mode)) +end + +""" + plan_conv_fft(signal_size::Int, kernel_size::Int, T::Type; mode=:full) + +Create an FFT convolution plan for 1D arrays of the specified sizes and element type. + +# Arguments +- `signal_size`: Length of signals to convolve +- `kernel_size`: Length of kernels to convolve +- `T`: Element type (Float32, Float16, ComplexF32, or ComplexF16) +- `mode`: Output mode (`:full`, `:same`, or `:valid`) + +# Returns +A `ConvFFTPlan` that can be used with `*` or `mul!` for efficient convolution. + +# Example +```julia +plan = plan_conv_fft(1000, 100, Float32) +signal = MtlVector(randn(Float32, 1000)) +kernel = MtlVector(randn(Float32, 100)) +result = conv_fft(plan, signal, kernel) # Uses pre-allocated buffers +``` +""" +function plan_conv_fft( + signal_size::Int, kernel_size::Int, ::Type{T}; mode::Symbol = :full + ) where {T <: Union{Float32, Float16}} + return _create_conv_plan((signal_size,), (kernel_size,), T, (1,), mode, true) +end + +function plan_conv_fft( + signal_size::Int, kernel_size::Int, ::Type{Complex{T}}; mode::Symbol = :full + ) where {T <: Union{Float32, Float16}} + return _create_conv_plan((signal_size,), (kernel_size,), Complex{T}, (1,), mode, false) +end + +""" + plan_conv_fft(signal_size::NTuple{N,Int}, kernel_size::NTuple{N,Int}, T::Type; dims=1, mode=:full) + +Create an FFT convolution plan for N-dimensional arrays. +""" +function plan_conv_fft( + signal_size::NTuple{N, Int}, kernel_size::NTuple{N, Int}, ::Type{T}; + dims::Union{Int, Tuple{Vararg{Int}}} = 1, mode::Symbol = :full + ) where {N, T <: Union{Float32, Float16}} + dims_tuple = dims isa Int ? (dims,) : Tuple(dims) + return _create_conv_plan(signal_size, kernel_size, T, dims_tuple, mode, true) +end + +function plan_conv_fft( + signal_size::NTuple{N, Int}, kernel_size::NTuple{N, Int}, ::Type{Complex{T}}; + dims::Union{Int, Tuple{Vararg{Int}}} = 1, mode::Symbol = :full + ) where {N, T <: Union{Float32, Float16}} + dims_tuple = dims isa Int ? (dims,) : Tuple(dims) + return _create_conv_plan(signal_size, kernel_size, Complex{T}, dims_tuple, mode, false) +end + +""" + plan_conv_fft(signal_size, kernel::MtlArray; dims=1, mode=:full) + +Create an FFT convolution plan with a pre-computed kernel FFT. + +This is the most efficient option when convolving many signals with the same kernel. +The kernel's FFT is computed once at plan creation time. + +# Example +```julia +kernel = MtlVector(randn(Float32, 100)) +plan = plan_conv_fft(1000, kernel) + +# Each convolution now only requires one FFT (signal) instead of two +for signal in signals + result = conv_fft(plan, signal) +end +``` +""" +function plan_conv_fft( + signal_size::Int, kernel::MtlVector{T}; mode::Symbol = :full + ) where {T <: Union{Float32, Float16}} + plan = _create_conv_plan((signal_size,), (length(kernel),), T, (1,), mode, true) + return _precompute_kernel_fft!(plan, kernel) +end + +function plan_conv_fft( + signal_size::Int, kernel::MtlVector{Complex{T}}; mode::Symbol = :full + ) where {T <: Union{Float32, Float16}} + plan = _create_conv_plan((signal_size,), (length(kernel),), Complex{T}, (1,), mode, false) + return _precompute_kernel_fft!(plan, kernel) +end + +function plan_conv_fft( + signal_size::NTuple{N, Int}, kernel::MtlArray{T, N}; + dims::Union{Int, Tuple{Vararg{Int}}} = 1, mode::Symbol = :full + ) where {N, T <: Union{Float32, Float16}} + dims_tuple = dims isa Int ? (dims,) : Tuple(dims) + plan = _create_conv_plan(signal_size, size(kernel), T, dims_tuple, mode, true) + return _precompute_kernel_fft!(plan, kernel) +end + +function plan_conv_fft( + signal_size::NTuple{N, Int}, kernel::MtlArray{Complex{T}, N}; + dims::Union{Int, Tuple{Vararg{Int}}} = 1, mode::Symbol = :full + ) where {N, T <: Union{Float32, Float16}} + dims_tuple = dims isa Int ? (dims,) : Tuple(dims) + plan = _create_conv_plan(signal_size, size(kernel), Complex{T}, dims_tuple, mode, false) + return _precompute_kernel_fft!(plan, kernel) +end + +""" +Internal function to create a convolution plan. +""" +function _create_conv_plan( + signal_size::NTuple{N, Int}, kernel_size::NTuple{N, Int}, + ::Type{T}, dims::Tuple{Vararg{Int}}, mode::Symbol, is_real::Bool + ) where {N, T} + # Validate dimensions + for d in dims + 1 <= d <= N || + throw(ArgumentError("Invalid dimension $d for array with $N dimensions")) + end + + # Compute sizes + full_size = ntuple(N) do i + if i in dims + signal_size[i] + kernel_size[i] - 1 + else + signal_size[i] + end + end + + output_size = ntuple(N) do i + if i in dims + _conv_output_size(signal_size[i], kernel_size[i], mode) + else + signal_size[i] + end + end + + fft_size = ntuple(N) do i + if i in dims + nextfastfft(full_size[i]) + else + signal_size[i] + end + end + + # Allocate buffers + signal_padded = MtlArray{T}(undef, fft_size) + kernel_padded = MtlArray{T}(undef, fft_size) + + # Zero-fill once (will be overwritten in parts during convolution) + fill!(signal_padded, zero(T)) + fill!(kernel_padded, zero(T)) + + return ConvFFTPlan{T, N, is_real}( + signal_size, kernel_size, output_size, full_size, fft_size, + dims, mode, signal_padded, kernel_padded, nothing + ) +end + +""" +Internal function to pre-compute kernel FFT for a plan. +""" +function _precompute_kernel_fft!(plan::ConvFFTPlan{T, N, IsReal}, kernel::MtlArray{T, N}) where {T, N, IsReal} + # Copy kernel to padded buffer + kernel_ranges = ntuple(i -> 1:plan.kernel_size[i], N) + fill!(plan.kernel_padded, zero(T)) + plan.kernel_padded[kernel_ranges...] = kernel + + # Compute kernel FFT + if IsReal + kernel_fft = rfft(plan.kernel_padded, plan.dims) + else + kernel_fft = fft(plan.kernel_padded, plan.dims) + end + + # Store in a new plan (since structs are immutable, we create a new one) + # Note: This is a bit awkward, but avoids making the struct mutable + return ConvFFTPlan{T, N, IsReal}( + plan.signal_size, plan.kernel_size, plan.output_size, + plan.full_size, plan.fft_size, plan.dims, plan.mode, + plan.signal_padded, plan.kernel_padded, kernel_fft + ) +end + +""" + conv_fft(plan::ConvFFTPlan, signal, kernel) + +Perform convolution using a pre-computed plan. + +Uses pre-allocated buffers from the plan, avoiding allocations. +""" +function conv_fft( + plan::ConvFFTPlan{T, N, true}, signal::MtlArray{T, N}, kernel::MtlArray{T, N} + ) where {T <: Union{Float32, Float16}, N} + @assert size(signal) == plan.signal_size "Signal size $(size(signal)) doesn't match plan $(plan.signal_size)" + @assert size(kernel) == plan.kernel_size "Kernel size $(size(kernel)) doesn't match plan $(plan.kernel_size)" + + # Copy signal to padded buffer + signal_ranges = ntuple(i -> 1:plan.signal_size[i], N) + fill!(plan.signal_padded, zero(T)) + plan.signal_padded[signal_ranges...] = signal + + # FFT signal + S = rfft(plan.signal_padded, plan.dims) + + # Kernel FFT (use cached if available, otherwise compute) + K = if plan.kernel_fft !== nothing + plan.kernel_fft + else + kernel_ranges = ntuple(i -> 1:plan.kernel_size[i], N) + fill!(plan.kernel_padded, zero(T)) + plan.kernel_padded[kernel_ranges...] = kernel + rfft(plan.kernel_padded, plan.dims) + end + + # Multiply and inverse FFT + Y = S .* K + first_dim = minimum(plan.dims) + y = irfft(Y, plan.fft_size[first_dim], plan.dims) + + # Extract result + if N == 1 + return _extract_conv_result(y, plan.output_size[1], plan.full_size[1], plan.mode) + else + return _extract_conv_result(y, plan.output_size, plan.full_size, plan.mode, plan.dims) + end +end + +# Complex version +function conv_fft( + plan::ConvFFTPlan{Complex{T}, N, false}, signal::MtlArray{Complex{T}, N}, + kernel::MtlArray{Complex{T}, N} + ) where {T <: Union{Float32, Float16}, N} + @assert size(signal) == plan.signal_size "Signal size $(size(signal)) doesn't match plan $(plan.signal_size)" + @assert size(kernel) == plan.kernel_size "Kernel size $(size(kernel)) doesn't match plan $(plan.kernel_size)" + + # Copy signal to padded buffer + signal_ranges = ntuple(i -> 1:plan.signal_size[i], N) + fill!(plan.signal_padded, zero(Complex{T})) + plan.signal_padded[signal_ranges...] = signal + + # FFT signal + S = fft(plan.signal_padded, plan.dims) + + # Kernel FFT (use cached if available) + K = if plan.kernel_fft !== nothing + plan.kernel_fft + else + kernel_ranges = ntuple(i -> 1:plan.kernel_size[i], N) + fill!(plan.kernel_padded, zero(Complex{T})) + plan.kernel_padded[kernel_ranges...] = kernel + fft(plan.kernel_padded, plan.dims) + end + + # Multiply and inverse FFT + Y = S .* K + y = ifft(Y, plan.dims) + + # Extract result + if N == 1 + return _extract_conv_result(y, plan.output_size[1], plan.full_size[1], plan.mode) + else + return _extract_conv_result(y, plan.output_size, plan.full_size, plan.mode, plan.dims) + end +end + +""" + conv_fft(plan::ConvFFTPlan, signal) + +Perform convolution using a plan with pre-computed kernel FFT. + +This is the fastest option - only one FFT (for the signal) is needed. +""" +function conv_fft( + plan::ConvFFTPlan{T, N, true}, signal::MtlArray{T, N} + ) where {T <: Union{Float32, Float16}, N} + plan.kernel_fft === nothing && + throw(ArgumentError("Plan has no pre-computed kernel FFT. Use conv_fft(plan, signal, kernel) or create plan with kernel.")) + + @assert size(signal) == plan.signal_size "Signal size $(size(signal)) doesn't match plan $(plan.signal_size)" + + # Copy signal to padded buffer + signal_ranges = ntuple(i -> 1:plan.signal_size[i], N) + fill!(plan.signal_padded, zero(T)) + plan.signal_padded[signal_ranges...] = signal + + # FFT signal and multiply with cached kernel FFT + S = rfft(plan.signal_padded, plan.dims) + Y = S .* plan.kernel_fft + + # Inverse FFT + first_dim = minimum(plan.dims) + y = irfft(Y, plan.fft_size[first_dim], plan.dims) + + # Extract result + if N == 1 + return _extract_conv_result(y, plan.output_size[1], plan.full_size[1], plan.mode) + else + return _extract_conv_result(y, plan.output_size, plan.full_size, plan.mode, plan.dims) + end +end + +# Complex version with pre-computed kernel +function conv_fft( + plan::ConvFFTPlan{Complex{T}, N, false}, signal::MtlArray{Complex{T}, N} + ) where {T <: Union{Float32, Float16}, N} + plan.kernel_fft === nothing && + throw(ArgumentError("Plan has no pre-computed kernel FFT. Use conv_fft(plan, signal, kernel) or create plan with kernel.")) + + @assert size(signal) == plan.signal_size "Signal size $(size(signal)) doesn't match plan $(plan.signal_size)" + + # Copy signal to padded buffer + signal_ranges = ntuple(i -> 1:plan.signal_size[i], N) + fill!(plan.signal_padded, zero(Complex{T})) + plan.signal_padded[signal_ranges...] = signal + + # FFT signal and multiply with cached kernel FFT + S = fft(plan.signal_padded, plan.dims) + Y = S .* plan.kernel_fft + + # Inverse FFT + y = ifft(Y, plan.dims) + + # Extract result + if N == 1 + return _extract_conv_result(y, plan.output_size[1], plan.full_size[1], plan.mode) + else + return _extract_conv_result(y, plan.output_size, plan.full_size, plan.mode, plan.dims) + end +end + +""" + get_cached_conv_plan(signal_size, kernel_size, T; dims=1, mode=:full) + +Get or create a cached convolution plan for the given parameters. + +Plans are cached globally and reused for repeated convolutions with the same sizes. +This is useful when array sizes are known in advance and convolutions are repeated. + +# Thread Safety +Plan cache access is thread-safe using a lock. + +# Cache Size +The cache holds up to $_CONV_PLAN_CACHE_MAX_SIZE plans. When full, the oldest plan +is evicted (FIFO). +""" +function get_cached_conv_plan( + signal_size::NTuple{N, Int}, kernel_size::NTuple{N, Int}, ::Type{T}; + dims::Union{Int, Tuple{Vararg{Int}}} = 1, mode::Symbol = :full + ) where {N, T} + dims_tuple = dims isa Int ? (dims,) : Tuple(dims) + key = _conv_plan_cache_key(signal_size, kernel_size, T, dims_tuple, mode) + + lock(_CONV_PLAN_CACHE_LOCK) do + if haskey(_CONV_PLAN_CACHE, key) + return _CONV_PLAN_CACHE[key] + else + # Create new plan + is_real = T <: Real + plan = _create_conv_plan(signal_size, kernel_size, T, dims_tuple, mode, is_real) + + # Evict oldest if cache is full + if length(_CONV_PLAN_CACHE) >= _CONV_PLAN_CACHE_MAX_SIZE + # Simple FIFO eviction - delete first key + first_key = first(keys(_CONV_PLAN_CACHE)) + delete!(_CONV_PLAN_CACHE, first_key) + end + + _CONV_PLAN_CACHE[key] = plan + return plan + end + end +end + +# 1D convenience +function get_cached_conv_plan( + signal_size::Int, kernel_size::Int, ::Type{T}; mode::Symbol = :full + ) where {T} + return get_cached_conv_plan((signal_size,), (kernel_size,), T; dims = 1, mode = mode) +end + +""" + clear_conv_plan_cache!() + +Clear the global convolution plan cache, freeing GPU memory. +""" +function clear_conv_plan_cache!() + lock(_CONV_PLAN_CACHE_LOCK) do + empty!(_CONV_PLAN_CACHE) + end + return nothing +end + +# ============================================================================ +# Fused MPSGraph Convolution (Single Graph Execution) +# ============================================================================ + +# Cache key for fused convolution graphs +struct FusedConvGraphKey + signal_fft_size::Tuple{Vararg{Int}} # Padded size for FFT + kernel_fft_size::Tuple{Vararg{Int}} # Should match signal_fft_size + output_size::Tuple{Vararg{Int}} # Output shape after extraction + eltype::DataType # Float32 or Float16 +end + +# Cached fused convolution graph +struct CachedFusedConvGraph + graph::MPSGraph + signal_placeholder::MPSGraphTensor + kernel_placeholder::MPSGraphTensor + result::MPSGraphTensor +end + +# Thread-safe cache for fused convolution graphs +const _fused_conv_graph_cache = Dict{FusedConvGraphKey, CachedFusedConvGraph}() +const _fused_conv_graph_cache_lock = ReentrantLock() + +# ============================================================================ +# Buffer Pooling for Fused Convolution +# ============================================================================ + +# Key for buffer pool: (fft_sizes, eltype) +struct BufferPoolKey + fft_sizes::Tuple{Vararg{Int}} + eltype::DataType +end + +# Cached buffers for fused convolution +mutable struct CachedFusedConvBuffers{T, N} + signal_padded::MtlArray{T, N} + kernel_padded::MtlArray{T, N} + output::MtlArray{T, N} +end + +# Thread-safe buffer pool +const _fused_conv_buffer_pool = Dict{BufferPoolKey, CachedFusedConvBuffers}() +const _fused_conv_buffer_pool_lock = ReentrantLock() + +""" +Get or create cached buffers for fused convolution. +Returns pre-allocated padded signal, kernel, and output buffers. +""" +function _get_cached_buffers(fft_sizes::NTuple{N, Int}, ::Type{T}) where {N, T} + key = BufferPoolKey(fft_sizes, T) + cached = get(_fused_conv_buffer_pool, key, nothing) + if cached !== nothing + return cached + end + lock(_fused_conv_buffer_pool_lock) do + cached = get(_fused_conv_buffer_pool, key, nothing) + if cached !== nothing + return cached + end + # Allocate new buffers + signal_padded = MtlArray{T, N}(undef, fft_sizes) + kernel_padded = MtlArray{T, N}(undef, fft_sizes) + output = MtlArray{T, N}(undef, fft_sizes) + cached = CachedFusedConvBuffers{T, N}(signal_padded, kernel_padded, output) + _fused_conv_buffer_pool[key] = cached + return cached + end +end + +""" + clear_fused_conv_buffer_pool!() + +Clear the fused convolution buffer pool, freeing GPU memory. +""" +function clear_fused_conv_buffer_pool!() + lock(_fused_conv_buffer_pool_lock) do + empty!(_fused_conv_buffer_pool) + end + return nothing +end + +# ============================================================================ +# Fast Padding Kernel (Single kernel for copy + zero-pad) +# ============================================================================ + +# Custom Metal kernel that copies source data and zero-pads in one operation +# This is ~4.5x faster than separate copyto! + broadcast zero operations +function _pad_copy_kernel_1d!(dest, src, src_len) + i = thread_position_in_grid_1d() + if i <= src_len + @inbounds dest[i] = src[i] + elseif i <= length(dest) + @inbounds dest[i] = zero(eltype(dest)) + end + return +end + +""" +Copy source array to destination with zero-padding using a single GPU kernel. +Much faster than separate copyto! + broadcast operations (~4.5x speedup). +""" +function _fast_pad_copy!(dest::MtlVector{T}, src::MtlVector{T}) where T + src_len = length(src) + dest_len = length(dest) + threads = min(256, dest_len) + groups = cld(dest_len, threads) + @metal threads=threads groups=groups _pad_copy_kernel_1d!(dest, src, src_len) + return dest +end + +# N-D version: pad along all dimensions (linearized) +function _pad_copy_kernel_nd!(dest, src, src_linear_len) + i = thread_position_in_grid_1d() + if i <= src_linear_len + @inbounds dest[i] = src[i] + elseif i <= length(dest) + @inbounds dest[i] = zero(eltype(dest)) + end + return +end + +""" +Fast N-D padding: copies source to destination buffer with zero-padding. +For N-D arrays, this only works correctly when source fits contiguously at the start. +For general N-D padding with different sizes per dimension, use _fast_pad_copy_nd!. +""" +function _fast_pad_copy_contiguous!(dest::MtlArray{T, N}, src::MtlArray{T, N}) where {T, N} + src_len = length(src) + dest_len = length(dest) + threads = min(256, dest_len) + groups = cld(dest_len, threads) + @metal threads=threads groups=groups _pad_copy_kernel_nd!(dest, src, src_len) + return dest +end + +# ============================================================================ +# In-Place Padding Graphs (Zero-copy padding inside MPSGraph) +# ============================================================================ + +# Cache key for graphs with inline padding +struct InlinePadConvGraphKey + signal_sizes::Tuple{Vararg{Int}} # Original signal shape + kernel_sizes::Tuple{Vararg{Int}} # Original kernel shape + fft_sizes::Tuple{Vararg{Int}} # Padded FFT size + output_sizes::Tuple{Vararg{Int}} # Output shape + eltype::DataType +end + +# Cached graph with inline padding +struct CachedInlinePadConvGraph + graph::MPSGraph + signal_placeholder::MPSGraphTensor + kernel_placeholder::MPSGraphTensor + result::MPSGraphTensor +end + +const _inline_pad_conv_cache = Dict{InlinePadConvGraphKey, CachedInlinePadConvGraph}() +const _inline_pad_conv_cache_lock = ReentrantLock() + +""" +Helper to create a zeros tensor of a given shape inside MPSGraph. +Uses constantWithScalar + broadcastTensor. +""" +function _create_zeros_tensor(graph::MPSGraph, shape::NTuple{N, Int}, ::Type{T}) where {N, T} + zero_scalar = constantWithScalar(graph, T(0), T) + mps_shape = MPSShape([NSNumber(Int32(s)) for s in reverse(shape)]) + return broadcastTensor(graph, zero_scalar, mps_shape, "zeros_$(join(shape, 'x'))") +end + +""" +Build a fused N-D convolution graph with inline padding. +Accepts unpadded signal and kernel, pads inside the graph using concat. +""" +function _build_inline_pad_conv_graph_nd( + signal_sizes::NTuple{N, Int}, kernel_sizes::NTuple{N, Int}, + fft_sizes::NTuple{N, Int}, ::Type{T} + ) where {N, T <: Union{Float32, Float16}} + graph = MPSGraph() + + # Placeholders for UNPADDED inputs + signal_ph = placeholderTensor(graph, signal_sizes, T) + kernel_ph = placeholderTensor(graph, kernel_sizes, T) + + # Create zeros tensors for padding (for each dimension) + # Pad signal: concat(signal, zeros) along each dimension + signal_padded = signal_ph + for dim in 1:N + pad_size = fft_sizes[dim] - signal_sizes[dim] + if pad_size > 0 + # Create zeros for this dimension + # Shape: same as current signal_padded but with pad_size in this dimension + current_shape = ntuple(N) do i + if i == dim + pad_size + elseif i < dim + fft_sizes[i] # Already padded dimensions + else + signal_sizes[i] # Not yet padded dimensions + end + end + zeros_tensor = _create_zeros_tensor(graph, current_shape, T) + # Concat along this dimension (Metal uses reversed axis order) + metal_dim = N - dim # Convert Julia dim to Metal axis + tensors = NSArray([signal_padded, zeros_tensor]) + signal_padded = concatTensors(graph, tensors, metal_dim, "signal_pad_dim$(dim)") + end + end + + # Pad kernel similarly + kernel_padded = kernel_ph + for dim in 1:N + pad_size = fft_sizes[dim] - kernel_sizes[dim] + if pad_size > 0 + current_shape = ntuple(N) do i + if i == dim + pad_size + elseif i < dim + fft_sizes[i] + else + kernel_sizes[i] + end + end + zeros_tensor = _create_zeros_tensor(graph, current_shape, T) + metal_dim = N - dim + tensors = NSArray([kernel_padded, zeros_tensor]) + kernel_padded = concatTensors(graph, tensors, metal_dim, "kernel_pad_dim$(dim)") + end + end + + # Now proceed with FFT convolution on padded tensors + fft_desc_fwd = MPSGraphFFTDescriptor(inverse = false) + axes = NSArray([NSNumber(Int32(i)) for i in (N-1):-1:0]) + + signal_fft = realToHermiteanFFTWithTensor(graph, signal_padded, axes, fft_desc_fwd, "signal_rfft") + kernel_fft = realToHermiteanFFTWithTensor(graph, kernel_padded, axes, fft_desc_fwd, "kernel_rfft") + + product = multiplicationWithPrimaryTensor(graph, signal_fft, kernel_fft, "freq_multiply") + + total_size = prod(fft_sizes) + fft_desc_inv = MPSGraphFFTDescriptor(inverse = true) + fft_desc_inv.roundToOddHermitean = isodd(fft_sizes[1]) + result_unscaled = HermiteanToRealFFTWithTensor(graph, product, axes, fft_desc_inv, "irfft") + + scale_factor = constantWithScalar(graph, T(1) / T(total_size), T) + result_scaled = multiplicationWithPrimaryTensor(graph, result_unscaled, scale_factor, "scale") + + return CachedInlinePadConvGraph(graph, signal_ph, kernel_ph, result_scaled) +end + +""" +Get or create a cached inline-padding convolution graph. +""" +function _get_cached_inline_pad_conv_graph(key::InlinePadConvGraphKey) + cached = get(_inline_pad_conv_cache, key, nothing) + if cached !== nothing + return cached + end + lock(_inline_pad_conv_cache_lock) do + cached = get(_inline_pad_conv_cache, key, nothing) + if cached !== nothing + return cached + end + cached = _build_inline_pad_conv_graph_nd( + key.signal_sizes, key.kernel_sizes, key.fft_sizes, key.eltype) + _inline_pad_conv_cache[key] = cached + return cached + end +end + +""" + conv_fft_inline_pad(signal::MtlArray{T,N}, kernel::MtlArray{T,N}; mode=:full) + +Compute N-D convolution using a fused MPSGraph with inline padding. +Eliminates Julia-side copy/zero kernel launches by moving padding into the graph. +""" +function conv_fft_inline_pad( + signal::MtlArray{T, N}, kernel::MtlArray{T, N}; mode::Symbol = :full + ) where {T <: Union{Float32, Float16}, N} + signal_sizes = size(signal) + kernel_sizes = size(kernel) + + # Compute sizes + full_sizes = ntuple(i -> signal_sizes[i] + kernel_sizes[i] - 1, N) + output_sizes = ntuple(i -> _conv_output_size(signal_sizes[i], kernel_sizes[i], mode), N) + fft_sizes = ntuple(i -> nextfastfft(full_sizes[i]), N) + + # Get cached graph with inline padding + key = InlinePadConvGraphKey(signal_sizes, kernel_sizes, fft_sizes, output_sizes, T) + cached = _get_cached_inline_pad_conv_graph(key) + + # Get output buffer from pool + buffers = _get_cached_buffers(fft_sizes, T) + output = buffers.output + + # Execute graph - no Julia-side padding needed! + @autoreleasepool begin + feeds = Dict{MPSGraphTensor, MPSGraphTensorData}( + cached.signal_placeholder => MPSGraphTensorData(signal), + cached.kernel_placeholder => MPSGraphTensorData(kernel) + ) + + resultdict = Dict{MPSGraphTensor, MPSGraphTensorData}( + cached.result => MPSGraphTensorData(output) + ) + + cmdbuf = MPSCommandBuffer(Metal.global_queue(current_device())) + encode!(cmdbuf, cached.graph, NSDictionary(feeds), NSDictionary(resultdict), nil, default_exec_desc()) + commit!(cmdbuf) + wait_completed(cmdbuf) + + return _extract_conv_result_nd(output, output_sizes, full_sizes, mode) + end +end + +""" + clear_inline_pad_conv_cache!() + +Clear the inline padding convolution graph cache. +""" +function clear_inline_pad_conv_cache!() + lock(_inline_pad_conv_cache_lock) do + empty!(_inline_pad_conv_cache) + end + return nothing +end + +""" +Build a fused N-D convolution graph: rfft(signal) * rfft(kernel) → irfft → scale +All operations in a single MPSGraph for minimal command submission overhead. +Works for any dimensionality (1D, 2D, 3D, etc). +""" +function _build_fused_conv_graph_nd(fft_sizes::NTuple{N, Int}, ::Type{T}) where {N, T <: Union{Float32, Float16}} + graph = MPSGraph() + + # Placeholders for padded signal and kernel (same shape) + signal_ph = placeholderTensor(graph, fft_sizes, T) + kernel_ph = placeholderTensor(graph, fft_sizes, T) + + # FFT descriptor for forward transform + fft_desc_fwd = MPSGraphFFTDescriptor(inverse = false) + + # Metal uses reversed axis ordering: axis 0 in Metal = last axis in Julia + # For N-D, we transform all dimensions + axes = NSArray([NSNumber(Int32(i)) for i in (N-1):-1:0]) + + # Forward rfft on both inputs + signal_fft = realToHermiteanFFTWithTensor(graph, signal_ph, axes, fft_desc_fwd, "signal_rfft") + kernel_fft = realToHermiteanFFTWithTensor(graph, kernel_ph, axes, fft_desc_fwd, "kernel_rfft") + + # Element-wise multiplication in frequency domain + product = multiplicationWithPrimaryTensor(graph, signal_fft, kernel_fft, "freq_multiply") + + # Inverse rfft - scale factor is product of all FFT dimensions + total_size = prod(fft_sizes) + fft_desc_inv = MPSGraphFFTDescriptor(inverse = true) + # For irfft, roundToOddHermitean depends on the first transformed dimension (last in Julia order) + fft_desc_inv.roundToOddHermitean = isodd(fft_sizes[1]) + result_unscaled = HermiteanToRealFFTWithTensor(graph, product, axes, fft_desc_inv, "irfft") + + # Apply scaling: divide by total FFT size + scale_factor = constantWithScalar(graph, T(1) / T(total_size), T) + result_scaled = multiplicationWithPrimaryTensor(graph, result_unscaled, scale_factor, "scale") + + # Return full result - slicing done in Julia for flexibility + return CachedFusedConvGraph(graph, signal_ph, kernel_ph, result_scaled) +end + +""" +Get or create a cached fused convolution graph. +""" +function _get_cached_fused_conv_graph(key::FusedConvGraphKey) + cached = get(_fused_conv_graph_cache, key, nothing) + if cached !== nothing + return cached + end + lock(_fused_conv_graph_cache_lock) do + cached = get(_fused_conv_graph_cache, key, nothing) + if cached !== nothing + return cached + end + # Build N-D fused graph + cached = _build_fused_conv_graph_nd(key.signal_fft_size, key.eltype) + _fused_conv_graph_cache[key] = cached + return cached + end +end + +""" + conv_fft_fused(signal::MtlArray{T,N}, kernel::MtlArray{T,N}; mode=:full) + +Compute N-D convolution using a fused MPSGraph that executes rfft → multiply → irfft → scale +in a single graph execution, minimizing command submission overhead. + +Convolution is performed along all dimensions. For 1D, 2D, 3D, etc. +This is optimized for throughput when processing many convolutions. +""" +function conv_fft_fused( + signal::MtlArray{T, N}, kernel::MtlArray{T, N}; mode::Symbol = :full + ) where {T <: Union{Float32, Float16}, N} + signal_sizes = size(signal) + kernel_sizes = size(kernel) + + # Compute full and output sizes for each dimension + full_sizes = ntuple(i -> signal_sizes[i] + kernel_sizes[i] - 1, N) + output_sizes = ntuple(i -> _conv_output_size(signal_sizes[i], kernel_sizes[i], mode), N) + fft_sizes = ntuple(i -> nextfastfft(full_sizes[i]), N) + + # Get cached fused graph + key = FusedConvGraphKey(fft_sizes, fft_sizes, output_sizes, T) + cached = _get_cached_fused_conv_graph(key) + + # Get cached buffers (avoids GPU memory allocation per call) + buffers = _get_cached_buffers(fft_sizes, T) + signal_padded = buffers.signal_padded + kernel_padded = buffers.kernel_padded + output = buffers.output + + # Copy data to padded arrays with zero-padding + # Use fast single-kernel approach for 1D (4.5x faster than separate copy + zero-fill) + if N == 1 + _fast_pad_copy!(signal_padded, signal) + _fast_pad_copy!(kernel_padded, kernel) + else + # For N-D, use the standard approach (could be optimized further) + signal_ranges = ntuple(i -> 1:signal_sizes[i], N) + kernel_ranges = ntuple(i -> 1:kernel_sizes[i], N) + signal_padded[signal_ranges...] = signal + kernel_padded[kernel_ranges...] = kernel + _zero_padding_regions_fused!(signal_padded, signal_sizes, fft_sizes) + _zero_padding_regions_fused!(kernel_padded, kernel_sizes, fft_sizes) + end + + # Execute fused graph + @autoreleasepool begin + feeds = Dict{MPSGraphTensor, MPSGraphTensorData}( + cached.signal_placeholder => MPSGraphTensorData(signal_padded), + cached.kernel_placeholder => MPSGraphTensorData(kernel_padded) + ) + + resultdict = Dict{MPSGraphTensor, MPSGraphTensorData}( + cached.result => MPSGraphTensorData(output) + ) + + cmdbuf = MPSCommandBuffer(Metal.global_queue(current_device())) + encode!(cmdbuf, cached.graph, NSDictionary(feeds), NSDictionary(resultdict), nil, default_exec_desc()) + commit!(cmdbuf) + wait_completed(cmdbuf) + + # Extract appropriate region based on mode (must copy since output buffer is reused) + return _extract_conv_result_nd(output, output_sizes, full_sizes, mode) + end +end + +# Helper to zero padding regions for N-D arrays (all dimensions) +function _zero_padding_regions_fused!(arr::MtlArray{T, N}, data_sizes::NTuple{N, Int}, fft_sizes::NTuple{N, Int}) where {T, N} + for dim in 1:N + if data_sizes[dim] < fft_sizes[dim] + # Build ranges: full range for other dims, padding range for this dim + ranges = ntuple(N) do i + if i == dim + (data_sizes[i] + 1):fft_sizes[i] + else + 1:fft_sizes[i] + end + end + @view(arr[ranges...]) .= zero(T) + end + end +end + +# Helper for extracting result based on mode +function _extract_conv_result_nd(y::MtlArray{T, N}, output_sizes::NTuple{N, Int}, full_sizes::NTuple{N, Int}, mode::Symbol) where {T, N} + if mode == :full + # Just take the first output_size elements in each dimension + ranges = ntuple(i -> 1:output_sizes[i], N) + return y[ranges...] + elseif mode == :same + # Center the output + ranges = ntuple(N) do i + offset = (full_sizes[i] - output_sizes[i]) ÷ 2 + (offset + 1):(offset + output_sizes[i]) + end + return y[ranges...] + else # :valid + # Only the fully-overlapping region, centered within the full convolution + # (matches _extract_conv_result and the direct path) + ranges = ntuple(N) do i + offset = full_sizes[i] - output_sizes[i] + (offset ÷ 2 + 1):(offset ÷ 2 + output_sizes[i]) + end + return y[ranges...] + end +end + +""" + clear_fused_conv_cache!() + +Clear the fused convolution graph cache and buffer pool, freeing GPU memory. +""" +function clear_fused_conv_cache!() + lock(_fused_conv_graph_cache_lock) do + empty!(_fused_conv_graph_cache) + end + clear_fused_conv_buffer_pool!() + return nothing +end + +# Export the new function +# (added to exports at top of file) + +# ============================================================================ +# Efficient Padding Helpers +# ============================================================================ + +""" + _zero_padding_regions!(arr, data_sizes, padded_sizes, dims) + +Zero only the padding regions of an N-D array, avoiding unnecessary writes. +For each dimension in `dims`, zeros elements from (data_size+1) to padded_size. +""" +function _zero_padding_regions!( + arr::MtlArray{T, N}, data_sizes::NTuple{N, Int}, + padded_sizes::NTuple{N, Int}, dims::Tuple + ) where {T, N} + # For each dimension that needs padding + for d in dims + if data_sizes[d] < padded_sizes[d] + # Build ranges for the padding strip in dimension d + ranges = ntuple(N) do i + if i == d + # Padding region in this dimension + (data_sizes[i]+1):padded_sizes[i] + elseif i < d + # For earlier dimensions, include entire padded size + # (to cover corner regions that previous strips may have missed) + 1:padded_sizes[i] + else + # For later dimensions, include only data region + # (corners will be covered by later strips) + 1:data_sizes[i] + end + end + @view(arr[ranges...]) .= zero(T) + end + end + return nothing +end + +# ============================================================================ +# 1D FFT Convolution (Real Inputs - Optimized) +# ============================================================================ + +""" + conv_fft(signal::MtlVector, kernel::MtlVector; mode=:full) + +Compute the 1D linear convolution of `signal` and `kernel` using FFT. + +# Arguments +- `signal`: Input signal (1D MtlArray) +- `kernel`: Convolution kernel (1D MtlArray) +- `mode`: Output mode + - `:full` (default): Full convolution, output length = length(signal) + length(kernel) - 1 + - `:same`: Output has same length as signal (centered) + - `:valid`: Only fully overlapping region, output length = max(length(signal) - length(kernel) + 1, 0) + +# Returns +MtlArray with the convolution result. + +# Example +```julia +using Metal + +signal = MtlVector(randn(Float32, 1000)) +kernel = MtlVector(randn(Float32, 100)) +result = conv_fft(signal, kernel) # length = 1099 +result_same = conv_fft(signal, kernel; mode=:same) # length = 1000 +``` + +# Notes +- Uses `rfft`/`irfft` for real inputs (2x memory savings vs full FFT) +- Pads to next fast FFT size for optimal performance +- For complex inputs, use the complex-valued method +""" +function conv_fft( + signal::MtlVector{T}, kernel::MtlVector{T}; mode::Symbol = :full + ) where {T <: Union{Float32, Float16}} + # Delegate to fused implementation for better performance + # (single MPSGraph execution instead of 4+ separate operations) + return conv_fft_fused(signal, kernel; mode=mode) +end + +# ============================================================================ +# 1D FFT Convolution (Complex Inputs) +# ============================================================================ + +""" + conv_fft(signal::MtlVector{Complex{T}}, kernel::MtlVector{Complex{T}}; mode=:full) + +Compute the 1D linear convolution of complex `signal` and `kernel` using FFT. +""" +function conv_fft( + signal::MtlVector{Complex{T}}, kernel::MtlVector{Complex{T}}; mode::Symbol = :full + ) where {T <: Union{Float32, Float16}} + ns = length(signal) + nk = length(kernel) + + # Compute sizes + full_size = ns + nk - 1 + output_size = _conv_output_size(ns, nk, mode) + + # Find optimal FFT size + nfft = nextfastfft(full_size) + + # Allocate padded arrays + signal_padded = MtlArray{Complex{T}}(undef, nfft) + kernel_padded = MtlArray{Complex{T}}(undef, nfft) + + # Copy data first, then zero only the padding region (not the entire buffer) + copyto!(signal_padded, 1, signal, 1, ns) + copyto!(kernel_padded, 1, kernel, 1, nk) + + # Zero only the padding regions + if ns < nfft + @view(signal_padded[(ns+1):nfft]) .= zero(Complex{T}) + end + if nk < nfft + @view(kernel_padded[(nk+1):nfft]) .= zero(Complex{T}) + end + + # FFT + S = fft(signal_padded) + K = fft(kernel_padded) + + # Multiply in frequency domain (in-place to avoid allocation) + S .*= K + + # Inverse FFT + y = ifft(S) + + # Extract appropriate region + return _extract_conv_result(y, output_size, full_size, mode) +end + +# ============================================================================ +# N-D FFT Convolution (along specified dimensions) +# ============================================================================ + +""" + conv_fft(signal::MtlArray, kernel::MtlArray; dims=1, mode=:full) + +Compute N-dimensional linear convolution along specified dimensions using FFT. + +# Arguments +- `signal`: Input signal (N-dimensional MtlArray) +- `kernel`: Convolution kernel (same number of dimensions as signal) +- `dims`: Dimension(s) along which to convolve (default: 1). Can be an integer or tuple. +- `mode`: Output mode (`:full`, `:same`, or `:valid`) + +# Returns +MtlArray with the convolution result. + +# Example +```julia +# 2D convolution along both dimensions +signal = MtlArray(randn(Float32, 100, 100)) +kernel = MtlArray(randn(Float32, 5, 5)) +result = conv_fft(signal, kernel; dims=(1,2)) + +# 1D convolution along rows only +result_rows = conv_fft(signal, kernel; dims=1) +``` +""" +function conv_fft( + signal::MtlArray{T, N}, kernel::MtlArray{T, N}; + dims::Union{Int, Tuple{Vararg{Int}}} = 1, mode::Symbol = :full + ) where {T <: Union{Float32, Float16}, N} + dims_tuple = dims isa Int ? (dims,) : Tuple(dims) + + # Validate dimensions + for d in dims_tuple + 1 <= d <= N || + throw(ArgumentError("Invalid dimension $d for array with $N dimensions")) + end + + # Use fused implementation when convolving along ALL dimensions (faster single-graph execution) + if length(dims_tuple) == N && Set(dims_tuple) == Set(1:N) + return conv_fft_fused(signal, kernel; mode=mode) + end + + # Compute output sizes for each convolved dimension + signal_sizes = size(signal) + kernel_sizes = size(kernel) + + full_sizes = ntuple(N) do i + if i in dims_tuple + signal_sizes[i] + kernel_sizes[i] - 1 + else + signal_sizes[i] + end + end + + output_sizes = ntuple(N) do i + if i in dims_tuple + _conv_output_size(signal_sizes[i], kernel_sizes[i], mode) + else + signal_sizes[i] + end + end + + # Compute FFT sizes + fft_sizes = ntuple(N) do i + if i in dims_tuple + nextfastfft(full_sizes[i]) + else + signal_sizes[i] + end + end + + # Pad signal and kernel with efficient memory operations + signal_padded = MtlArray{T}(undef, fft_sizes) + kernel_padded = MtlArray{T}(undef, fft_sizes) + + # Copy data to padded arrays first + signal_ranges = ntuple(i -> 1:signal_sizes[i], N) + kernel_ranges = ntuple(i -> 1:kernel_sizes[i], N) + + signal_padded[signal_ranges...] = signal + kernel_padded[kernel_ranges...] = kernel + + # Zero only the padding regions (not the entire buffer) + _zero_padding_regions!(signal_padded, signal_sizes, fft_sizes, dims_tuple) + _zero_padding_regions!(kernel_padded, kernel_sizes, fft_sizes, dims_tuple) + + # FFT along specified dimensions (use rfft for real inputs) + S = rfft(signal_padded, dims_tuple) + K = rfft(kernel_padded, dims_tuple) + + # Multiply in frequency domain (in-place to avoid allocation) + S .*= K + + # Inverse FFT + # For irfft, we need the output size of the first transformed dimension + first_dim = minimum(dims_tuple) + y = irfft(S, fft_sizes[first_dim], dims_tuple) + + # Extract appropriate region + return _extract_conv_result(y, output_sizes, full_sizes, mode, dims_tuple) +end + +# Complex N-D version +function conv_fft( + signal::MtlArray{Complex{T}, N}, kernel::MtlArray{Complex{T}, N}; + dims::Union{Int, Tuple{Vararg{Int}}} = 1, mode::Symbol = :full + ) where {T <: Union{Float32, Float16}, N} + dims_tuple = dims isa Int ? (dims,) : Tuple(dims) + + for d in dims_tuple + 1 <= d <= N || + throw(ArgumentError("Invalid dimension $d for array with $N dimensions")) + end + + signal_sizes = size(signal) + kernel_sizes = size(kernel) + + full_sizes = ntuple(N) do i + if i in dims_tuple + signal_sizes[i] + kernel_sizes[i] - 1 + else + signal_sizes[i] + end + end + + output_sizes = ntuple(N) do i + if i in dims_tuple + _conv_output_size(signal_sizes[i], kernel_sizes[i], mode) + else + signal_sizes[i] + end + end + + fft_sizes = ntuple(N) do i + if i in dims_tuple + nextfastfft(full_sizes[i]) + else + signal_sizes[i] + end + end + + # Pad signal and kernel with efficient memory operations + signal_padded = MtlArray{Complex{T}}(undef, fft_sizes) + kernel_padded = MtlArray{Complex{T}}(undef, fft_sizes) + + # Copy data to padded arrays first + signal_ranges = ntuple(i -> 1:signal_sizes[i], N) + kernel_ranges = ntuple(i -> 1:kernel_sizes[i], N) + + signal_padded[signal_ranges...] = signal + kernel_padded[kernel_ranges...] = kernel + + # Zero only the padding regions (not the entire buffer) + _zero_padding_regions!(signal_padded, signal_sizes, fft_sizes, dims_tuple) + _zero_padding_regions!(kernel_padded, kernel_sizes, fft_sizes, dims_tuple) + + S = fft(signal_padded, dims_tuple) + K = fft(kernel_padded, dims_tuple) + + # Multiply in frequency domain (in-place to avoid allocation) + S .*= K + + y = ifft(S, dims_tuple) + + return _extract_conv_result(y, output_sizes, full_sizes, mode, dims_tuple) +end + +# ============================================================================ +# Cross-correlation +# ============================================================================ + +# GPU-friendly reverse along specified dimensions +# Uses broadcasting to avoid scalar indexing +function _gpu_reverse(v::MtlArray{T, 1}) where {T} + n = length(v) + return v[n:-1:1] +end + +function _gpu_reverse(v::MtlArray{T, N}, dims::Tuple) where {T, N} + # Build index arrays for each dimension + indices = ntuple(N) do i + if i in dims + size(v, i):-1:1 + else + 1:size(v, i) + end + end + return v[indices...] +end + +""" + xcorr(u::MtlArray, v::MtlArray; dims=1, mode=:full) + +Compute cross-correlation of `u` and `v` using FFT. + +Cross-correlation is related to convolution by: + xcorr(u, v) = conv(u, reverse(conj(v))) + +For real signals, this simplifies to: + xcorr(u, v) = conv(u, reverse(v)) + +# Arguments +- `u`, `v`: Input arrays +- `dims`: Dimension(s) along which to compute correlation +- `mode`: Output mode (`:full`, `:same`, or `:valid`) + +# Example +```julia +u = MtlVector(randn(Float32, 1000)) +v = MtlVector(randn(Float32, 100)) +r = xcorr(u, v) # Cross-correlation +``` +""" +function xcorr( + u::MtlArray{T, N}, v::MtlArray{T, N}; + dims::Union{Int, Tuple{Vararg{Int}}} = 1, mode::Symbol = :full + ) where {T <: Union{Float32, Float16}, N} + # For real signals: xcorr(u, v) = conv(u, reverse(v, dims=dims)) + dims_tuple = dims isa Int ? (dims,) : Tuple(dims) + v_reversed = N == 1 ? _gpu_reverse(v) : _gpu_reverse(v, dims_tuple) + # 1D conv_fft doesn't take dims argument + if N == 1 + return conv_fft(u, v_reversed; mode = mode) + else + return conv_fft(u, v_reversed; dims = dims, mode = mode) + end +end + +function xcorr( + u::MtlArray{Complex{T}, N}, v::MtlArray{Complex{T}, N}; + dims::Union{Int, Tuple{Vararg{Int}}} = 1, mode::Symbol = :full + ) where {T <: Union{Float32, Float16}, N} + # For complex signals: xcorr(u, v) = conv(u, reverse(conj(v), dims=dims)) + dims_tuple = dims isa Int ? (dims,) : Tuple(dims) + v_conj = conj(v) + v_conj_reversed = N == 1 ? _gpu_reverse(v_conj) : _gpu_reverse(v_conj, dims_tuple) + # 1D conv_fft doesn't take dims argument + if N == 1 + return conv_fft(u, v_conj_reversed; mode = mode) + else + return conv_fft(u, v_conj_reversed; dims = dims, mode = mode) + end +end + +# ============================================================================ +# In-place convolution (output pre-allocated) +# ============================================================================ + +""" + conv_fft!(output, signal, kernel; dims=1, mode=:full) + +Compute convolution and store result in pre-allocated `output` array. + +The output array must have the correct size for the specified mode. +""" +function conv_fft!( + output::MtlArray{T, N}, signal::MtlArray{T, N}, kernel::MtlArray{T, N}; + dims::Union{Int, Tuple{Vararg{Int}}} = 1, mode::Symbol = :full + ) where {T, N} + result = conv_fft(signal, kernel; dims = dims, mode = mode) + @assert size(output) == size(result) "Output size $(size(output)) does not match expected size $(size(result))" + copyto!(output, result) + return output +end + +# ============================================================================ +# MPS Direct Convolution (for small kernels) +# ============================================================================ +# +# Uses MPSGraph's convolution2D operation for direct (non-FFT) convolution. +# This is optimized for small kernels (3×3, 5×5, 7×7) where FFT overhead dominates. +# +# Note: MPSGraph convolution expects 4D tensors in NHWC or NCHW format: +# - N = batch size +# - H = height +# - W = width +# - C = channels +# +# For signal processing, we treat 2D arrays as single-channel images with batch size 1. + +export conv_direct, imfilter + +""" + conv_direct(image::MtlMatrix, kernel::MtlMatrix; mode=:same, padding=:zeros) + +Compute 2D convolution using MPS direct convolution (optimized for small kernels). + +This function is optimized for small kernels (3×3, 5×5, 7×7) where it outperforms +FFT-based convolution. For large kernels, use `conv_fft` instead. + +# Arguments +- `image`: 2D input image (H×W) +- `kernel`: 2D convolution kernel (Kh×Kw) +- `mode`: Output size mode + - `:same` (default): Output has same size as input + - `:valid`: Only fully overlapping region + - `:full`: Full convolution output (not natively supported, falls back to FFT) +- `padding`: Padding type for `:same` mode + - `:zeros` (default): Zero padding + +# Returns +MtlMatrix with the convolution result. + +# Example +```julia +image = MtlMatrix(randn(Float32, 256, 256)) +kernel = MtlMatrix(Float32[ + 1 0 -1 + 2 0 -2 + 1 0 -1 +] ./ 8) # Sobel edge detector +edges = conv_direct(image, kernel) +``` + +# Notes +- For 3×3 kernels on 256×256 images, expect ~10-50x speedup over FFT +- Kernel is flipped internally to match mathematical convolution definition +- Currently supports Float32 and Float16 only +""" +function conv_direct( + image::MtlMatrix{T}, kernel::MtlMatrix{T}; + mode::Symbol = :same, padding::Symbol = :zeros + ) where {T <: Union{Float32, Float16}} + # For :full mode, fall back to FFT + if mode == :full + return conv_fft(image, kernel; dims = (1, 2), mode = :full) + end + + H, W = size(image) + Kh, Kw = size(kernel) + + # Validate kernel size (MPS works best with odd-sized kernels) + if Kh % 2 == 0 || Kw % 2 == 0 + @warn "Even-sized kernels may have unexpected centering. Odd sizes (3×3, 5×5, 7×7) recommended." maxlog = 1 + end + + # Flip kernel for mathematical convolution (MPS does correlation by default) + kernel_flipped = kernel[end:-1:1, end:-1:1] + + # Compute padding for :same mode + # Note: MPSGraph padding is (top, bottom) for Y, (left, right) for X + # In NHWC layout: H is the 2nd dim (padTop/padBottom), W is the 3rd dim (padLeft/padRight) + if mode == :same + # Symmetric padding to maintain size + pad_top = (Kh - 1) ÷ 2 + pad_bottom = Kh - 1 - pad_top + pad_left = (Kw - 1) ÷ 2 + pad_right = Kw - 1 - pad_left + elseif mode == :valid + pad_top = pad_bottom = pad_left = pad_right = 0 + else + throw(ArgumentError("Unknown mode: $mode. Use :same, :valid, or :full")) + end + + # Convert 2D arrays to 4D tensors for MPSGraph convolution + # Due to shape reversal in placeholderTensor (Julia shape is reversed for MPSGraph), + # we need to create 4D arrays where the Julia dimensions map correctly after reversal. + # + # We'll use NHWC layout with shapes that account for the reversal: + # - Julia shape (a, b, c, d) → MPSGraph shape (d, c, b, a) + # + # For NHWC image (N=1, H, W, C=1), MPSGraph expects shape (1, H, W, 1) + # So Julia must have shape (1, W, H, 1) which after reversal gives MPSGraph (1, H, W, 1) + # + # For HWIO kernel (Kh, Kw, Cin=1, Cout=1), MPSGraph expects shape (Kh, Kw, 1, 1) + # So Julia must have shape (1, 1, Kw, Kh) which after reversal gives MPSGraph (Kh, Kw, 1, 1) + + # Transpose the image (H, W) → (W, H) so that after reshape and reversal it matches + # Then add batch and channel dimensions + image_transposed = permutedims(image, (2, 1)) # (W, H) + image_4d = reshape(image_transposed, 1, W, H, 1) # Julia: (1, W, H, 1) → MPSGraph: (1, H, W, 1) + + # For kernel: transpose (Kh, Kw) → (Kw, Kh) then reshape + kernel_transposed = permutedims(kernel_flipped, (2, 1)) # (Kw, Kh) + kernel_4d = reshape(kernel_transposed, 1, 1, Kw, Kh) # Julia: (1, 1, Kw, Kh) → MPSGraph: (Kh, Kw, 1, 1) + + # Output size + if mode == :same + out_h, out_w = H, W + else # :valid + out_h = H - Kh + 1 + out_w = W - Kw + 1 + end + + # Create output array with reversed dimensions + # MPSGraph will produce (1, out_h, out_w, 1), which we specify as Julia (1, out_w, out_h, 1) + output = MtlArray{T}(undef, 1, out_w, out_h, 1) + + # Build and execute MPSGraph + @autoreleasepool begin + _conv2d_mpsgraph!(output, image_4d, kernel_4d, pad_top, pad_bottom, pad_left, pad_right) + end + + # Extract 2D result and transpose back to (H, W) + result_transposed = reshape(output, out_w, out_h) # (out_w, out_h) + return permutedims(result_transposed, (2, 1)) # (out_h, out_w) = (H, W) +end + +""" +Internal function to execute MPSGraph 2D convolution. +""" +function _conv2d_mpsgraph!( + output::MtlArray{T, 4}, image::MtlArray{T, 4}, kernel::MtlArray{T, 4}, + pad_top::Int, pad_bottom::Int, pad_left::Int, pad_right::Int + ) where {T} + graph = MPSGraph() + + # Create placeholders + placeImage = placeholderTensor(graph, size(image), T) + placeKernel = placeholderTensor(graph, size(kernel), T) + + feeds = Dict{MPSGraphTensor, MPSGraphTensorData}( + placeImage => MPSGraphTensorData(image), + placeKernel => MPSGraphTensorData(kernel) + ) + + # Create convolution descriptor + descriptor = MPSGraphConvolution2DOpDescriptor(; + strideX = 1, strideY = 1, + dilationX = 1, dilationY = 1, + paddingLeft = pad_left, paddingRight = pad_right, + paddingTop = pad_top, paddingBottom = pad_bottom, + paddingStyle = MPSGraphPaddingStyleExplicit, + dataLayout = MPSGraphTensorNamedDataLayoutNHWC, + weightsLayout = MPSGraphTensorNamedDataLayoutHWIO, + groups = 1 + ) + + # Perform convolution + convResult = convolution2DWithSourceTensor(graph, placeImage, placeKernel, descriptor, "conv2d") + + # Create result dictionary + resultdict = Dict{MPSGraphTensor, MPSGraphTensorData}( + convResult => MPSGraphTensorData(output) + ) + + # Execute + cmdbuf = MPSCommandBuffer(Metal.global_queue(device())) + encode!(cmdbuf, graph, NSDictionary(feeds), NSDictionary(resultdict), nil, default_exec_desc()) + commit!(cmdbuf) + wait_completed(cmdbuf) + + return output +end + +""" + imfilter(image::MtlMatrix, kernel::MtlMatrix) + +Apply a filter kernel to an image using direct convolution. + +This is a convenience function following the ImageFiltering.jl interface. +It automatically selects between MPS direct convolution (for small kernels) +and FFT convolution (for large kernels). + +# Arguments +- `image`: 2D input image +- `kernel`: 2D filter kernel + +# Returns +Filtered image with same size as input (`:same` mode). + +# Example +```julia +using Metal + +# Create test image +image = MtlMatrix(randn(Float32, 512, 512)) + +# Gaussian blur (5×5 approximation) +gaussian = MtlMatrix(Float32[ + 1 4 6 4 1 + 4 16 24 16 4 + 6 24 36 24 6 + 4 16 24 16 4 + 1 4 6 4 1 +] ./ 256) + +blurred = imfilter(image, gaussian) + +# Sobel edge detection +sobel_x = MtlMatrix(Float32[-1 0 1; -2 0 2; -1 0 1] ./ 8) +sobel_y = MtlMatrix(Float32[-1 -2 -1; 0 0 0; 1 2 1] ./ 8) +edges_x = imfilter(image, sobel_x) +edges_y = imfilter(image, sobel_y) +edges = sqrt.(edges_x.^2 .+ edges_y.^2) +``` + +# Notes +- For kernels ≤ 11×11, uses MPS direct convolution +- For larger kernels, automatically falls back to FFT convolution +- The kernel is centered on each pixel (like ImageFiltering.jl's `imfilter`) +""" +# Threshold for switching between direct and FFT convolution +# MPS direct convolution is faster for small kernels +const _DIRECT_CONV_THRESHOLD = 11 + +function imfilter(image::MtlMatrix{T}, kernel::MtlMatrix{T}) where {T <: Union{Float32, Float16}} + Kh, Kw = size(kernel) + + if Kh <= _DIRECT_CONV_THRESHOLD && Kw <= _DIRECT_CONV_THRESHOLD + return conv_direct(image, kernel; mode = :same) + else + return conv_fft(image, kernel; dims = (1, 2), mode = :same) + end +end + +# ============================================================================ +# Unified Convolution API (with automatic algorithm selection) +# ============================================================================ +# +# The unified `conv()` function automatically selects the best algorithm: +# - For 2D arrays with small kernels: MPS direct convolution (faster) +# - For 2D arrays with large kernels: FFT convolution +# - For 1D arrays: FFT convolution (no MPS direct 1D support) +# - For N-D arrays: FFT convolution along specified dimensions + +""" + conv(signal::MtlArray, kernel::MtlArray; mode=:full, dims=nothing, algorithm=:auto) + +Compute linear convolution of `signal` and `kernel` with automatic algorithm selection. + +This is the recommended entry point for convolution operations. It automatically +selects between MPS direct convolution (optimized for small kernels) and FFT-based +convolution (better for large kernels or higher dimensions). + +# Arguments +- `signal`: Input signal (1D, 2D, or N-D MtlArray) +- `kernel`: Convolution kernel (same dimensions as signal) +- `mode`: Output size mode + - `:full` (default): Full convolution output + - `:same`: Output has same size as signal (centered) + - `:valid`: Only fully overlapping region +- `dims`: Dimensions along which to convolve + - `nothing` (default): All dimensions for 1D/2D, dim 1 for N-D + - Integer or tuple: Specific dimension(s) +- `algorithm`: Algorithm selection + - `:auto` (default): Automatically select best algorithm + - `:fft`: Force FFT-based convolution + - `:direct`: Force MPS direct convolution (2D only, small kernels) + +# Returns +MtlArray with the convolution result. + +# Algorithm Selection (when `algorithm=:auto`) + +For **2D matrices** with `:same` or `:valid` mode: +- Kernels ≤ 11×11: Uses MPS direct convolution (~8x faster for 3×3) +- Larger kernels: Uses FFT convolution + +For **1D vectors**, **N-D arrays**, or `:full` mode: +- Always uses FFT convolution + +# Examples + +```julia +using Metal, Metal.MPSGraphs + +# 1D signal processing +signal = MtlVector(randn(Float32, 10000)) +kernel = MtlVector(Float32[0.25, 0.5, 0.25]) # Simple smoothing +smoothed = conv(signal, kernel; mode=:same) + +# 2D image filtering (auto-selects direct convolution) +image = MtlMatrix(randn(Float32, 512, 512)) +sobel_x = MtlMatrix(Float32[-1 0 1; -2 0 2; -1 0 1] ./ 8) +edges = conv(image, sobel_x; mode=:same) + +# 2D with large kernel (auto-selects FFT) +large_kernel = MtlMatrix(randn(Float32, 33, 33)) +result = conv(image, large_kernel; mode=:same) + +# Force specific algorithm +result_fft = conv(image, sobel_x; mode=:same, algorithm=:fft) +result_direct = conv(image, sobel_x; mode=:same, algorithm=:direct) +``` + +# Performance Tips + +1. For repeated convolutions with same sizes, use `plan_conv_fft()` or + `get_cached_conv_plan()` for even better performance with FFT. + +2. For small kernels (3×3, 5×5, 7×7), direct convolution is typically + 8-50x faster than FFT. + +3. For large kernels (>15×15), FFT becomes more efficient due to O(n log n) + vs O(n×m) complexity. + +# See Also +- `conv_fft`: Force FFT-based convolution +- `conv_direct`: Force MPS direct convolution (2D only) +- `imfilter`: ImageFiltering.jl-compatible API for 2D filtering +- `xcorr`: Cross-correlation +- `plan_conv_fft`: Pre-computed FFT plan for repeated convolutions +""" +function conv( + signal::MtlVector{T}, kernel::MtlVector{T}; + mode::Symbol = :full, dims = nothing, algorithm::Symbol = :auto + ) where {T <: Union{Float32, Float16}} + # 1D always uses FFT (no MPS direct 1D support) + if algorithm == :direct + throw(ArgumentError("Direct convolution not supported for 1D arrays. Use :auto or :fft.")) + end + return conv_fft(signal, kernel; mode = mode) +end + +# Complex 1D +function conv( + signal::MtlVector{Complex{T}}, kernel::MtlVector{Complex{T}}; + mode::Symbol = :full, dims = nothing, algorithm::Symbol = :auto + ) where {T <: Union{Float32, Float16}} + if algorithm == :direct + throw(ArgumentError("Direct convolution not supported for complex 1D arrays. Use :auto or :fft.")) + end + return conv_fft(signal, kernel; mode = mode) +end + +# 2D with automatic algorithm selection +function conv( + signal::MtlMatrix{T}, kernel::MtlMatrix{T}; + mode::Symbol = :full, dims = nothing, algorithm::Symbol = :auto + ) where {T <: Union{Float32, Float16}} + # Determine dims for FFT (default: both dimensions for 2D) + conv_dims = dims === nothing ? (1, 2) : (dims isa Int ? (dims,) : Tuple(dims)) + + # Check if we should use direct convolution + Kh, Kw = size(kernel) + use_direct = false + + if algorithm == :auto + # Auto-select: use direct for small kernels with :same or :valid mode + # Direct convolution is only supported when convolving all dimensions + if conv_dims == (1, 2) && mode != :full + use_direct = Kh <= _DIRECT_CONV_THRESHOLD && Kw <= _DIRECT_CONV_THRESHOLD + end + elseif algorithm == :direct + # User requested direct convolution + if mode == :full + @warn "Direct convolution doesn't support :full mode. Falling back to FFT." maxlog = 1 + use_direct = false + elseif conv_dims != (1, 2) + throw(ArgumentError("Direct convolution requires convolving all dimensions (dims=(1,2) or nothing).")) + else + use_direct = true + end + elseif algorithm == :fft + use_direct = false + else + throw(ArgumentError("Unknown algorithm: $algorithm. Use :auto, :fft, or :direct.")) + end + + if use_direct + return conv_direct(signal, kernel; mode = mode) + else + return conv_fft(signal, kernel; dims = conv_dims, mode = mode) + end +end + +# Complex 2D +function conv( + signal::MtlMatrix{Complex{T}}, kernel::MtlMatrix{Complex{T}}; + mode::Symbol = :full, dims = nothing, algorithm::Symbol = :auto + ) where {T <: Union{Float32, Float16}} + if algorithm == :direct + throw(ArgumentError("Direct convolution not supported for complex arrays. Use :auto or :fft.")) + end + conv_dims = dims === nothing ? (1, 2) : (dims isa Int ? (dims,) : Tuple(dims)) + return conv_fft(signal, kernel; dims = conv_dims, mode = mode) +end + +# N-D generic (N > 2) +function conv( + signal::MtlArray{T, N}, kernel::MtlArray{T, N}; + mode::Symbol = :full, dims = nothing, algorithm::Symbol = :auto + ) where {T <: Union{Float32, Float16}, N} + # N-D always uses FFT + if algorithm == :direct && N > 2 + throw(ArgumentError("Direct convolution only supported for 2D arrays. Use :auto or :fft.")) + end + conv_dims = dims === nothing ? 1 : (dims isa Int ? (dims,) : Tuple(dims)) + return conv_fft(signal, kernel; dims = conv_dims, mode = mode) +end + +# Complex N-D +function conv( + signal::MtlArray{Complex{T}, N}, kernel::MtlArray{Complex{T}, N}; + mode::Symbol = :full, dims = nothing, algorithm::Symbol = :auto + ) where {T <: Union{Float32, Float16}, N} + if algorithm == :direct + throw(ArgumentError("Direct convolution not supported for complex arrays. Use :auto or :fft.")) + end + conv_dims = dims === nothing ? 1 : (dims isa Int ? (dims,) : Tuple(dims)) + return conv_fft(signal, kernel; dims = conv_dims, mode = mode) +end + +# Note: Batched convolution for 4D tensors can be added later if needed. +# The current implementation focuses on 2D images which covers most use cases. diff --git a/lib/mpsgraphs/operations.jl b/lib/mpsgraphs/operations.jl index 107c9ae31..e88d0a990 100644 --- a/lib/mpsgraphs/operations.jl +++ b/lib/mpsgraphs/operations.jl @@ -74,3 +74,95 @@ Dumps the `graph`. This function is undocumented from Apple so it may stop working at any time. """ dump_graph(graph::MPSGraph) = @objc [graph::id{MPSGraph} dump]::Nothing ## COV_EXCL_LINE + +## Convolution support (used by convolution.jl) + +function concatTensors(graph::MPSGraph, tensors::NSArray, dimension::Int, name = "concat") + obj = @objc [graph::id{MPSGraph} concatTensors:tensors::id{NSArray} + dimension:dimension::NSInteger + name:name::id{NSString}]::id{MPSGraphTensor} + MPSGraphTensor(obj) +end + +""" + convolution2DWithSourceTensor(graph, source, weights, descriptor, name="conv2d") + +2D convolution operation using MPSGraph. + +# Arguments +- `graph`: MPSGraph instance +- `source`: Input tensor in NHWC or NCHW format (depending on descriptor) +- `weights`: Convolution kernel/weights in OIHW or HWIO format (depending on descriptor) +- `descriptor`: MPSGraphConvolution2DOpDescriptor configuring stride, padding, dilation, etc. +- `name`: Operation name for debugging + +# Returns +MPSGraphTensor with the convolution result. +""" +function convolution2DWithSourceTensor( + graph::MPSGraph, source::MPSGraphTensor, weights::MPSGraphTensor, + descriptor::MPSGraphConvolution2DOpDescriptor, name = "conv2d" + ) + obj = @objc [graph::id{MPSGraph} convolution2DWithSourceTensor:source::id{MPSGraphTensor} + weightsTensor:weights::id{MPSGraphTensor} + descriptor:descriptor::id{MPSGraphConvolution2DOpDescriptor} + name:name::id{NSString}]::id{MPSGraphTensor} + MPSGraphTensor(obj) +end + +""" + MPSGraphConvolution2DOpDescriptor(; + strideX=1, strideY=1, + dilationX=1, dilationY=1, + paddingLeft=0, paddingRight=0, paddingTop=0, paddingBottom=0, + paddingStyle=MPSGraphPaddingStyleExplicit, + dataLayout=MPSGraphTensorNamedDataLayoutNHWC, + weightsLayout=MPSGraphTensorNamedDataLayoutHWIO, + groups=1 + ) + +Create a 2D convolution operation descriptor. + +# Arguments +- `strideX`, `strideY`: Stride in X and Y directions +- `dilationX`, `dilationY`: Dilation rate in X and Y directions +- `paddingLeft/Right/Top/Bottom`: Explicit padding values +- `paddingStyle`: One of: + - `MPSGraphPaddingStyleExplicit` (default) - use explicit padding values + - `MPSGraphPaddingStyleTF_VALID` - no padding + - `MPSGraphPaddingStyleTF_SAME` - pad to keep output same size as input +- `dataLayout`: Input/output tensor layout (NHWC or NCHW) +- `weightsLayout`: Kernel tensor layout (HWIO or OIHW) +- `groups`: Number of groups for grouped convolution +""" +function MPSGraphConvolution2DOpDescriptor(; + strideX::Integer = 1, strideY::Integer = 1, + dilationX::Integer = 1, dilationY::Integer = 1, + paddingLeft::Integer = 0, paddingRight::Integer = 0, + paddingTop::Integer = 0, paddingBottom::Integer = 0, + paddingStyle::MPSGraphPaddingStyle = MPSGraphPaddingStyleExplicit, + dataLayout::MPSGraphTensorNamedDataLayout = MPSGraphTensorNamedDataLayoutNHWC, + weightsLayout::MPSGraphTensorNamedDataLayout = MPSGraphTensorNamedDataLayoutHWIO, + groups::Integer = 1 + ) + # Create descriptor via alloc/init + desc = @objc [MPSGraphConvolution2DOpDescriptor alloc]::id{MPSGraphConvolution2DOpDescriptor} + desc = @objc [desc::id{MPSGraphConvolution2DOpDescriptor} init]::id{MPSGraphConvolution2DOpDescriptor} + descriptor = MPSGraphConvolution2DOpDescriptor(desc) + + # Set properties + descriptor.strideInX = UInt64(strideX) + descriptor.strideInY = UInt64(strideY) + descriptor.dilationRateInX = UInt64(dilationX) + descriptor.dilationRateInY = UInt64(dilationY) + descriptor.paddingLeft = UInt64(paddingLeft) + descriptor.paddingRight = UInt64(paddingRight) + descriptor.paddingTop = UInt64(paddingTop) + descriptor.paddingBottom = UInt64(paddingBottom) + descriptor.paddingStyle = paddingStyle + descriptor.dataLayout = dataLayout + descriptor.weightsLayout = weightsLayout + descriptor.groups = UInt64(groups) + + return descriptor +end diff --git a/perf/Project.toml b/perf/Project.toml index decfbe75f..4abbeca91 100644 --- a/perf/Project.toml +++ b/perf/Project.toml @@ -1,6 +1,9 @@ [deps] BenchmarkTools = "6e4b80f9-dd63-53aa-95a3-0cdb28fa8baf" +DSP = "717857b8-e6f2-59f4-9121-6e50c889abd2" HTTP = "cd3eb016-35fb-5094-929b-558a96fad6f3" +ImageCore = "a09fc81d-aa75-5fe9-8630-4744c3626534" +ImageFiltering = "6a3955dd-da59-5b1f-98d4-e7296123deb5" JSON = "682c06a0-de6a-54ab-a142-c8b1cf79cde6" Metal = "dde4c033-4e86-420c-a63e-0dd931031962" StableRNGs = "860ef19b-820b-49d6-a774-d7a799459cd3" diff --git a/perf/benchmark_fused_comprehensive.jl b/perf/benchmark_fused_comprehensive.jl new file mode 100644 index 000000000..8f4ef910a --- /dev/null +++ b/perf/benchmark_fused_comprehensive.jl @@ -0,0 +1,324 @@ +using Metal +using Metal.MPSGraphs: conv_fft, conv_fft_fused +using DSP +using Statistics + +println("=" ^ 70) +println("COMPREHENSIVE FUSED CONVOLUTION BENCHMARKS") +println("=" ^ 70) +println("Note: conv_fft() now uses fused implementation automatically") +println() + +# ============================================================================ +# 1D BENCHMARKS +# ============================================================================ + +println("=" ^ 70) +println("1D CONVOLUTION: GPU (Fused) vs CPU (DSP.jl)") +println("=" ^ 70) + +configs_1d = [ + (10_000, 100), + (50_000, 100), + (100_000, 100), + (100_000, 500), + (500_000, 500), + (1_000_000, 500), + (1_000_000, 1000), + (5_000_000, 1000), +] + +results_1d = [] + +for (signal_size, kernel_size) in configs_1d + # Create test data + signal_cpu = rand(Float32, signal_size) + kernel_cpu = rand(Float32, kernel_size) + signal_gpu = MtlVector(signal_cpu) + kernel_gpu = MtlVector(kernel_cpu) + + # Warmup + _ = conv_fft(signal_gpu, kernel_gpu) + Metal.synchronize() + + # Benchmark GPU + n_iters = 10 + times_gpu = Float64[] + for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + _ = conv_fft(signal_gpu, kernel_gpu) + Metal.synchronize() + end + push!(times_gpu, t * 1000) + end + + # Benchmark CPU + cpu_iters = signal_size > 1_000_000 ? 3 : 5 + times_cpu = Float64[] + for _ in 1:cpu_iters + t = @elapsed begin + _ = DSP.conv(signal_cpu, kernel_cpu) + end + push!(times_cpu, t * 1000) + end + + med_gpu = median(times_gpu) + med_cpu = median(times_cpu) + speedup = med_cpu / med_gpu + winner = speedup >= 1.0 ? "GPU" : "CPU" + + push!(results_1d, (signal_size, kernel_size, med_gpu, med_cpu, speedup, winner)) + println("Signal=$(signal_size), Kernel=$(kernel_size): GPU=$(round(med_gpu, digits=2))ms, CPU=$(round(med_cpu, digits=2))ms, Speedup=$(round(speedup, digits=2))x [$winner]") +end + +# ============================================================================ +# 2D BENCHMARKS +# ============================================================================ + +println("\n" * "=" ^ 70) +println("2D CONVOLUTION: GPU (Fused) vs CPU (imfilter-style)") +println("=" ^ 70) + +configs_2d = [ + ((128, 128), (5, 5)), + ((256, 256), (5, 5)), + ((256, 256), (15, 15)), + ((512, 512), (5, 5)), + ((512, 512), (15, 15)), + ((1024, 1024), (5, 5)), + ((1024, 1024), (15, 15)), + ((2048, 2048), (15, 15)), +] + +results_2d = [] + +for (image_size, kernel_size) in configs_2d + # Create test data + image_cpu = rand(Float32, image_size...) + kernel_cpu = rand(Float32, kernel_size...) + image_gpu = MtlMatrix(image_cpu) + kernel_gpu = MtlMatrix(kernel_cpu) + + # Warmup + _ = conv_fft(image_gpu, kernel_gpu; dims=(1,2)) + Metal.synchronize() + + # Benchmark GPU + n_iters = 10 + times_gpu = Float64[] + for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + _ = conv_fft(image_gpu, kernel_gpu; dims=(1,2)) + Metal.synchronize() + end + push!(times_gpu, t * 1000) + end + + # Benchmark CPU (using DSP.conv for 2D as reference) + cpu_iters = prod(image_size) > 1_000_000 ? 3 : 5 + times_cpu = Float64[] + for _ in 1:cpu_iters + t = @elapsed begin + _ = DSP.conv(image_cpu, kernel_cpu) + end + push!(times_cpu, t * 1000) + end + + med_gpu = median(times_gpu) + med_cpu = median(times_cpu) + speedup = med_cpu / med_gpu + winner = speedup >= 1.0 ? "GPU" : "CPU" + + push!(results_2d, (image_size, kernel_size, med_gpu, med_cpu, speedup, winner)) + println("Image=$(image_size), Kernel=$(kernel_size): GPU=$(round(med_gpu, digits=2))ms, CPU=$(round(med_cpu, digits=2))ms, Speedup=$(round(speedup, digits=2))x [$winner]") +end + +# ============================================================================ +# 3D BENCHMARKS +# ============================================================================ + +println("\n" * "=" ^ 70) +println("3D CONVOLUTION: GPU (Fused) vs CPU (DSP.conv)") +println("=" ^ 70) + +configs_3d = [ + ((32, 32, 32), (3, 3, 3)), + ((64, 64, 64), (3, 3, 3)), + ((64, 64, 64), (5, 5, 5)), + ((128, 128, 128), (3, 3, 3)), + ((128, 128, 128), (5, 5, 5)), + ((256, 256, 64), (5, 5, 5)), + ((256, 256, 128), (5, 5, 5)), +] + +results_3d = [] + +for (volume_size, kernel_size) in configs_3d + # Create test data + volume_cpu = rand(Float32, volume_size...) + kernel_cpu = rand(Float32, kernel_size...) + volume_gpu = MtlArray(volume_cpu) + kernel_gpu = MtlArray(kernel_cpu) + + # Warmup + _ = conv_fft(volume_gpu, kernel_gpu; dims=(1,2,3)) + Metal.synchronize() + + # Benchmark GPU + n_iters = 8 + times_gpu = Float64[] + for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + _ = conv_fft(volume_gpu, kernel_gpu; dims=(1,2,3)) + Metal.synchronize() + end + push!(times_gpu, t * 1000) + end + + # Benchmark CPU + cpu_iters = prod(volume_size) > 500_000 ? 2 : 3 + times_cpu = Float64[] + for _ in 1:cpu_iters + t = @elapsed begin + _ = DSP.conv(volume_cpu, kernel_cpu) + end + push!(times_cpu, t * 1000) + end + + med_gpu = median(times_gpu) + med_cpu = median(times_cpu) + speedup = med_cpu / med_gpu + winner = speedup >= 1.0 ? "GPU" : "CPU" + + push!(results_3d, (volume_size, kernel_size, med_gpu, med_cpu, speedup, winner)) + println("Volume=$(volume_size), Kernel=$(kernel_size): GPU=$(round(med_gpu, digits=2))ms, CPU=$(round(med_cpu, digits=2))ms, Speedup=$(round(speedup, digits=2))x [$winner]") +end + +# ============================================================================ +# THROUGHPUT BENCHMARK +# ============================================================================ + +println("\n" * "=" ^ 70) +println("THROUGHPUT BENCHMARK (Async Pipeline, batch=50)") +println("=" ^ 70) + +# 1D throughput +signal_1d = MtlVector(rand(Float32, 1_000_000)) +kernel_1d = MtlVector(rand(Float32, 500)) +signal_1d_cpu = rand(Float32, 1_000_000) +kernel_1d_cpu = rand(Float32, 500) + +batch = 50 +Metal.synchronize() +t_gpu_1d = @elapsed begin + for _ in 1:batch + _ = conv_fft(signal_1d, kernel_1d) + end + Metal.synchronize() +end +tput_gpu_1d = batch / t_gpu_1d + +t_cpu_1d = @elapsed begin + for _ in 1:batch + _ = DSP.conv(signal_1d_cpu, kernel_1d_cpu) + end +end +tput_cpu_1d = batch / t_cpu_1d + +println("\n1D (1M × 500):") +println(" GPU: $(round(tput_gpu_1d, digits=1)) ops/sec ($(round(1000/tput_gpu_1d, digits=2)) ms/op)") +println(" CPU: $(round(tput_cpu_1d, digits=1)) ops/sec ($(round(1000/tput_cpu_1d, digits=2)) ms/op)") +println(" Speedup: $(round(tput_gpu_1d / tput_cpu_1d, digits=2))x") + +# 2D throughput +image_2d = MtlMatrix(rand(Float32, 512, 512)) +kernel_2d = MtlMatrix(rand(Float32, 15, 15)) +image_2d_cpu = rand(Float32, 512, 512) +kernel_2d_cpu = rand(Float32, 15, 15) + +Metal.synchronize() +t_gpu_2d = @elapsed begin + for _ in 1:batch + _ = conv_fft(image_2d, kernel_2d; dims=(1,2)) + end + Metal.synchronize() +end +tput_gpu_2d = batch / t_gpu_2d + +t_cpu_2d = @elapsed begin + for _ in 1:batch + _ = DSP.conv(image_2d_cpu, kernel_2d_cpu) + end +end +tput_cpu_2d = batch / t_cpu_2d + +println("\n2D (512×512 × 15×15):") +println(" GPU: $(round(tput_gpu_2d, digits=1)) ops/sec ($(round(1000/tput_gpu_2d, digits=2)) ms/op)") +println(" CPU: $(round(tput_cpu_2d, digits=1)) ops/sec ($(round(1000/tput_cpu_2d, digits=2)) ms/op)") +println(" Speedup: $(round(tput_gpu_2d / tput_cpu_2d, digits=2))x") + +# 3D throughput +volume_3d = MtlArray(rand(Float32, 64, 64, 64)) +kernel_3d = MtlArray(rand(Float32, 5, 5, 5)) +volume_3d_cpu = rand(Float32, 64, 64, 64) +kernel_3d_cpu = rand(Float32, 5, 5, 5) + +batch_3d = 20 +Metal.synchronize() +t_gpu_3d = @elapsed begin + for _ in 1:batch_3d + _ = conv_fft(volume_3d, kernel_3d; dims=(1,2,3)) + end + Metal.synchronize() +end +tput_gpu_3d = batch_3d / t_gpu_3d + +t_cpu_3d = @elapsed begin + for _ in 1:batch_3d + _ = DSP.conv(volume_3d_cpu, kernel_3d_cpu) + end +end +tput_cpu_3d = batch_3d / t_cpu_3d + +println("\n3D (64³ × 5³):") +println(" GPU: $(round(tput_gpu_3d, digits=1)) ops/sec ($(round(1000/tput_gpu_3d, digits=2)) ms/op)") +println(" CPU: $(round(tput_cpu_3d, digits=1)) ops/sec ($(round(1000/tput_cpu_3d, digits=2)) ms/op)") +println(" Speedup: $(round(tput_gpu_3d / tput_cpu_3d, digits=2))x") + +# ============================================================================ +# SUMMARY TABLES +# ============================================================================ + +println("\n" * "=" ^ 70) +println("SUMMARY TABLES (for documentation)") +println("=" ^ 70) + +println("\n### 1D Convolution Results") +println("| Signal | Kernel | GPU (ms) | CPU (ms) | Speedup | Winner |") +println("|--------|--------|----------|----------|---------|--------|") +for (ss, ks, gpu, cpu, speedup, winner) in results_1d + speedup_str = speedup >= 1.0 ? "**$(round(speedup, digits=2))x**" : "$(round(speedup, digits=2))x" + winner_str = winner == "GPU" ? "**GPU**" : "CPU" + println("| $(ss) | $(ks) | $(round(gpu, digits=2)) | $(round(cpu, digits=2)) | $(speedup_str) | $(winner_str) |") +end + +println("\n### 2D Convolution Results") +println("| Image | Kernel | GPU (ms) | CPU (ms) | Speedup | Winner |") +println("|-------|--------|----------|----------|---------|--------|") +for (is, ks, gpu, cpu, speedup, winner) in results_2d + speedup_str = speedup >= 1.0 ? "**$(round(speedup, digits=2))x**" : "$(round(speedup, digits=2))x" + winner_str = winner == "GPU" ? "**GPU**" : "CPU" + println("| $(is[1])×$(is[2]) | $(ks[1])×$(ks[2]) | $(round(gpu, digits=2)) | $(round(cpu, digits=2)) | $(speedup_str) | $(winner_str) |") +end + +println("\n### 3D Convolution Results") +println("| Volume | Kernel | GPU (ms) | CPU (ms) | Speedup | Winner |") +println("|--------|--------|----------|----------|---------|--------|") +for (vs, ks, gpu, cpu, speedup, winner) in results_3d + speedup_str = speedup >= 1.0 ? "**$(round(speedup, digits=2))x**" : "$(round(speedup, digits=2))x" + winner_str = winner == "GPU" ? "**GPU**" : "CPU" + println("| $(vs[1])×$(vs[2])×$(vs[3]) | $(ks[1])×$(ks[2])×$(ks[3]) | $(round(gpu, digits=2)) | $(round(cpu, digits=2)) | $(speedup_str) | $(winner_str) |") +end diff --git a/perf/benchmark_fused_conv.jl b/perf/benchmark_fused_conv.jl new file mode 100644 index 000000000..7388fc804 --- /dev/null +++ b/perf/benchmark_fused_conv.jl @@ -0,0 +1,96 @@ +using Metal +using Metal.MPSGraphs: conv_fft, conv_fft_fused +using DSP +using Statistics + +println("=" ^ 60) +println("LATENCY BENCHMARK: Fused vs Existing vs CPU") +println("=" ^ 60) + +# Test configurations +configs = [ + (10_000, 100), + (100_000, 100), + (100_000, 500), + (500_000, 500), + (1_000_000, 500), + (1_000_000, 1000), +] + +results = [] + +for (signal_size, kernel_size) in configs + println("\n--- Signal: $(signal_size), Kernel: $(kernel_size) ---") + + # Create test data + signal_cpu = rand(Float32, signal_size) + kernel_cpu = rand(Float32, kernel_size) + signal_gpu = MtlVector(signal_cpu) + kernel_gpu = MtlVector(kernel_cpu) + + # Warmup + _ = conv_fft(signal_gpu, kernel_gpu) + _ = conv_fft_fused(signal_gpu, kernel_gpu) + Metal.synchronize() + + # Benchmark existing implementation + n_iters = 10 + times_existing = Float64[] + for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + _ = conv_fft(signal_gpu, kernel_gpu) + Metal.synchronize() + end + push!(times_existing, t * 1000) # ms + end + + # Benchmark fused implementation + times_fused = Float64[] + for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + _ = conv_fft_fused(signal_gpu, kernel_gpu) + Metal.synchronize() + end + push!(times_fused, t * 1000) # ms + end + + # Benchmark CPU (fewer iterations for large sizes) + cpu_iters = signal_size > 500_000 ? 3 : 5 + times_cpu = Float64[] + for _ in 1:cpu_iters + t = @elapsed begin + _ = DSP.conv(signal_cpu, kernel_cpu) + end + push!(times_cpu, t * 1000) # ms + end + + med_existing = median(times_existing) + med_fused = median(times_fused) + med_cpu = median(times_cpu) + + speedup_fused_vs_existing = med_existing / med_fused + speedup_fused_vs_cpu = med_cpu / med_fused + speedup_existing_vs_cpu = med_cpu / med_existing + + println(" Existing GPU: $(round(med_existing, digits=3)) ms") + println(" Fused GPU: $(round(med_fused, digits=3)) ms") + println(" CPU (DSP): $(round(med_cpu, digits=3)) ms") + println(" Fused vs Existing: $(round(speedup_fused_vs_existing, digits=2))x") + println(" Fused vs CPU: $(round(speedup_fused_vs_cpu, digits=2))x") + println(" Existing vs CPU: $(round(speedup_existing_vs_cpu, digits=2))x") + + push!(results, (signal_size, kernel_size, med_existing, med_fused, med_cpu)) +end + +println("\n" * "=" ^ 60) +println("SUMMARY TABLE") +println("=" ^ 60) +println("Signal | Kernel | Existing | Fused | CPU | Fused/Exist | Fused/CPU") +println("-" ^ 80) +for (ss, ks, exist, fused, cpu) in results + speedup1 = round(exist / fused, digits=2) + speedup2 = round(cpu / fused, digits=2) + println("$(lpad(ss, 9)) | $(lpad(ks, 6)) | $(lpad(round(exist, digits=2), 8)) | $(lpad(round(fused, digits=2), 7)) | $(lpad(round(cpu, digits=2), 7)) | $(lpad(speedup1, 11))x | $(lpad(speedup2, 9))x") +end diff --git a/perf/benchmark_inline_pad.jl b/perf/benchmark_inline_pad.jl new file mode 100644 index 000000000..e6701f55b --- /dev/null +++ b/perf/benchmark_inline_pad.jl @@ -0,0 +1,200 @@ +using Metal +using Metal.MPSGraphs: conv_fft_inline_pad, conv_fft_fused +using DSP +using Statistics + +println("=" ^ 70) +println("INLINE PADDING BENCHMARK: Fused vs Inline-Pad vs CPU") +println("=" ^ 70) + +configs_1d = [ + (10_000, 100), + (50_000, 100), + (100_000, 100), + (100_000, 500), + (500_000, 500), + (1_000_000, 500), +] + +println("\n1D CONVOLUTION") +println("-" ^ 70) +println("Signal | Kernel | Fused(ms) | Inline(ms) | Speedup | CPU(ms)") +println("-" ^ 70) + +for (signal_size, kernel_size) in configs_1d + signal_cpu = rand(Float32, signal_size) + kernel_cpu = rand(Float32, kernel_size) + signal_gpu = MtlVector(signal_cpu) + kernel_gpu = MtlVector(kernel_cpu) + + # Warmup + _ = conv_fft_fused(signal_gpu, kernel_gpu) + _ = conv_fft_inline_pad(signal_gpu, kernel_gpu) + Metal.synchronize() + + n_iters = 10 + + # Benchmark fused + times_fused = Float64[] + for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + _ = conv_fft_fused(signal_gpu, kernel_gpu) + Metal.synchronize() + end + Base.push!(times_fused, t * 1000) + end + + # Benchmark inline pad + times_inline = Float64[] + for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + _ = conv_fft_inline_pad(signal_gpu, kernel_gpu) + Metal.synchronize() + end + Base.push!(times_inline, t * 1000) + end + + # Benchmark CPU + cpu_iters = signal_size > 500_000 ? 3 : 5 + times_cpu = Float64[] + for _ in 1:cpu_iters + t = @elapsed begin + _ = DSP.conv(signal_cpu, kernel_cpu) + end + Base.push!(times_cpu, t * 1000) + end + + med_fused = median(times_fused) + med_inline = median(times_inline) + med_cpu = median(times_cpu) + speedup = med_fused / med_inline + + println("$(lpad(signal_size, 10)) | $(lpad(kernel_size, 6)) | $(lpad(round(med_fused, digits=2), 9)) | $(lpad(round(med_inline, digits=2), 10)) | $(lpad(round(speedup, digits=2), 7))x | $(lpad(round(med_cpu, digits=2), 6))") +end + +println("\n2D CONVOLUTION") +println("-" ^ 70) +println("Image | Kernel | Fused(ms) | Inline(ms) | Speedup | CPU(ms)") +println("-" ^ 70) + +configs_2d = [ + ((128, 128), (5, 5)), + ((256, 256), (5, 5)), + ((256, 256), (15, 15)), + ((512, 512), (15, 15)), + ((1024, 1024), (15, 15)), +] + +for (image_size, kernel_size) in configs_2d + image_cpu = rand(Float32, image_size...) + kernel_cpu = rand(Float32, kernel_size...) + image_gpu = MtlMatrix(image_cpu) + kernel_gpu = MtlMatrix(kernel_cpu) + + # Warmup + _ = conv_fft_fused(image_gpu, kernel_gpu) + _ = conv_fft_inline_pad(image_gpu, kernel_gpu) + Metal.synchronize() + + n_iters = 10 + + times_fused = Float64[] + for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + _ = conv_fft_fused(image_gpu, kernel_gpu) + Metal.synchronize() + end + Base.push!(times_fused, t * 1000) + end + + times_inline = Float64[] + for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + _ = conv_fft_inline_pad(image_gpu, kernel_gpu) + Metal.synchronize() + end + Base.push!(times_inline, t * 1000) + end + + cpu_iters = prod(image_size) > 500_000 ? 3 : 5 + times_cpu = Float64[] + for _ in 1:cpu_iters + t = @elapsed DSP.conv(image_cpu, kernel_cpu) + Base.push!(times_cpu, t * 1000) + end + + med_fused = median(times_fused) + med_inline = median(times_inline) + med_cpu = median(times_cpu) + speedup = med_fused / med_inline + + image_str = "$(image_size[1])x$(image_size[2])" + kernel_str = "$(kernel_size[1])x$(kernel_size[2])" + println("$(lpad(image_str, 10)) | $(lpad(kernel_str, 6)) | $(lpad(round(med_fused, digits=2), 9)) | $(lpad(round(med_inline, digits=2), 10)) | $(lpad(round(speedup, digits=2), 7))x | $(lpad(round(med_cpu, digits=2), 6))") +end + +println("\n3D CONVOLUTION") +println("-" ^ 70) +println("Volume | Kernel | Fused(ms) | Inline(ms) | Speedup | CPU(ms)") +println("-" ^ 70) + +configs_3d = [ + ((32, 32, 32), (3, 3, 3)), + ((64, 64, 64), (3, 3, 3)), + ((64, 64, 64), (5, 5, 5)), + ((128, 128, 128), (5, 5, 5)), +] + +for (vol_size, kernel_size) in configs_3d + vol_cpu = rand(Float32, vol_size...) + kernel_cpu = rand(Float32, kernel_size...) + vol_gpu = MtlArray(vol_cpu) + kernel_gpu = MtlArray(kernel_cpu) + + # Warmup + _ = conv_fft_fused(vol_gpu, kernel_gpu) + _ = conv_fft_inline_pad(vol_gpu, kernel_gpu) + Metal.synchronize() + + n_iters = 8 + + times_fused = Float64[] + for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + _ = conv_fft_fused(vol_gpu, kernel_gpu) + Metal.synchronize() + end + Base.push!(times_fused, t * 1000) + end + + times_inline = Float64[] + for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + _ = conv_fft_inline_pad(vol_gpu, kernel_gpu) + Metal.synchronize() + end + Base.push!(times_inline, t * 1000) + end + + cpu_iters = prod(vol_size) > 500_000 ? 2 : 3 + times_cpu = Float64[] + for _ in 1:cpu_iters + t = @elapsed DSP.conv(vol_cpu, kernel_cpu) + Base.push!(times_cpu, t * 1000) + end + + med_fused = median(times_fused) + med_inline = median(times_inline) + med_cpu = median(times_cpu) + speedup = med_fused / med_inline + + vol_str = "$(vol_size[1])x$(vol_size[2])x$(vol_size[3])" + kernel_str = "$(kernel_size[1])x$(kernel_size[2])x$(kernel_size[3])" + println("$(lpad(vol_str, 10)) | $(lpad(kernel_str, 6)) | $(lpad(round(med_fused, digits=2), 9)) | $(lpad(round(med_inline, digits=2), 10)) | $(lpad(round(speedup, digits=2), 7))x | $(lpad(round(med_cpu, digits=2), 6))") +end diff --git a/perf/benchmark_throughput.jl b/perf/benchmark_throughput.jl new file mode 100644 index 000000000..c3d2c994d --- /dev/null +++ b/perf/benchmark_throughput.jl @@ -0,0 +1,129 @@ +using Metal +using Metal.MPSGraphs: conv_fft, conv_fft_fused +using DSP +using Statistics + +println("=" ^ 70) +println("THROUGHPUT BENCHMARK: Async Pipeline Performance") +println("=" ^ 70) + +# Test configuration: moderate size where GPU wins +signal_size = 1_000_000 +kernel_size = 500 + +println("\nTest config: Signal=$(signal_size), Kernel=$(kernel_size)") +println("-" ^ 70) + +# Create test data +signal_cpu = rand(Float32, signal_size) +kernel_cpu = rand(Float32, kernel_size) +signal_gpu = MtlVector(signal_cpu) +kernel_gpu = MtlVector(kernel_cpu) + +# Warmup +for _ in 1:3 + _ = conv_fft_fused(signal_gpu, kernel_gpu) +end +Metal.synchronize() + +# Test different batch sizes for throughput +batch_sizes = [1, 5, 10, 20, 50] + +println("\n--- FUSED GPU: Sync vs Async Pipeline ---") +println("Batch | Sync (ms/op) | Async (ms/op) | Speedup") +println("-" ^ 50) + +for batch in batch_sizes + # Sync mode: wait after each operation + Metal.synchronize() + t_sync = @elapsed begin + for _ in 1:batch + _ = conv_fft_fused(signal_gpu, kernel_gpu) + Metal.synchronize() + end + end + ms_sync = (t_sync / batch) * 1000 + + # Async mode: queue all, sync once + Metal.synchronize() + t_async = @elapsed begin + for _ in 1:batch + _ = conv_fft_fused(signal_gpu, kernel_gpu) + end + Metal.synchronize() + end + ms_async = (t_async / batch) * 1000 + + speedup = ms_sync / ms_async + println("$(lpad(batch, 5)) | $(lpad(round(ms_sync, digits=3), 12)) | $(lpad(round(ms_async, digits=3), 13)) | $(round(speedup, digits=2))x") +end + +println("\n--- EXISTING GPU: Sync vs Async Pipeline ---") +println("Batch | Sync (ms/op) | Async (ms/op) | Speedup") +println("-" ^ 50) + +for batch in batch_sizes + # Sync mode + Metal.synchronize() + t_sync = @elapsed begin + for _ in 1:batch + _ = conv_fft(signal_gpu, kernel_gpu) + Metal.synchronize() + end + end + ms_sync = (t_sync / batch) * 1000 + + # Async mode + Metal.synchronize() + t_async = @elapsed begin + for _ in 1:batch + _ = conv_fft(signal_gpu, kernel_gpu) + end + Metal.synchronize() + end + ms_async = (t_async / batch) * 1000 + + speedup = ms_sync / ms_async + println("$(lpad(batch, 5)) | $(lpad(round(ms_sync, digits=3), 12)) | $(lpad(round(ms_async, digits=3), 13)) | $(round(speedup, digits=2))x") +end + +# Compare throughput: fused async vs existing async vs CPU +println("\n" * "=" ^ 70) +println("MAXIMUM THROUGHPUT COMPARISON (batch=50, async)") +println("=" ^ 70) + +batch = 50 + +# Fused async +Metal.synchronize() +t_fused = @elapsed begin + for _ in 1:batch + _ = conv_fft_fused(signal_gpu, kernel_gpu) + end + Metal.synchronize() +end +tput_fused = batch / t_fused + +# Existing async +Metal.synchronize() +t_existing = @elapsed begin + for _ in 1:batch + _ = conv_fft(signal_gpu, kernel_gpu) + end + Metal.synchronize() +end +tput_existing = batch / t_existing + +# CPU +t_cpu = @elapsed begin + for _ in 1:batch + _ = DSP.conv(signal_cpu, kernel_cpu) + end +end +tput_cpu = batch / t_cpu + +println("\nFused GPU: $(round(tput_fused, digits=1)) ops/sec ($(round(1000/tput_fused, digits=2)) ms/op)") +println("Existing GPU: $(round(tput_existing, digits=1)) ops/sec ($(round(1000/tput_existing, digits=2)) ms/op)") +println("CPU (DSP): $(round(tput_cpu, digits=1)) ops/sec ($(round(1000/tput_cpu, digits=2)) ms/op)") +println("\nFused vs CPU: $(round(tput_fused / tput_cpu, digits=2))x throughput") +println("Fused vs Existing: $(round(tput_fused / tput_existing, digits=2))x throughput") diff --git a/perf/bottleneck_analysis.jl b/perf/bottleneck_analysis.jl new file mode 100644 index 000000000..90b61f360 --- /dev/null +++ b/perf/bottleneck_analysis.jl @@ -0,0 +1,495 @@ +using Metal +using Metal.MPSGraphs: conv_fft, conv_fft_fused, conv_fft_inline_pad +using Metal.MPSGraphs: _get_cached_fused_conv_graph, FusedConvGraphKey, nextfastfft, _conv_output_size +using Metal.MPSGraphs: MPSGraphTensorData, MPSCommandBuffer, NSDictionary, encode!, commit!, wait_completed, default_exec_desc, nil +using Metal.MPSGraphs: _get_cached_buffers +using Statistics +using Printf + +println("=" ^ 80) +println("COMPREHENSIVE BOTTLENECK ANALYSIS") +println("=" ^ 80) + +# ============================================================================ +# SECTION 1: Component-level breakdown for FUSED implementation +# ============================================================================ + +function analyze_fused_components(signal_size, kernel_size; n_iters=50) + println("\n" * "-" ^ 80) + println("FUSED IMPLEMENTATION: Signal=$signal_size, Kernel=$kernel_size") + println("-" ^ 80) + + # Setup + signal = MtlVector(rand(Float32, signal_size)) + kernel = MtlVector(rand(Float32, kernel_size)) + + # Warmup + _ = conv_fft_fused(signal, kernel) + Metal.synchronize() + + # Get sizes + ns, nk = length(signal), length(kernel) + full_size = ns + nk - 1 + output_size = _conv_output_size(ns, nk, :full) + nfft = nextfastfft(full_size) + + println("FFT size: $nfft, Output size: $output_size") + + # Get cached graph and buffers + key = FusedConvGraphKey((nfft,), (nfft,), (output_size,), Float32) + cached = _get_cached_fused_conv_graph(key) + buffers = _get_cached_buffers((nfft,), Float32) + + results = Dict{String, Float64}() + + # 1. Cache lookup (graph) + times = Float64[] + for _ in 1:n_iters + t = @elapsed _get_cached_fused_conv_graph(key) + push!(times, t * 1e6) + end + results["Graph cache lookup"] = median(times) + + # 2. Buffer pool lookup + times = Float64[] + for _ in 1:n_iters + t = @elapsed _get_cached_buffers((nfft,), Float32) + push!(times, t * 1e6) + end + results["Buffer pool lookup"] = median(times) + + # 3. copyto! for signal + times = Float64[] + for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + copyto!(buffers.signal_padded, 1, signal, 1, ns) + Metal.synchronize() + end + push!(times, t * 1e6) + end + results["copyto! signal (sync)"] = median(times) + + # 4. copyto! for kernel + times = Float64[] + for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + copyto!(buffers.kernel_padded, 1, kernel, 1, nk) + Metal.synchronize() + end + push!(times, t * 1e6) + end + results["copyto! kernel (sync)"] = median(times) + + # 5. Zero-padding signal + times = Float64[] + for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + @view(buffers.signal_padded[(ns+1):nfft]) .= 0f0 + Metal.synchronize() + end + push!(times, t * 1e6) + end + results["Zero-pad signal (sync)"] = median(times) + + # 6. Zero-padding kernel + times = Float64[] + for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + @view(buffers.kernel_padded[(nk+1):nfft]) .= 0f0 + Metal.synchronize() + end + push!(times, t * 1e6) + end + results["Zero-pad kernel (sync)"] = median(times) + + # 7. MPSGraphTensorData creation + times = Float64[] + for _ in 1:n_iters + t = @elapsed begin + td1 = MPSGraphTensorData(buffers.signal_padded) + td2 = MPSGraphTensorData(buffers.kernel_padded) + td3 = MPSGraphTensorData(buffers.output) + end + push!(times, t * 1e6) + end + results["MPSGraphTensorData (3x)"] = median(times) + + # 8. NSDictionary creation + td1 = MPSGraphTensorData(buffers.signal_padded) + td2 = MPSGraphTensorData(buffers.kernel_padded) + td3 = MPSGraphTensorData(buffers.output) + times = Float64[] + for _ in 1:n_iters + t = @elapsed begin + feeds = NSDictionary(Dict(cached.signal_placeholder => td1, cached.kernel_placeholder => td2)) + results_dict = NSDictionary(Dict(cached.result => td3)) + end + push!(times, t * 1e6) + end + results["NSDictionary creation"] = median(times) + + # 9. MPSCommandBuffer creation + times = Float64[] + for _ in 1:n_iters + t = @elapsed MPSCommandBuffer(Metal.global_queue(Metal.current_device())) + push!(times, t * 1e6) + end + results["MPSCommandBuffer"] = median(times) + + # 10. encode! + times = Float64[] + for _ in 1:n_iters + td1 = MPSGraphTensorData(buffers.signal_padded) + td2 = MPSGraphTensorData(buffers.kernel_padded) + td3 = MPSGraphTensorData(buffers.output) + feeds = NSDictionary(Dict(cached.signal_placeholder => td1, cached.kernel_placeholder => td2)) + results_ns = NSDictionary(Dict(cached.result => td3)) + cmdbuf = MPSCommandBuffer(Metal.global_queue(Metal.current_device())) + t = @elapsed encode!(cmdbuf, cached.graph, feeds, results_ns, nil, default_exec_desc()) + push!(times, t * 1e6) + end + results["encode!"] = median(times) + + # 11. commit! + times = Float64[] + for _ in 1:n_iters + td1 = MPSGraphTensorData(buffers.signal_padded) + td2 = MPSGraphTensorData(buffers.kernel_padded) + td3 = MPSGraphTensorData(buffers.output) + feeds = NSDictionary(Dict(cached.signal_placeholder => td1, cached.kernel_placeholder => td2)) + results_ns = NSDictionary(Dict(cached.result => td3)) + cmdbuf = MPSCommandBuffer(Metal.global_queue(Metal.current_device())) + encode!(cmdbuf, cached.graph, feeds, results_ns, nil, default_exec_desc()) + t = @elapsed commit!(cmdbuf) + push!(times, t * 1e6) + end + results["commit!"] = median(times) + + # 12. wait_completed + times = Float64[] + for _ in 1:n_iters + td1 = MPSGraphTensorData(buffers.signal_padded) + td2 = MPSGraphTensorData(buffers.kernel_padded) + td3 = MPSGraphTensorData(buffers.output) + feeds = NSDictionary(Dict(cached.signal_placeholder => td1, cached.kernel_placeholder => td2)) + results_ns = NSDictionary(Dict(cached.result => td3)) + cmdbuf = MPSCommandBuffer(Metal.global_queue(Metal.current_device())) + encode!(cmdbuf, cached.graph, feeds, results_ns, nil, default_exec_desc()) + commit!(cmdbuf) + t = @elapsed wait_completed(cmdbuf) + push!(times, t * 1e6) + end + results["wait_completed"] = median(times) + + # 13. Output slice copy + output_arr = MtlVector{Float32}(undef, output_size) + times = Float64[] + for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + copyto!(output_arr, 1, buffers.output, 1, output_size) + Metal.synchronize() + end + push!(times, t * 1e6) + end + results["Output slice copy (sync)"] = median(times) + + # Full conv_fft_fused + times = Float64[] + for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + _ = conv_fft_fused(signal, kernel) + Metal.synchronize() + end + push!(times, t * 1e6) + end + results["TOTAL conv_fft_fused"] = median(times) + + # Print results + println("\nComponent breakdown (μs):") + total_components = 0.0 + for (name, time) in sort(collect(results), by=x->x[2], rev=true) + if name != "TOTAL conv_fft_fused" + total_components += time + pct = 100 * time / results["TOTAL conv_fft_fused"] + @printf(" %-30s %8.1f μs (%5.1f%%)\n", name, time, pct) + end + end + println() + @printf(" %-30s %8.1f μs\n", "Sum of components:", total_components) + @printf(" %-30s %8.1f μs\n", "Actual total:", results["TOTAL conv_fft_fused"]) + + return results +end + +# ============================================================================ +# SECTION 2: Compare all three implementations +# ============================================================================ + +function compare_implementations(signal_size, kernel_size; n_iters=30) + println("\n" * "=" ^ 80) + println("IMPLEMENTATION COMPARISON: Signal=$signal_size, Kernel=$kernel_size") + println("=" ^ 80) + + signal_cpu = rand(Float32, signal_size) + kernel_cpu = rand(Float32, kernel_size) + signal = MtlVector(signal_cpu) + kernel = MtlVector(kernel_cpu) + + # Warmup all + _ = conv_fft(signal, kernel) + _ = conv_fft_fused(signal, kernel) + _ = conv_fft_inline_pad(signal, kernel) + Metal.synchronize() + + results = Dict{String, Float64}() + + # conv_fft (original, now uses fused internally) + times = Float64[] + for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + _ = conv_fft(signal, kernel) + Metal.synchronize() + end + push!(times, t * 1000) + end + results["conv_fft"] = median(times) + + # conv_fft_fused (explicit) + times = Float64[] + for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + _ = conv_fft_fused(signal, kernel) + Metal.synchronize() + end + push!(times, t * 1000) + end + results["conv_fft_fused"] = median(times) + + # conv_fft_inline_pad + times = Float64[] + for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + _ = conv_fft_inline_pad(signal, kernel) + Metal.synchronize() + end + push!(times, t * 1000) + end + results["conv_fft_inline_pad"] = median(times) + + println("\nLatency comparison (ms):") + for (name, time) in sort(collect(results), by=x->x[2]) + @printf(" %-25s %8.3f ms\n", name, time) + end + + fastest = minimum(values(results)) + println("\nSpeedups vs fastest:") + for (name, time) in sort(collect(results), by=x->x[2]) + @printf(" %-25s %5.2fx\n", name, time / fastest) + end + + return results +end + +# ============================================================================ +# SECTION 3: Async pipeline analysis +# ============================================================================ + +function analyze_async_pipeline(signal_size, kernel_size; batch_sizes=[1, 5, 10, 20, 50]) + println("\n" * "=" ^ 80) + println("ASYNC PIPELINE ANALYSIS: Signal=$signal_size, Kernel=$kernel_size") + println("=" ^ 80) + + signal = MtlVector(rand(Float32, signal_size)) + kernel = MtlVector(rand(Float32, kernel_size)) + + # Warmup + for _ in 1:3 + _ = conv_fft_fused(signal, kernel) + end + Metal.synchronize() + + println("\nBatch | Sync (ms/op) | Async (ms/op) | Pipeline Speedup") + println("-" ^ 60) + + for batch in batch_sizes + # Sync mode: wait after each + Metal.synchronize() + t_sync = @elapsed begin + for _ in 1:batch + _ = conv_fft_fused(signal, kernel) + Metal.synchronize() + end + end + ms_sync = (t_sync / batch) * 1000 + + # Async mode: queue all, wait once + Metal.synchronize() + t_async = @elapsed begin + for _ in 1:batch + _ = conv_fft_fused(signal, kernel) + end + Metal.synchronize() + end + ms_async = (t_async / batch) * 1000 + + speedup = ms_sync / ms_async + @printf("%5d | %12.3f | %13.3f | %5.2fx\n", batch, ms_sync, ms_async, speedup) + end +end + +# ============================================================================ +# SECTION 4: Memory operation breakdown +# ============================================================================ + +function analyze_memory_operations(signal_size, kernel_size; n_iters=50) + println("\n" * "=" ^ 80) + println("MEMORY OPERATION DEEP DIVE: Signal=$signal_size, Kernel=$kernel_size") + println("=" ^ 80) + + ns, nk = signal_size, kernel_size + full_size = ns + nk - 1 + nfft = nextfastfft(full_size) + + println("Sizes: signal=$ns, kernel=$nk, FFT=$nfft, padding_signal=$(nfft-ns), padding_kernel=$(nfft-nk)") + + signal = MtlVector(rand(Float32, signal_size)) + kernel = MtlVector(rand(Float32, kernel_size)) + + # Pre-allocate + signal_padded = MtlVector{Float32}(undef, nfft) + kernel_padded = MtlVector{Float32}(undef, nfft) + + # Warmup + copyto!(signal_padded, 1, signal, 1, ns) + Metal.synchronize() + + results = Dict{String, Float64}() + + # Test 1: copyto! with sync + times = Float64[] + for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + copyto!(signal_padded, 1, signal, 1, ns) + Metal.synchronize() + end + push!(times, t * 1e6) + end + results["copyto! signal (with sync)"] = median(times) + + # Test 2: copyto! without sync (just queue time) + times = Float64[] + for _ in 1:n_iters + Metal.synchronize() + t = @elapsed copyto!(signal_padded, 1, signal, 1, ns) + push!(times, t * 1e6) + end + results["copyto! signal (queue only)"] = median(times) + + # Test 3: broadcast zero-fill with sync + times = Float64[] + for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + @view(signal_padded[(ns+1):nfft]) .= 0f0 + Metal.synchronize() + end + push!(times, t * 1e6) + end + results["Zero-fill broadcast (with sync)"] = median(times) + + # Test 4: fill! for zeros + times = Float64[] + for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + fill!(@view(signal_padded[(ns+1):nfft]), 0f0) + Metal.synchronize() + end + push!(times, t * 1e6) + end + results["fill! zeros (with sync)"] = median(times) + + # Test 5: Full array fill + times = Float64[] + for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + fill!(signal_padded, 0f0) + Metal.synchronize() + end + push!(times, t * 1e6) + end + results["fill! full array (with sync)"] = median(times) + + # Test 6: MtlArray allocation + times = Float64[] + for _ in 1:n_iters + GC.gc(false) + t = @elapsed MtlVector{Float32}(undef, nfft) + push!(times, t * 1e6) + end + results["MtlArray allocation"] = median(times) + + # Test 7: Combined copy + zero in one sync + times = Float64[] + for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + copyto!(signal_padded, 1, signal, 1, ns) + @view(signal_padded[(ns+1):nfft]) .= 0f0 + Metal.synchronize() + end + push!(times, t * 1e6) + end + results["copyto! + zero (one sync)"] = median(times) + + # Test 8: Metal.synchronize() alone + times = Float64[] + for _ in 1:n_iters + Metal.synchronize() + t = @elapsed Metal.synchronize() + push!(times, t * 1e6) + end + results["Metal.synchronize() alone"] = median(times) + + println("\nMemory operation timings (μs):") + for (name, time) in sort(collect(results), by=x->x[2], rev=true) + @printf(" %-35s %8.1f μs\n", name, time) + end + + return results +end + +# ============================================================================ +# RUN ANALYSIS +# ============================================================================ + +# Small signal (where GPU struggles) +analyze_fused_components(10_000, 100) +compare_implementations(10_000, 100) +analyze_memory_operations(10_000, 100) + +# Medium signal +analyze_fused_components(100_000, 500) +compare_implementations(100_000, 500) + +# Large signal (where GPU wins) +analyze_fused_components(1_000_000, 1000) +compare_implementations(1_000_000, 1000) + +# Async pipeline analysis +analyze_async_pipeline(100_000, 500) + +println("\n" * "=" ^ 80) +println("ANALYSIS COMPLETE") +println("=" ^ 80) diff --git a/perf/padding_alternatives.jl b/perf/padding_alternatives.jl new file mode 100644 index 000000000..9fdb81e09 --- /dev/null +++ b/perf/padding_alternatives.jl @@ -0,0 +1,265 @@ +using Metal +using Statistics +using Printf + +println("=" ^ 80) +println("PADDING ALTERNATIVES: Can we eliminate/reduce memory overhead?") +println("=" ^ 80) + +signal_size = 100_000 +kernel_size = 500 +full_size = signal_size + kernel_size - 1 +nfft = Metal.MPSGraphs.nextfastfft(full_size) + +println("\nSetup: signal=$signal_size, kernel=$kernel_size, FFT size=$nfft") +println("Padding needed: signal=$(nfft - signal_size), kernel=$(nfft - kernel_size)") + +signal = MtlVector(rand(Float32, signal_size)) +kernel = MtlVector(rand(Float32, kernel_size)) + +n_iters = 30 + +# ============================================================================ +# CURRENT APPROACH: Separate copy + zero-fill +# ============================================================================ + +println("\n" * "=" ^ 60) +println("METHOD 1: Current approach (copyto! + broadcast zero)") +println("=" ^ 60) + +signal_padded = MtlVector{Float32}(undef, nfft) +kernel_padded = MtlVector{Float32}(undef, nfft) + +# Warmup +copyto!(signal_padded, 1, signal, 1, signal_size) +@view(signal_padded[(signal_size+1):nfft]) .= 0f0 +Metal.synchronize() + +times = Float64[] +for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + copyto!(signal_padded, 1, signal, 1, signal_size) + @view(signal_padded[(signal_size+1):nfft]) .= 0f0 + copyto!(kernel_padded, 1, kernel, 1, kernel_size) + @view(kernel_padded[(kernel_size+1):nfft]) .= 0f0 + Metal.synchronize() + end + push!(times, t * 1e6) +end +current_time = median(times) +@printf(" Time: %.1f μs\n", current_time) + +# ============================================================================ +# ALTERNATIVE 1: Pre-fill with zeros, then copy +# ============================================================================ + +println("\n" * "=" ^ 60) +println("METHOD 2: Pre-fill zeros, then copy (fill! + copyto!)") +println("=" ^ 60) + +times = Float64[] +for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + fill!(signal_padded, 0f0) + fill!(kernel_padded, 0f0) + copyto!(signal_padded, 1, signal, 1, signal_size) + copyto!(kernel_padded, 1, kernel, 1, kernel_size) + Metal.synchronize() + end + push!(times, t * 1e6) +end +prefill_time = median(times) +@printf(" Time: %.1f μs (%.2fx vs current)\n", prefill_time, prefill_time / current_time) + +# ============================================================================ +# ALTERNATIVE 2: Use zeros() then copy +# ============================================================================ + +println("\n" * "=" ^ 60) +println("METHOD 3: Create with zeros() then copy") +println("=" ^ 60) + +times = Float64[] +for _ in 1:n_iters + Metal.synchronize() + GC.gc(false) + t = @elapsed begin + sp = Metal.zeros(Float32, nfft) + kp = Metal.zeros(Float32, nfft) + copyto!(sp, 1, signal, 1, signal_size) + copyto!(kp, 1, kernel, 1, kernel_size) + Metal.synchronize() + end + push!(times, t * 1e6) +end +zeros_time = median(times) +@printf(" Time: %.1f μs (%.2fx vs current)\n", zeros_time, zeros_time / current_time) + +# ============================================================================ +# ALTERNATIVE 3: Single kernel pad-copy (custom kernel) +# ============================================================================ + +println("\n" * "=" ^ 60) +println("METHOD 4: Custom Metal kernel for pad+copy") +println("=" ^ 60) + +# Define a custom kernel that copies and pads in one operation +function pad_copy_kernel(dest, src, src_len) + i = thread_position_in_grid_1d() + if i <= src_len + @inbounds dest[i] = src[i] + elseif i <= length(dest) + @inbounds dest[i] = 0f0 + end + return +end + +# Warmup +@metal threads=256 groups=cld(nfft, 256) pad_copy_kernel(signal_padded, signal, signal_size) +Metal.synchronize() + +times = Float64[] +for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + @metal threads=256 groups=cld(nfft, 256) pad_copy_kernel(signal_padded, signal, signal_size) + @metal threads=256 groups=cld(nfft, 256) pad_copy_kernel(kernel_padded, kernel, kernel_size) + Metal.synchronize() + end + push!(times, t * 1e6) +end +custom_kernel_time = median(times) +@printf(" Time: %.1f μs (%.2fx vs current)\n", custom_kernel_time, custom_kernel_time / current_time) + +# ============================================================================ +# ALTERNATIVE 4: Batch operations without intermediate sync +# ============================================================================ + +println("\n" * "=" ^ 60) +println("METHOD 5: Queue all ops, single sync at end") +println("=" ^ 60) + +times = Float64[] +for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + # Queue all operations without waiting + copyto!(signal_padded, 1, signal, 1, signal_size) + copyto!(kernel_padded, 1, kernel, 1, kernel_size) + @view(signal_padded[(signal_size+1):nfft]) .= 0f0 + @view(kernel_padded[(kernel_size+1):nfft]) .= 0f0 + # Single sync at end + Metal.synchronize() + end + push!(times, t * 1e6) +end +batch_time = median(times) +@printf(" Time: %.1f μs (%.2fx vs current)\n", batch_time, batch_time / current_time) + +# ============================================================================ +# ALTERNATIVE 5: No padding at all - what if data comes pre-padded? +# ============================================================================ + +println("\n" * "=" ^ 60) +println("METHOD 6: Pre-padded data (best case scenario)") +println("=" ^ 60) + +# Simulate pre-padded data +signal_prepadded = Metal.zeros(Float32, nfft) +copyto!(signal_prepadded, 1, signal, 1, signal_size) +kernel_prepadded = Metal.zeros(Float32, nfft) +copyto!(kernel_prepadded, 1, kernel, 1, kernel_size) +Metal.synchronize() + +times = Float64[] +for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + # Just copy pre-padded data (simulating zero-cost padding) + copyto!(signal_padded, signal_prepadded) + copyto!(kernel_padded, kernel_prepadded) + Metal.synchronize() + end + push!(times, t * 1e6) +end +prepadded_time = median(times) +@printf(" Time: %.1f μs (%.2fx vs current)\n", prepadded_time, prepadded_time / current_time) + +# ============================================================================ +# ALTERNATIVE 6: Skip padding entirely - use inline padding in MPSGraph +# ============================================================================ + +println("\n" * "=" ^ 60) +println("METHOD 7: MPSGraph inline padding (no Julia-side padding)") +println("=" ^ 60) + +# With inline padding, we pass the original unpadded arrays +# The MPSGraph handles padding internally via concat operations + +using Metal.MPSGraphs: conv_fft_fused, conv_fft_inline_pad + +# conv_fft_fused (needs pre-padded buffers, so we measure padding + execution) +times_fused = Float64[] +for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + _ = conv_fft_fused(signal, kernel) + Metal.synchronize() + end + push!(times_fused, t * 1000) # ms +end +fused_time = median(times_fused) + +# conv_fft_inline_pad (no Julia-side padding needed) +times_inline = Float64[] +for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + _ = conv_fft_inline_pad(signal, kernel) + Metal.synchronize() + end + push!(times_inline, t * 1000) # ms +end +inline_time = median(times_inline) + +@printf(" conv_fft_fused: %.3f ms (includes padding)\n", fused_time) +@printf(" conv_fft_inline_pad: %.3f ms (no Julia-side padding)\n", inline_time) +@printf(" Difference: %.3f ms saved\n", fused_time - inline_time) + +# ============================================================================ +# SUMMARY +# ============================================================================ + +println("\n" * "=" ^ 80) +println("SUMMARY: Padding time comparison") +println("=" ^ 80) + +results = [ + ("Current (copyto! + broadcast)", current_time), + ("Pre-fill zeros + copy", prefill_time), + ("Create zeros() + copy", zeros_time), + ("Custom Metal kernel", custom_kernel_time), + ("Batch ops (single sync)", batch_time), + ("Pre-padded data", prepadded_time), +] + +sort!(results, by=x->x[2]) + +println("\nRanked by speed:") +for (i, (name, time)) in enumerate(results) + speedup = current_time / time + @printf("%d. %-30s %7.1f μs (%5.2fx vs current)\n", i, name, time, speedup) +end + +println("\n" * "=" ^ 80) +println("KEY FINDINGS") +println("=" ^ 80) +println(""" +1. The padding overhead (~700 μs) comes from GPU kernel launch costs +2. Each GPU operation (copyto!, broadcast) has ~200 μs overhead +3. Custom kernels can potentially reduce this by combining operations +4. The inline-pad implementation in MPSGraph eliminates Julia-side padding +5. Pre-padded data would be fastest but requires user workflow changes +""") diff --git a/perf/simple_bottleneck.jl b/perf/simple_bottleneck.jl new file mode 100644 index 000000000..791a47533 --- /dev/null +++ b/perf/simple_bottleneck.jl @@ -0,0 +1,362 @@ +using Metal +using Metal.MPSGraphs: conv_fft, conv_fft_fused, conv_fft_inline_pad +using DSP +using Statistics +using Printf + +println("=" ^ 80) +println("BOTTLENECK PROFILING: Current Implementation Analysis") +println("=" ^ 80) + +# ============================================================================ +# Test 1: Implementation comparison across sizes +# ============================================================================ + +function benchmark_latency(f, args...; n_warmup=3, n_iters=20) + # Warmup + for _ in 1:n_warmup + f(args...) + end + Metal.synchronize() + + # Benchmark + times = Float64[] + for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + f(args...) + Metal.synchronize() + end + push!(times, t * 1000) # ms + end + return median(times), minimum(times), maximum(times) +end + +function benchmark_cpu(signal_cpu, kernel_cpu; n_iters=5) + times = Float64[] + for _ in 1:n_iters + t = @elapsed DSP.conv(signal_cpu, kernel_cpu) + push!(times, t * 1000) + end + return median(times) +end + +println("\n" * "=" ^ 80) +println("SECTION 1: LATENCY COMPARISON ACROSS SIZES") +println("=" ^ 80) + +configs = [ + (10_000, 100, "Small"), + (50_000, 250, "Medium-Small"), + (100_000, 500, "Medium"), + (500_000, 500, "Medium-Large"), + (1_000_000, 1000, "Large"), +] + +println("\n| Size | Signal | Kernel | conv_fft | conv_fused | conv_inline | CPU (ms) | Best GPU |") +println("|------------|----------|--------|----------|------------|-------------|----------|----------|") + +for (signal_size, kernel_size, label) in configs + signal_cpu = rand(Float32, signal_size) + kernel_cpu = rand(Float32, kernel_size) + signal = MtlVector(signal_cpu) + kernel = MtlVector(kernel_cpu) + + # Benchmark each + t_fft, _, _ = benchmark_latency(conv_fft, signal, kernel) + t_fused, _, _ = benchmark_latency(conv_fft_fused, signal, kernel) + t_inline, _, _ = benchmark_latency(conv_fft_inline_pad, signal, kernel) + t_cpu = benchmark_cpu(signal_cpu, kernel_cpu) + + # Find best GPU + times = [("conv_fft", t_fft), ("fused", t_fused), ("inline", t_inline)] + best_name, best_time = sort(times, by=x->x[2])[1] + + @printf("| %-10s | %8d | %6d | %8.3f | %10.3f | %11.3f | %8.2f | %-8s |\n", + label, signal_size, kernel_size, t_fft, t_fused, t_inline, t_cpu, best_name) +end + +# ============================================================================ +# Test 2: Deep dive into memory operations +# ============================================================================ + +println("\n" * "=" ^ 80) +println("SECTION 2: MEMORY OPERATION OVERHEAD ANALYSIS") +println("=" ^ 80) + +function analyze_memory_overhead(signal_size, kernel_size; n_iters=30) + println("\n--- Signal=$signal_size, Kernel=$kernel_size ---") + + signal = MtlVector(rand(Float32, signal_size)) + kernel = MtlVector(rand(Float32, kernel_size)) + + # Compute FFT size + full_size = signal_size + kernel_size - 1 + nfft = Metal.MPSGraphs.nextfastfft(full_size) + output_size = full_size + + println("FFT size: $nfft, Padding needed: $(nfft - signal_size) (signal), $(nfft - kernel_size) (kernel)") + + # Pre-allocate + signal_padded = MtlVector{Float32}(undef, nfft) + kernel_padded = MtlVector{Float32}(undef, nfft) + + # Test 1: copyto! queued (no sync) + times_copyto_queue = Float64[] + for _ in 1:n_iters + Metal.synchronize() + t = @elapsed copyto!(signal_padded, 1, signal, 1, signal_size) + push!(times_copyto_queue, t * 1e6) + end + med_copyto_queue = median(times_copyto_queue) + + # Test 2: copyto! with sync + times_copyto_sync = Float64[] + for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + copyto!(signal_padded, 1, signal, 1, signal_size) + Metal.synchronize() + end + push!(times_copyto_sync, t * 1e6) + end + med_copyto_sync = median(times_copyto_sync) + + # Test 3: broadcast zero-fill queued + times_zero_queue = Float64[] + for _ in 1:n_iters + Metal.synchronize() + t = @elapsed (@view(signal_padded[(signal_size+1):nfft]) .= 0f0) + push!(times_zero_queue, t * 1e6) + end + med_zero_queue = median(times_zero_queue) + + # Test 4: broadcast zero-fill with sync + times_zero_sync = Float64[] + for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + @view(signal_padded[(signal_size+1):nfft]) .= 0f0 + Metal.synchronize() + end + push!(times_zero_sync, t * 1e6) + end + med_zero_sync = median(times_zero_sync) + + # Test 5: Combined copy + zero with one sync + times_combined = Float64[] + for _ in 1:n_iters + Metal.synchronize() + t = @elapsed begin + copyto!(signal_padded, 1, signal, 1, signal_size) + @view(signal_padded[(signal_size+1):nfft]) .= 0f0 + copyto!(kernel_padded, 1, kernel, 1, kernel_size) + @view(kernel_padded[(kernel_size+1):nfft]) .= 0f0 + Metal.synchronize() + end + push!(times_combined, t * 1e6) + end + med_combined = median(times_combined) + + # Test 6: MtlArray allocation + times_alloc = Float64[] + for _ in 1:n_iters + GC.gc(false) + t = @elapsed MtlVector{Float32}(undef, nfft) + push!(times_alloc, t * 1e6) + end + med_alloc = median(times_alloc) + + # Test 7: Metal.synchronize() alone (baseline) + times_sync_alone = Float64[] + for _ in 1:n_iters + Metal.synchronize() + t = @elapsed Metal.synchronize() + push!(times_sync_alone, t * 1e6) + end + med_sync_alone = median(times_sync_alone) + + @printf(" copyto! (queue only): %7.1f μs\n", med_copyto_queue) + @printf(" copyto! (with sync): %7.1f μs\n", med_copyto_sync) + @printf(" Zero-fill (queue only): %7.1f μs\n", med_zero_queue) + @printf(" Zero-fill (with sync): %7.1f μs\n", med_zero_sync) + @printf(" Both inputs padded (sync): %7.1f μs\n", med_combined) + @printf(" MtlArray allocation: %7.1f μs\n", med_alloc) + @printf(" Metal.synchronize() alone: %7.1f μs\n", med_sync_alone) + + # Calculate actual GPU work time + gpu_work_time = med_copyto_sync - med_sync_alone + @printf(" => Actual GPU copy time: %7.1f μs\n", gpu_work_time) + + return med_combined +end + +analyze_memory_overhead(10_000, 100) +analyze_memory_overhead(100_000, 500) +analyze_memory_overhead(1_000_000, 1000) + +# ============================================================================ +# Test 3: Async pipeline benefits +# ============================================================================ + +println("\n" * "=" ^ 80) +println("SECTION 3: ASYNC PIPELINE BENEFITS") +println("=" ^ 80) + +function test_async_pipeline(signal_size, kernel_size; batch_sizes=[1, 5, 10, 20, 50]) + println("\n--- Signal=$signal_size, Kernel=$kernel_size ---") + + signal = MtlVector(rand(Float32, signal_size)) + kernel = MtlVector(rand(Float32, kernel_size)) + + # Warmup + for _ in 1:3 + _ = conv_fft_fused(signal, kernel) + end + Metal.synchronize() + + println("Batch | Sync (ms/op) | Async (ms/op) | Speedup | Throughput") + println("-" ^ 60) + + for batch in batch_sizes + # Sync mode + Metal.synchronize() + t_sync = @elapsed begin + for _ in 1:batch + _ = conv_fft_fused(signal, kernel) + Metal.synchronize() + end + end + ms_sync = (t_sync / batch) * 1000 + + # Async mode + Metal.synchronize() + t_async = @elapsed begin + for _ in 1:batch + _ = conv_fft_fused(signal, kernel) + end + Metal.synchronize() + end + ms_async = (t_async / batch) * 1000 + + speedup = ms_sync / ms_async + throughput = batch / t_async + + @printf("%5d | %12.3f | %13.3f | %6.2fx | %7.1f ops/s\n", + batch, ms_sync, ms_async, speedup, throughput) + end +end + +test_async_pipeline(100_000, 500) +test_async_pipeline(1_000_000, 1000) + +# ============================================================================ +# Test 4: GPU vs CPU crossover point +# ============================================================================ + +println("\n" * "=" ^ 80) +println("SECTION 4: GPU vs CPU CROSSOVER ANALYSIS") +println("=" ^ 80) + +sizes = [1_000, 2_500, 5_000, 7_500, 10_000, 15_000, 25_000, 50_000, 100_000] +kernel_size = 100 + +println("\nKernel size fixed at $kernel_size") +println("Signal Size | GPU (ms) | CPU (ms) | GPU/CPU | Winner") +println("-" ^ 60) + +for signal_size in sizes + signal_cpu = rand(Float32, signal_size) + kernel_cpu = rand(Float32, kernel_size) + signal = MtlVector(signal_cpu) + kernel = MtlVector(kernel_cpu) + + t_gpu, _, _ = benchmark_latency(conv_fft_fused, signal, kernel; n_iters=15) + t_cpu = benchmark_cpu(signal_cpu, kernel_cpu; n_iters=5) + + ratio = t_gpu / t_cpu + winner = t_gpu < t_cpu ? "GPU" : "CPU" + + @printf("%11d | %8.3f | %8.3f | %7.2fx | %s\n", + signal_size, t_gpu, t_cpu, ratio, winner) +end + +# ============================================================================ +# Test 5: 2D and 3D analysis +# ============================================================================ + +println("\n" * "=" ^ 80) +println("SECTION 5: 2D CONVOLUTION ANALYSIS") +println("=" ^ 80) + +configs_2d = [ + ((64, 64), (5, 5)), + ((128, 128), (5, 5)), + ((256, 256), (5, 5)), + ((256, 256), (15, 15)), + ((512, 512), (15, 15)), + ((1024, 1024), (15, 15)), +] + +println("\n| Image Size | Kernel | conv_fft | conv_fused | conv_inline | CPU (ms) | Speedup |") +println("|------------|--------|----------|------------|-------------|----------|---------|") + +for (img_size, kern_size) in configs_2d + img_cpu = rand(Float32, img_size...) + kern_cpu = rand(Float32, kern_size...) + img = MtlMatrix(img_cpu) + kern = MtlMatrix(kern_cpu) + + t_fft, _, _ = benchmark_latency(x -> conv_fft(x, kern; dims=(1,2)), img; n_iters=15) + t_fused, _, _ = benchmark_latency(conv_fft_fused, img, kern; n_iters=15) + t_inline, _, _ = benchmark_latency(conv_fft_inline_pad, img, kern; n_iters=15) + + cpu_iters = prod(img_size) > 500_000 ? 3 : 5 + t_cpu = benchmark_cpu(img_cpu, kern_cpu; n_iters=cpu_iters) + + best_gpu = min(t_fft, t_fused, t_inline) + speedup = t_cpu / best_gpu + + @printf("| %4dx%-5d | %2dx%-3d | %8.3f | %10.3f | %11.3f | %8.2f | %6.2fx |\n", + img_size[1], img_size[2], kern_size[1], kern_size[2], + t_fft, t_fused, t_inline, t_cpu, speedup) +end + +println("\n" * "=" ^ 80) +println("SECTION 6: 3D CONVOLUTION ANALYSIS") +println("=" ^ 80) + +configs_3d = [ + ((32, 32, 32), (3, 3, 3)), + ((64, 64, 64), (3, 3, 3)), + ((64, 64, 64), (5, 5, 5)), + ((128, 128, 64), (5, 5, 5)), +] + +println("\n| Volume Size | Kernel | conv_fft | conv_fused | conv_inline | CPU (ms) | Speedup |") +println("|----------------|--------|----------|------------|-------------|----------|---------|") + +for (vol_size, kern_size) in configs_3d + vol_cpu = rand(Float32, vol_size...) + kern_cpu = rand(Float32, kern_size...) + vol = MtlArray(vol_cpu) + kern = MtlArray(kern_cpu) + + t_fft, _, _ = benchmark_latency(x -> conv_fft(x, kern; dims=(1,2,3)), vol; n_iters=10) + t_fused, _, _ = benchmark_latency(conv_fft_fused, vol, kern; n_iters=10) + t_inline, _, _ = benchmark_latency(conv_fft_inline_pad, vol, kern; n_iters=10) + + cpu_iters = prod(vol_size) > 500_000 ? 2 : 3 + t_cpu = benchmark_cpu(vol_cpu, kern_cpu; n_iters=cpu_iters) + + best_gpu = min(t_fft, t_fused, t_inline) + speedup = t_cpu / best_gpu + + @printf("| %3dx%3dx%-5d | %dx%dx%-1d | %8.3f | %10.3f | %11.3f | %8.2f | %6.2fx |\n", + vol_size[1], vol_size[2], vol_size[3], kern_size[1], kern_size[2], kern_size[3], + t_fft, t_fused, t_inline, t_cpu, speedup) +end + +println("\n" * "=" ^ 80) +println("ANALYSIS COMPLETE") +println("=" ^ 80) diff --git a/src/Metal.jl b/src/Metal.jl index 2d31a2686..01b6cd769 100644 --- a/src/Metal.jl +++ b/src/Metal.jl @@ -60,6 +60,10 @@ include("../lib/mps/MPS.jl") export MPS include("../lib/mpsgraphs/MPSGraphs.jl") export MPSGraphs +# Re-export the public convolution API to the top-level Metal namespace. +# Internal helpers (conv_direct, imfilter, plan-cache management) stay in MPSGraphs. +using .MPSGraphs: conv, conv_fft, conv_fft!, conv_fft_fused, xcorr, plan_conv_fft, ConvFFTPlan +export conv, conv_fft, conv_fft!, conv_fft_fused, xcorr, plan_conv_fft, ConvFFTPlan # LinearAlgebra include("linalg.jl") diff --git a/test/mpsgraphs/convolution.jl b/test/mpsgraphs/convolution.jl new file mode 100644 index 000000000..cd78b3148 --- /dev/null +++ b/test/mpsgraphs/convolution.jl @@ -0,0 +1,401 @@ +# Tests for convolution operations in MPSGraphs +# +# This tests: +# - FFT-based convolution (conv_fft) +# - MPS direct convolution (conv_direct) +# - Unified conv() API with auto-selection +# - Convolution plan caching + +# Internal (non-exported) convolution symbols, accessed directly from the submodule. +# Metal exports only the core conv API; these stay internal. +using Metal.MPSGraphs: conv_direct, get_cached_conv_plan, clear_conv_plan_cache!, imfilter + +# Simple reference convolution for verification (CPU) +function ref_conv(u::Vector{T}, v::Vector{T}) where {T} + nu = length(u) + nv = length(v) + n = nu + nv - 1 + result = zeros(T, n) + for i in 1:nu + for j in 1:nv + result[i + j - 1] += u[i] * v[j] + end + end + return result +end + +# Tolerance functions based on type precision +rtol(::Type{Float16}) = 1.0e-2 +rtol(::Type{Float32}) = 1.0e-4 +rtol(::Type{ComplexF16}) = 1.0e-2 +rtol(::Type{ComplexF32}) = 1.0e-4 + +if MPS.is_supported(device()) + + # ============================================================================ + # FFT Convolution Tests + # ============================================================================ + + @testset "FFT Convolution" begin + @testset "1D Real Convolution" begin + @testset for T in [Float32, Float16] + # Basic convolution + signal = rand(T, 100) + kernel = rand(T, 10) + + d_signal = MtlVector(signal) + d_kernel = MtlVector(kernel) + + # Full mode + result = Array(conv_fft(d_signal, d_kernel; mode = :full)) + expected = T.(ref_conv(Float64.(signal), Float64.(kernel))) + @test isapprox(result, expected, rtol = rtol(T)) + + # Same mode + result_same = Array(conv_fft(d_signal, d_kernel; mode = :same)) + @test length(result_same) == length(signal) + + # Valid mode + result_valid = Array(conv_fft(d_signal, d_kernel; mode = :valid)) + @test length(result_valid) == length(signal) - length(kernel) + 1 + end + end + + @testset "1D Complex Convolution" begin + @testset for T in [ComplexF32, ComplexF16] + signal = rand(T, 100) + kernel = rand(T, 10) + + d_signal = MtlVector(signal) + d_kernel = MtlVector(kernel) + + result = Array(conv_fft(d_signal, d_kernel; mode = :full)) + # For complex, just verify output size (reference conv doesn't support complex) + @test length(result) == length(signal) + length(kernel) - 1 + end + end + + @testset "2D Convolution" begin + @testset for T in [Float32, Float16] + signal = rand(T, 64, 64) + kernel = rand(T, 5, 5) + + d_signal = MtlMatrix(signal) + d_kernel = MtlMatrix(kernel) + + # Full mode along both dimensions + result = Array(conv_fft(d_signal, d_kernel; dims = (1, 2), mode = :full)) + @test size(result) == (68, 68) + + # Same mode + result_same = Array(conv_fft(d_signal, d_kernel; dims = (1, 2), mode = :same)) + @test size(result_same) == size(signal) + + # Valid mode + result_valid = Array(conv_fft(d_signal, d_kernel; dims = (1, 2), mode = :valid)) + @test size(result_valid) == (60, 60) + end + end + + @testset "Batched Convolution" begin + # 3D array, convolve along dims 1 and 2 + signal = rand(Float32, 32, 32, 4) + kernel = rand(Float32, 3, 3, 4) + + d_signal = MtlArray(signal) + d_kernel = MtlArray(kernel) + + result = conv_fft(d_signal, d_kernel; dims = (1, 2), mode = :same) + @test size(result) == size(signal) + end + end + + # ============================================================================ + # Cross-correlation Tests + # ============================================================================ + + @testset "Cross-correlation" begin + @testset for T in [Float32, Float16] + u = rand(T, 100) + v = rand(T, 10) + + d_u = MtlVector(u) + d_v = MtlVector(v) + + result = Array(xcorr(d_u, d_v; mode = :full)) + # Cross-correlation is convolution with reversed kernel + expected = Array(conv_fft(d_u, MtlVector(reverse(v)); mode = :full)) + @test isapprox(result, expected, rtol = rtol(T)) + end + end + + # ============================================================================ + # Convolution Plan Tests + # ============================================================================ + + @testset "Convolution Plans" begin + @testset "Basic Plan Usage" begin + signal_size = 1000 + kernel_size = 100 + + plan = plan_conv_fft(signal_size, kernel_size, Float32; mode = :full) + + signal = MtlVector(rand(Float32, signal_size)) + kernel = MtlVector(rand(Float32, kernel_size)) + + result = conv_fft(plan, signal, kernel) + expected = conv_fft(signal, kernel; mode = :full) + + @test isapprox(Array(result), Array(expected), rtol = 1.0e-4) + end + + @testset "Plan with Pre-computed Kernel" begin + signal_size = 1000 + kernel = MtlVector(rand(Float32, 100)) + + plan = plan_conv_fft(signal_size, kernel; mode = :full) + + signal1 = MtlVector(rand(Float32, signal_size)) + signal2 = MtlVector(rand(Float32, signal_size)) + + result1 = conv_fft(plan, signal1) + result2 = conv_fft(plan, signal2) + + # Verify correctness + expected1 = conv_fft(signal1, kernel; mode = :full) + expected2 = conv_fft(signal2, kernel; mode = :full) + + @test isapprox(Array(result1), Array(expected1), rtol = 1.0e-4) + @test isapprox(Array(result2), Array(expected2), rtol = 1.0e-4) + end + + @testset "Cached Plan" begin + signal_size = (64, 64) + kernel_size = (5, 5) + + # Get same plan twice + plan1 = get_cached_conv_plan(signal_size, kernel_size, Float32; dims = (1, 2)) + plan2 = get_cached_conv_plan(signal_size, kernel_size, Float32; dims = (1, 2)) + + # Should be same object (cached) + @test plan1 === plan2 + + # Clean up + clear_conv_plan_cache!() + end + end + + # ============================================================================ + # MPS Direct Convolution Tests + # ============================================================================ + + @testset "MPS Direct Convolution" begin + @testset "Basic 2D Convolution" begin + @testset for T in [Float32, Float16] + image = rand(T, 64, 64) + kernel = rand(T, 3, 3) + + d_image = MtlMatrix(image) + d_kernel = MtlMatrix(kernel) + + # Direct convolution + result_direct = Array(conv_direct(d_image, d_kernel; mode = :same)) + + # FFT convolution (for comparison) + result_fft = Array(conv_fft(d_image, d_kernel; dims = (1, 2), mode = :same)) + + @test size(result_direct) == size(image) + @test isapprox(result_direct, result_fft, rtol = rtol(T)) + end + end + + @testset "Different Kernel Sizes" begin + image = rand(Float32, 128, 128) + d_image = MtlMatrix(image) + + for ks in [3, 5, 7, 9, 11] + kernel = rand(Float32, ks, ks) + d_kernel = MtlMatrix(kernel) + + result_direct = Array(conv_direct(d_image, d_kernel; mode = :same)) + result_fft = Array(conv_fft(d_image, d_kernel; dims = (1, 2), mode = :same)) + + @test isapprox(result_direct, result_fft, rtol = 1.0e-4) + end + end + + @testset "Valid Mode" begin + image = rand(Float32, 64, 64) + kernel = rand(Float32, 5, 5) + + d_image = MtlMatrix(image) + d_kernel = MtlMatrix(kernel) + + result = Array(conv_direct(d_image, d_kernel; mode = :valid)) + @test size(result) == (60, 60) + + # Compare with FFT + result_fft = Array(conv_fft(d_image, d_kernel; dims = (1, 2), mode = :valid)) + @test isapprox(result, result_fft, rtol = 1.0e-4) + end + + @testset "Full Mode Falls Back to FFT" begin + image = rand(Float32, 64, 64) + kernel = rand(Float32, 3, 3) + + d_image = MtlMatrix(image) + d_kernel = MtlMatrix(kernel) + + # Full mode should work (falls back to FFT internally) + result = Array(conv_direct(d_image, d_kernel; mode = :full)) + @test size(result) == (66, 66) + end + end + + # ============================================================================ + # imfilter Tests + # ============================================================================ + + @testset "imfilter" begin + @testset "Small Kernel (uses direct)" begin + image = rand(Float32, 256, 256) + kernel = Float32[ + -1 0 1 + -2 0 2 + -1 0 1 + ] ./ 8 # Sobel + + d_image = MtlMatrix(image) + d_kernel = MtlMatrix(kernel) + + result = Array(imfilter(d_image, d_kernel)) + @test size(result) == size(image) + + # Verify correctness against FFT + result_fft = Array(conv_fft(d_image, d_kernel; dims = (1, 2), mode = :same)) + @test isapprox(result, result_fft, rtol = 1.0e-4) + end + + @testset "Large Kernel (uses FFT)" begin + image = rand(Float32, 256, 256) + kernel = rand(Float32, 15, 15) # Larger than threshold + + d_image = MtlMatrix(image) + d_kernel = MtlMatrix(kernel) + + result = Array(imfilter(d_image, d_kernel)) + @test size(result) == size(image) + end + end + + # ============================================================================ + # Unified conv() API Tests + # ============================================================================ + + @testset "Unified conv() API" begin + @testset "1D Convolution" begin + signal = MtlVector(rand(Float32, 1000)) + kernel = MtlVector(rand(Float32, 10)) + + result_full = conv(signal, kernel; mode = :full) + result_same = conv(signal, kernel; mode = :same) + result_valid = conv(signal, kernel; mode = :valid) + + @test length(result_full) == 1009 + @test length(result_same) == 1000 + @test length(result_valid) == 991 + end + + @testset "2D Auto-Selection" begin + image = MtlMatrix(rand(Float32, 128, 128)) + small_kernel = MtlMatrix(rand(Float32, 3, 3)) + large_kernel = MtlMatrix(rand(Float32, 15, 15)) + + # Small kernel should auto-select direct + result_small = conv(image, small_kernel; mode = :same) + @test size(result_small) == size(image) + + # Large kernel should auto-select FFT + result_large = conv(image, large_kernel; mode = :same) + @test size(result_large) == size(image) + + # Both should give same result as explicit FFT + expected_small = conv_fft(image, small_kernel; dims = (1, 2), mode = :same) + expected_large = conv_fft(image, large_kernel; dims = (1, 2), mode = :same) + + @test isapprox(Array(result_small), Array(expected_small), rtol = 1.0e-4) + @test isapprox(Array(result_large), Array(expected_large), rtol = 1.0e-4) + end + + @testset "Algorithm Forcing" begin + image = MtlMatrix(rand(Float32, 64, 64)) + kernel = MtlMatrix(rand(Float32, 3, 3)) + + result_auto = conv(image, kernel; mode = :same, algorithm = :auto) + result_fft = conv(image, kernel; mode = :same, algorithm = :fft) + result_direct = conv(image, kernel; mode = :same, algorithm = :direct) + + # All should give numerically similar results + @test isapprox(Array(result_auto), Array(result_direct), rtol = 1.0e-6) + @test isapprox(Array(result_auto), Array(result_fft), rtol = 1.0e-4) + end + + @testset "Error Handling" begin + signal_1d = MtlVector(rand(Float32, 100)) + kernel_1d = MtlVector(rand(Float32, 10)) + + # 1D doesn't support direct + @test_throws ArgumentError conv(signal_1d, kernel_1d; algorithm = :direct) + + # Invalid algorithm + image = MtlMatrix(rand(Float32, 64, 64)) + kernel = MtlMatrix(rand(Float32, 3, 3)) + @test_throws ArgumentError conv(image, kernel; algorithm = :invalid) + end + + @testset "Complex Arrays" begin + signal = MtlVector(rand(ComplexF32, 100)) + kernel = MtlVector(rand(ComplexF32, 10)) + + result = conv(signal, kernel; mode = :full) + expected = conv_fft(signal, kernel; mode = :full) + @test isapprox(Array(result), Array(expected), rtol = 1.0e-4) + + # Complex doesn't support direct + @test_throws ArgumentError conv(signal, kernel; algorithm = :direct) + end + end + + # ============================================================================ + # Edge Cases + # ============================================================================ + + @testset "Edge Cases" begin + @testset "Single Element Kernel" begin + signal = MtlVector(rand(Float32, 100)) + kernel = MtlVector(Float32[2.0]) + + result = Array(conv_fft(signal, kernel; mode = :full)) + expected = Float32.(Array(signal) .* 2.0) + @test isapprox(result, expected, rtol = 1.0e-4) + end + + @testset "Large Arrays" begin + # Test with larger arrays to ensure stability + signal = MtlVector(rand(Float32, 10000)) + kernel = MtlVector(rand(Float32, 500)) + + result = conv_fft(signal, kernel; mode = :full) + @test length(result) == 10499 + end + + @testset "Non-Square 2D Arrays" begin + image = MtlMatrix(rand(Float32, 64, 128)) + kernel = MtlMatrix(rand(Float32, 3, 5)) + + result = conv(image, kernel; mode = :same) + @test size(result) == (64, 128) + end + end + +end # MPS.is_supported(device()) From 94fe22bdb8bae882b523898157d7a8a6ccc1bcdf Mon Sep 17 00:00:00 2001 From: Kaan Kesgin Date: Wed, 27 May 2026 17:19:20 +0200 Subject: [PATCH 02/12] Add DSP.jl extension for GPU convolution Provide DSP.conv / DSP.xcorr methods for MtlArray via a package extension (MetalDSPExt), dispatching to the FFT-based convolution engine. This is the idiomatic JuliaGPU pattern (extend the standard interface, like AbstractFFTs) and avoids shadowing DSP's exports with a bespoke Metal.conv. Validated against CPU DSP.conv/xcorr for 1-D, 2-D, and 3-D inputs. --- Project.toml | 3 +++ ext/MetalDSPExt.jl | 44 ++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 47 insertions(+) create mode 100644 ext/MetalDSPExt.jl diff --git a/Project.toml b/Project.toml index de50e845d..3da8ebb86 100644 --- a/Project.toml +++ b/Project.toml @@ -33,9 +33,11 @@ StaticArrays = "90137ffa-7385-5640-81b9-e52037218182" UUIDs = "cf7118a7-6976-5b1a-9a39-7adc72f591a4" [weakdeps] +DSP = "717857b8-e6f2-59f4-9121-6e50c889abd2" SpecialFunctions = "276daf66-3868-5448-9aa4-cd146d93841b" [extensions] +MetalDSPExt = "DSP" SpecialFunctionsExt = "SpecialFunctions" [compat] @@ -44,6 +46,7 @@ Adapt = "4.5" BFloat16s = "0.5, 0.6" CEnum = "0.4, 0.5" CodecBzip2 = "0.8.5" +DSP = "0.7, 0.8" ExprTools = "0.1" GPUArrays = "11.5" GPUCompiler = "1.13.2" diff --git a/ext/MetalDSPExt.jl b/ext/MetalDSPExt.jl new file mode 100644 index 000000000..e6321f1d4 --- /dev/null +++ b/ext/MetalDSPExt.jl @@ -0,0 +1,44 @@ +module MetalDSPExt + +# GPU-accelerated linear convolution and cross-correlation for `MtlArray`s. +# +# Dispatches `DSP.conv` / `DSP.xcorr` to Metal's FFT-based convolution engine +# (`Metal.MPSGraphs`), so `using DSP, Metal; conv(a, b)` runs on the GPU instead +# of falling back to scalar CPU indexing. This mirrors how the FFT support +# extends `AbstractFFTs` rather than introducing a parallel `Metal.conv`. + +using Metal +import DSP + +# Element types the MPSGraph FFT/convolution engine supports. +const MtlConvNumber = Union{Float32, Float16, ComplexF32, ComplexF16} + +""" + DSP.conv(u::MtlArray, v::MtlArray; algorithm = :auto) + +Full linear convolution of two `MtlArray`s on the GPU, computed via the FFT +convolution theorem (with an MPS direct-convolution fast path for small 2-D +kernels). Convolves over all dimensions, matching `DSP.conv` semantics. The +`algorithm` keyword accepts `:auto`, `:fft`, or `:direct`. +""" +function DSP.conv( + u::MtlArray{T, N}, v::MtlArray{T, N}; algorithm::Symbol = :auto + ) where {T <: MtlConvNumber, N} + alg = algorithm in (:fft, :direct) ? algorithm : :auto + return Metal.MPSGraphs.conv(u, v; dims = ntuple(identity, N), mode = :full, algorithm = alg) +end + +""" + DSP.xcorr(u::MtlVector, v::MtlVector; padmode = :none) + +Cross-correlation of two GPU vectors, conjugating `v` (the DSP/MATLAB +convention). Only `padmode = :none` (the full correlation) is supported. +""" +function DSP.xcorr( + u::MtlVector{T}, v::MtlVector{T}; padmode::Symbol = :none + ) where {T <: MtlConvNumber} + padmode === :none || throw(ArgumentError("MetalDSPExt only supports padmode = :none")) + return Metal.MPSGraphs.xcorr(u, v; mode = :full) +end + +end # module From a05e8541212901b29ae1791b52ceca6c9623a718 Mon Sep 17 00:00:00 2001 From: Kaan Kesgin Date: Wed, 27 May 2026 17:25:08 +0200 Subject: [PATCH 03/12] Expose convolution only through the DSP extension Remove the bespoke top-level conv/conv_fft/xcorr/... exports from Metal so DSP.conv / DSP.xcorr (via MetalDSPExt) become the sole public interface, mirroring how FFT support extends AbstractFFTs. The convolution engine stays internal to MPSGraphs. - src/Metal.jl: drop bespoke conv-API re-exports - test/mpsgraphs/convolution.jl: import engine symbols from MPSGraphs - test/dsp.jl: public-interface test for DSP.conv / DSP.xcorr - test/Project.toml: add DSP test dep --- src/Metal.jl | 6 ++---- test/Project.toml | 1 + test/dsp.jl | 37 +++++++++++++++++++++++++++++++++++ test/mpsgraphs/convolution.jl | 8 +++++--- 4 files changed, 45 insertions(+), 7 deletions(-) create mode 100644 test/dsp.jl diff --git a/src/Metal.jl b/src/Metal.jl index 01b6cd769..1ddadebc8 100644 --- a/src/Metal.jl +++ b/src/Metal.jl @@ -60,10 +60,8 @@ include("../lib/mps/MPS.jl") export MPS include("../lib/mpsgraphs/MPSGraphs.jl") export MPSGraphs -# Re-export the public convolution API to the top-level Metal namespace. -# Internal helpers (conv_direct, imfilter, plan-cache management) stay in MPSGraphs. -using .MPSGraphs: conv, conv_fft, conv_fft!, conv_fft_fused, xcorr, plan_conv_fft, ConvFFTPlan -export conv, conv_fft, conv_fft!, conv_fft_fused, xcorr, plan_conv_fft, ConvFFTPlan +# The convolution engine lives in MPSGraphs and is exposed publicly through the +# DSP.jl package extension (DSP.conv / DSP.xcorr), not as bespoke Metal exports. # LinearAlgebra include("linalg.jl") diff --git a/test/Project.toml b/test/Project.toml index 322e6edac..94718cdf8 100644 --- a/test/Project.toml +++ b/test/Project.toml @@ -1,6 +1,7 @@ [deps] AbstractFFTs = "621f4979-c628-5d54-868e-fcf4e3e8185c" Adapt = "79e6a3ab-5dfb-504d-930d-738a2a938a0e" +DSP = "717857b8-e6f2-59f4-9121-6e50c889abd2" FFTW = "7a1cc6ca-52ef-59f5-83cd-3a7055c09341" BFloat16s = "ab4f0b2a-ad5b-11e8-123f-65d77653426b" BenchmarkTools = "6e4b80f9-dd63-53aa-95a3-0cdb28fa8baf" diff --git a/test/dsp.jl b/test/dsp.jl new file mode 100644 index 000000000..19e0d35ff --- /dev/null +++ b/test/dsp.jl @@ -0,0 +1,37 @@ +using DSP + +# Public interface for GPU convolution: DSP.conv / DSP.xcorr on MtlArrays, provided +# by the MetalDSPExt package extension. The underlying engine (modes, dims, direct +# path, plan caching) is tested in mpsgraphs/convolution.jl. + +@testset "DSP.conv" begin + @testset "1-D" begin + a = rand(Float32, 64) + b = rand(Float32, 7) + @test Array(conv(MtlArray(a), MtlArray(b))) ≈ conv(a, b) rtol = 1.0f-3 + end + + @testset "2-D" begin + a = rand(Float32, 16, 16) + b = rand(Float32, 3, 3) + @test Array(conv(MtlArray(a), MtlArray(b))) ≈ conv(a, b) rtol = 1.0f-2 + end + + @testset "3-D" begin + a = rand(Float32, 8, 8, 8) + b = rand(Float32, 3, 3, 3) + @test Array(conv(MtlArray(a), MtlArray(b))) ≈ conv(a, b) rtol = 1.0f-2 + end + + @testset "algorithm = :fft" begin + a = rand(Float32, 32, 32) + b = rand(Float32, 5, 5) + @test Array(conv(MtlArray(a), MtlArray(b); algorithm = :fft)) ≈ conv(a, b) rtol = 1.0f-2 + end +end + +@testset "DSP.xcorr" begin + u = rand(Float32, 32) + v = rand(Float32, 20) + @test Array(xcorr(MtlArray(u), MtlArray(v))) ≈ xcorr(u, v) rtol = 1.0f-3 +end diff --git a/test/mpsgraphs/convolution.jl b/test/mpsgraphs/convolution.jl index cd78b3148..07d424a1e 100644 --- a/test/mpsgraphs/convolution.jl +++ b/test/mpsgraphs/convolution.jl @@ -6,9 +6,11 @@ # - Unified conv() API with auto-selection # - Convolution plan caching -# Internal (non-exported) convolution symbols, accessed directly from the submodule. -# Metal exports only the core conv API; these stay internal. -using Metal.MPSGraphs: conv_direct, get_cached_conv_plan, clear_conv_plan_cache!, imfilter +# The convolution engine is internal to MPSGraphs (public access is via DSP.conv). +# These tests exercise the engine directly, so import its symbols from the submodule. +using Metal.MPSGraphs: conv, conv_fft, conv_fft!, conv_fft_fused, xcorr, + plan_conv_fft, ConvFFTPlan, conv_direct, get_cached_conv_plan, + clear_conv_plan_cache!, imfilter # Simple reference convolution for verification (CPU) function ref_conv(u::Vector{T}, v::Vector{T}) where {T} From 9e199944b3e8add55aaa150016d43763a3943103 Mon Sep 17 00:00:00 2001 From: Kaan Kesgin Date: Wed, 27 May 2026 17:53:16 +0200 Subject: [PATCH 04/12] Document convolution performance; remove investigation scripts Add perf/CONVOLUTION_PERFORMANCE.md with measured fused-FFT vs MPS-direct vs ConvFFTPlan benchmarks (basis for the single-implementation trim: the FFT path is fastest in every regime, so the trim costs no performance). Remove the dev-only investigation scripts and revert the perf-only DSP/Image* deps. --- perf/CONVOLUTION_PERFORMANCE.md | 65 ++++ perf/Project.toml | 3 - perf/benchmark_fused_comprehensive.jl | 324 ----------------- perf/benchmark_fused_conv.jl | 96 ----- perf/benchmark_inline_pad.jl | 200 ----------- perf/benchmark_throughput.jl | 129 ------- perf/bottleneck_analysis.jl | 495 -------------------------- perf/padding_alternatives.jl | 265 -------------- perf/simple_bottleneck.jl | 362 ------------------- 9 files changed, 65 insertions(+), 1874 deletions(-) create mode 100644 perf/CONVOLUTION_PERFORMANCE.md delete mode 100644 perf/benchmark_fused_comprehensive.jl delete mode 100644 perf/benchmark_fused_conv.jl delete mode 100644 perf/benchmark_inline_pad.jl delete mode 100644 perf/benchmark_throughput.jl delete mode 100644 perf/bottleneck_analysis.jl delete mode 100644 perf/padding_alternatives.jl delete mode 100644 perf/simple_bottleneck.jl diff --git a/perf/CONVOLUTION_PERFORMANCE.md b/perf/CONVOLUTION_PERFORMANCE.md new file mode 100644 index 000000000..e18f28800 --- /dev/null +++ b/perf/CONVOLUTION_PERFORMANCE.md @@ -0,0 +1,65 @@ +# Convolution performance notes + +Benchmarks backing the design of the GPU convolution support (the `DSP.conv` / +`DSP.xcorr` extension and its internal FFT engine). Measured on Apple Silicon +(Metal), `Float32`, with warmup + `Metal.synchronize()`, reported as the best of +5 runs × 30 calls (ms/call). These numbers drive the "single-implementation" +decision and are intended as source material for the PR description. + +## Single 2-D image, `mode = :same` + +| image | kernel | fused FFT | MPS-direct | fused speedup | +|------:|-------:|----------:|-----------:|--------------:| +| 16² | 3² | 0.53 | 2.30 | 4.3× | +| 32² | 3² | 0.58 | 2.79 | 4.8× | +| 64² | 5² | 0.53 | 3.38 | 6.3× | +| 128² | 5² | 0.59 | 5.76 | 9.8× | +| 256² | 5² | 0.58 | 5.18 | 9.0× | +| 512² | 3²–63² | 0.6–0.8 | 2.4–11.1 | 2–18× | + +The fused FFT path is ~0.6 ms and essentially **kernel-size independent**, while +MPS-direct (`convolution2DWithSourceTensor`) is slower at every size measured — +including the small kernels the old `:auto` heuristic routed to it. + +## Repeated convolution: plan vs one-shot (512², `mode = :full`) + +| kernel | one-shot fused | `ConvFFTPlan` (reuse) | ratio | +|-------:|---------------:|----------------------:|------:| +| 7² | 0.65 | 1.23 | 0.52× | +| 15² | 0.62 | 1.19 | 0.52× | +| 31² | 0.66 | 1.19 | 0.56× | + +The `ConvFFTPlan` repeated-convolution path is **~2× slower** than simply calling +the one-shot fused convolution. The fused graph and its transforms are already +cached by the FFT graph cache, so the plan's separate rfft/irfft + buffer reuse +adds overhead without a payoff. + +## Design decisions + +- **Single implementation = fused FFT.** It is the fastest path in every regime + measured, so collapsing to it (per reviewer guidance on the FFT PR — "one + implementation, like CUDA.jl") costs **no performance**. It also removes the + `:auto` heuristic that mis-routed small kernels to the slower direct path, so + the trim is a net speed-up for that case. +- **Drop MPS-direct** from the signal-convolution engine: always slower here. +- **Drop `ConvFFTPlan` / `plan_conv_fft`:** slower than one-shot — negative value. +- **`imfilter`** routes to the fused FFT path (its old direct branch is removed). + +## Coordination with PR #745 and NNlib + +PR #745 (`MPSGraphs.graph_conv!`) implements **NN-style** convolution +(`convolution2DWithSourceTensor` with stride/dilation/padding/groups) for the +NNlib/Flux lane (issue #210). This work covers **signal-processing** convolution +(`DSP.conv` / `DSP.xcorr`, FFT-based). They are complementary: + +- Dropping MPS-direct here also drops our `convolution2DWithSourceTensor` / + `MPSGraphConvolution2DOpDescriptor` wrappers, **removing the overlap** with #745. +- We do **not** add `NNlib.conv` here — it would duplicate #745. The NN path + belongs in #745 (or a follow-up NNlib extension), not the DSP/signal PR. + +## Caveats + +Numbers are for single-array (all-dims) `Float32` signal convolution on one GPU. +NN-style batched/channelled convolution (many small filters, NCHW layout) is a +different workload where the MPS conv2d primitive is appropriate — that is #745's +domain, not this PR's. diff --git a/perf/Project.toml b/perf/Project.toml index 4abbeca91..decfbe75f 100644 --- a/perf/Project.toml +++ b/perf/Project.toml @@ -1,9 +1,6 @@ [deps] BenchmarkTools = "6e4b80f9-dd63-53aa-95a3-0cdb28fa8baf" -DSP = "717857b8-e6f2-59f4-9121-6e50c889abd2" HTTP = "cd3eb016-35fb-5094-929b-558a96fad6f3" -ImageCore = "a09fc81d-aa75-5fe9-8630-4744c3626534" -ImageFiltering = "6a3955dd-da59-5b1f-98d4-e7296123deb5" JSON = "682c06a0-de6a-54ab-a142-c8b1cf79cde6" Metal = "dde4c033-4e86-420c-a63e-0dd931031962" StableRNGs = "860ef19b-820b-49d6-a774-d7a799459cd3" diff --git a/perf/benchmark_fused_comprehensive.jl b/perf/benchmark_fused_comprehensive.jl deleted file mode 100644 index 8f4ef910a..000000000 --- a/perf/benchmark_fused_comprehensive.jl +++ /dev/null @@ -1,324 +0,0 @@ -using Metal -using Metal.MPSGraphs: conv_fft, conv_fft_fused -using DSP -using Statistics - -println("=" ^ 70) -println("COMPREHENSIVE FUSED CONVOLUTION BENCHMARKS") -println("=" ^ 70) -println("Note: conv_fft() now uses fused implementation automatically") -println() - -# ============================================================================ -# 1D BENCHMARKS -# ============================================================================ - -println("=" ^ 70) -println("1D CONVOLUTION: GPU (Fused) vs CPU (DSP.jl)") -println("=" ^ 70) - -configs_1d = [ - (10_000, 100), - (50_000, 100), - (100_000, 100), - (100_000, 500), - (500_000, 500), - (1_000_000, 500), - (1_000_000, 1000), - (5_000_000, 1000), -] - -results_1d = [] - -for (signal_size, kernel_size) in configs_1d - # Create test data - signal_cpu = rand(Float32, signal_size) - kernel_cpu = rand(Float32, kernel_size) - signal_gpu = MtlVector(signal_cpu) - kernel_gpu = MtlVector(kernel_cpu) - - # Warmup - _ = conv_fft(signal_gpu, kernel_gpu) - Metal.synchronize() - - # Benchmark GPU - n_iters = 10 - times_gpu = Float64[] - for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - _ = conv_fft(signal_gpu, kernel_gpu) - Metal.synchronize() - end - push!(times_gpu, t * 1000) - end - - # Benchmark CPU - cpu_iters = signal_size > 1_000_000 ? 3 : 5 - times_cpu = Float64[] - for _ in 1:cpu_iters - t = @elapsed begin - _ = DSP.conv(signal_cpu, kernel_cpu) - end - push!(times_cpu, t * 1000) - end - - med_gpu = median(times_gpu) - med_cpu = median(times_cpu) - speedup = med_cpu / med_gpu - winner = speedup >= 1.0 ? "GPU" : "CPU" - - push!(results_1d, (signal_size, kernel_size, med_gpu, med_cpu, speedup, winner)) - println("Signal=$(signal_size), Kernel=$(kernel_size): GPU=$(round(med_gpu, digits=2))ms, CPU=$(round(med_cpu, digits=2))ms, Speedup=$(round(speedup, digits=2))x [$winner]") -end - -# ============================================================================ -# 2D BENCHMARKS -# ============================================================================ - -println("\n" * "=" ^ 70) -println("2D CONVOLUTION: GPU (Fused) vs CPU (imfilter-style)") -println("=" ^ 70) - -configs_2d = [ - ((128, 128), (5, 5)), - ((256, 256), (5, 5)), - ((256, 256), (15, 15)), - ((512, 512), (5, 5)), - ((512, 512), (15, 15)), - ((1024, 1024), (5, 5)), - ((1024, 1024), (15, 15)), - ((2048, 2048), (15, 15)), -] - -results_2d = [] - -for (image_size, kernel_size) in configs_2d - # Create test data - image_cpu = rand(Float32, image_size...) - kernel_cpu = rand(Float32, kernel_size...) - image_gpu = MtlMatrix(image_cpu) - kernel_gpu = MtlMatrix(kernel_cpu) - - # Warmup - _ = conv_fft(image_gpu, kernel_gpu; dims=(1,2)) - Metal.synchronize() - - # Benchmark GPU - n_iters = 10 - times_gpu = Float64[] - for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - _ = conv_fft(image_gpu, kernel_gpu; dims=(1,2)) - Metal.synchronize() - end - push!(times_gpu, t * 1000) - end - - # Benchmark CPU (using DSP.conv for 2D as reference) - cpu_iters = prod(image_size) > 1_000_000 ? 3 : 5 - times_cpu = Float64[] - for _ in 1:cpu_iters - t = @elapsed begin - _ = DSP.conv(image_cpu, kernel_cpu) - end - push!(times_cpu, t * 1000) - end - - med_gpu = median(times_gpu) - med_cpu = median(times_cpu) - speedup = med_cpu / med_gpu - winner = speedup >= 1.0 ? "GPU" : "CPU" - - push!(results_2d, (image_size, kernel_size, med_gpu, med_cpu, speedup, winner)) - println("Image=$(image_size), Kernel=$(kernel_size): GPU=$(round(med_gpu, digits=2))ms, CPU=$(round(med_cpu, digits=2))ms, Speedup=$(round(speedup, digits=2))x [$winner]") -end - -# ============================================================================ -# 3D BENCHMARKS -# ============================================================================ - -println("\n" * "=" ^ 70) -println("3D CONVOLUTION: GPU (Fused) vs CPU (DSP.conv)") -println("=" ^ 70) - -configs_3d = [ - ((32, 32, 32), (3, 3, 3)), - ((64, 64, 64), (3, 3, 3)), - ((64, 64, 64), (5, 5, 5)), - ((128, 128, 128), (3, 3, 3)), - ((128, 128, 128), (5, 5, 5)), - ((256, 256, 64), (5, 5, 5)), - ((256, 256, 128), (5, 5, 5)), -] - -results_3d = [] - -for (volume_size, kernel_size) in configs_3d - # Create test data - volume_cpu = rand(Float32, volume_size...) - kernel_cpu = rand(Float32, kernel_size...) - volume_gpu = MtlArray(volume_cpu) - kernel_gpu = MtlArray(kernel_cpu) - - # Warmup - _ = conv_fft(volume_gpu, kernel_gpu; dims=(1,2,3)) - Metal.synchronize() - - # Benchmark GPU - n_iters = 8 - times_gpu = Float64[] - for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - _ = conv_fft(volume_gpu, kernel_gpu; dims=(1,2,3)) - Metal.synchronize() - end - push!(times_gpu, t * 1000) - end - - # Benchmark CPU - cpu_iters = prod(volume_size) > 500_000 ? 2 : 3 - times_cpu = Float64[] - for _ in 1:cpu_iters - t = @elapsed begin - _ = DSP.conv(volume_cpu, kernel_cpu) - end - push!(times_cpu, t * 1000) - end - - med_gpu = median(times_gpu) - med_cpu = median(times_cpu) - speedup = med_cpu / med_gpu - winner = speedup >= 1.0 ? "GPU" : "CPU" - - push!(results_3d, (volume_size, kernel_size, med_gpu, med_cpu, speedup, winner)) - println("Volume=$(volume_size), Kernel=$(kernel_size): GPU=$(round(med_gpu, digits=2))ms, CPU=$(round(med_cpu, digits=2))ms, Speedup=$(round(speedup, digits=2))x [$winner]") -end - -# ============================================================================ -# THROUGHPUT BENCHMARK -# ============================================================================ - -println("\n" * "=" ^ 70) -println("THROUGHPUT BENCHMARK (Async Pipeline, batch=50)") -println("=" ^ 70) - -# 1D throughput -signal_1d = MtlVector(rand(Float32, 1_000_000)) -kernel_1d = MtlVector(rand(Float32, 500)) -signal_1d_cpu = rand(Float32, 1_000_000) -kernel_1d_cpu = rand(Float32, 500) - -batch = 50 -Metal.synchronize() -t_gpu_1d = @elapsed begin - for _ in 1:batch - _ = conv_fft(signal_1d, kernel_1d) - end - Metal.synchronize() -end -tput_gpu_1d = batch / t_gpu_1d - -t_cpu_1d = @elapsed begin - for _ in 1:batch - _ = DSP.conv(signal_1d_cpu, kernel_1d_cpu) - end -end -tput_cpu_1d = batch / t_cpu_1d - -println("\n1D (1M × 500):") -println(" GPU: $(round(tput_gpu_1d, digits=1)) ops/sec ($(round(1000/tput_gpu_1d, digits=2)) ms/op)") -println(" CPU: $(round(tput_cpu_1d, digits=1)) ops/sec ($(round(1000/tput_cpu_1d, digits=2)) ms/op)") -println(" Speedup: $(round(tput_gpu_1d / tput_cpu_1d, digits=2))x") - -# 2D throughput -image_2d = MtlMatrix(rand(Float32, 512, 512)) -kernel_2d = MtlMatrix(rand(Float32, 15, 15)) -image_2d_cpu = rand(Float32, 512, 512) -kernel_2d_cpu = rand(Float32, 15, 15) - -Metal.synchronize() -t_gpu_2d = @elapsed begin - for _ in 1:batch - _ = conv_fft(image_2d, kernel_2d; dims=(1,2)) - end - Metal.synchronize() -end -tput_gpu_2d = batch / t_gpu_2d - -t_cpu_2d = @elapsed begin - for _ in 1:batch - _ = DSP.conv(image_2d_cpu, kernel_2d_cpu) - end -end -tput_cpu_2d = batch / t_cpu_2d - -println("\n2D (512×512 × 15×15):") -println(" GPU: $(round(tput_gpu_2d, digits=1)) ops/sec ($(round(1000/tput_gpu_2d, digits=2)) ms/op)") -println(" CPU: $(round(tput_cpu_2d, digits=1)) ops/sec ($(round(1000/tput_cpu_2d, digits=2)) ms/op)") -println(" Speedup: $(round(tput_gpu_2d / tput_cpu_2d, digits=2))x") - -# 3D throughput -volume_3d = MtlArray(rand(Float32, 64, 64, 64)) -kernel_3d = MtlArray(rand(Float32, 5, 5, 5)) -volume_3d_cpu = rand(Float32, 64, 64, 64) -kernel_3d_cpu = rand(Float32, 5, 5, 5) - -batch_3d = 20 -Metal.synchronize() -t_gpu_3d = @elapsed begin - for _ in 1:batch_3d - _ = conv_fft(volume_3d, kernel_3d; dims=(1,2,3)) - end - Metal.synchronize() -end -tput_gpu_3d = batch_3d / t_gpu_3d - -t_cpu_3d = @elapsed begin - for _ in 1:batch_3d - _ = DSP.conv(volume_3d_cpu, kernel_3d_cpu) - end -end -tput_cpu_3d = batch_3d / t_cpu_3d - -println("\n3D (64³ × 5³):") -println(" GPU: $(round(tput_gpu_3d, digits=1)) ops/sec ($(round(1000/tput_gpu_3d, digits=2)) ms/op)") -println(" CPU: $(round(tput_cpu_3d, digits=1)) ops/sec ($(round(1000/tput_cpu_3d, digits=2)) ms/op)") -println(" Speedup: $(round(tput_gpu_3d / tput_cpu_3d, digits=2))x") - -# ============================================================================ -# SUMMARY TABLES -# ============================================================================ - -println("\n" * "=" ^ 70) -println("SUMMARY TABLES (for documentation)") -println("=" ^ 70) - -println("\n### 1D Convolution Results") -println("| Signal | Kernel | GPU (ms) | CPU (ms) | Speedup | Winner |") -println("|--------|--------|----------|----------|---------|--------|") -for (ss, ks, gpu, cpu, speedup, winner) in results_1d - speedup_str = speedup >= 1.0 ? "**$(round(speedup, digits=2))x**" : "$(round(speedup, digits=2))x" - winner_str = winner == "GPU" ? "**GPU**" : "CPU" - println("| $(ss) | $(ks) | $(round(gpu, digits=2)) | $(round(cpu, digits=2)) | $(speedup_str) | $(winner_str) |") -end - -println("\n### 2D Convolution Results") -println("| Image | Kernel | GPU (ms) | CPU (ms) | Speedup | Winner |") -println("|-------|--------|----------|----------|---------|--------|") -for (is, ks, gpu, cpu, speedup, winner) in results_2d - speedup_str = speedup >= 1.0 ? "**$(round(speedup, digits=2))x**" : "$(round(speedup, digits=2))x" - winner_str = winner == "GPU" ? "**GPU**" : "CPU" - println("| $(is[1])×$(is[2]) | $(ks[1])×$(ks[2]) | $(round(gpu, digits=2)) | $(round(cpu, digits=2)) | $(speedup_str) | $(winner_str) |") -end - -println("\n### 3D Convolution Results") -println("| Volume | Kernel | GPU (ms) | CPU (ms) | Speedup | Winner |") -println("|--------|--------|----------|----------|---------|--------|") -for (vs, ks, gpu, cpu, speedup, winner) in results_3d - speedup_str = speedup >= 1.0 ? "**$(round(speedup, digits=2))x**" : "$(round(speedup, digits=2))x" - winner_str = winner == "GPU" ? "**GPU**" : "CPU" - println("| $(vs[1])×$(vs[2])×$(vs[3]) | $(ks[1])×$(ks[2])×$(ks[3]) | $(round(gpu, digits=2)) | $(round(cpu, digits=2)) | $(speedup_str) | $(winner_str) |") -end diff --git a/perf/benchmark_fused_conv.jl b/perf/benchmark_fused_conv.jl deleted file mode 100644 index 7388fc804..000000000 --- a/perf/benchmark_fused_conv.jl +++ /dev/null @@ -1,96 +0,0 @@ -using Metal -using Metal.MPSGraphs: conv_fft, conv_fft_fused -using DSP -using Statistics - -println("=" ^ 60) -println("LATENCY BENCHMARK: Fused vs Existing vs CPU") -println("=" ^ 60) - -# Test configurations -configs = [ - (10_000, 100), - (100_000, 100), - (100_000, 500), - (500_000, 500), - (1_000_000, 500), - (1_000_000, 1000), -] - -results = [] - -for (signal_size, kernel_size) in configs - println("\n--- Signal: $(signal_size), Kernel: $(kernel_size) ---") - - # Create test data - signal_cpu = rand(Float32, signal_size) - kernel_cpu = rand(Float32, kernel_size) - signal_gpu = MtlVector(signal_cpu) - kernel_gpu = MtlVector(kernel_cpu) - - # Warmup - _ = conv_fft(signal_gpu, kernel_gpu) - _ = conv_fft_fused(signal_gpu, kernel_gpu) - Metal.synchronize() - - # Benchmark existing implementation - n_iters = 10 - times_existing = Float64[] - for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - _ = conv_fft(signal_gpu, kernel_gpu) - Metal.synchronize() - end - push!(times_existing, t * 1000) # ms - end - - # Benchmark fused implementation - times_fused = Float64[] - for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - _ = conv_fft_fused(signal_gpu, kernel_gpu) - Metal.synchronize() - end - push!(times_fused, t * 1000) # ms - end - - # Benchmark CPU (fewer iterations for large sizes) - cpu_iters = signal_size > 500_000 ? 3 : 5 - times_cpu = Float64[] - for _ in 1:cpu_iters - t = @elapsed begin - _ = DSP.conv(signal_cpu, kernel_cpu) - end - push!(times_cpu, t * 1000) # ms - end - - med_existing = median(times_existing) - med_fused = median(times_fused) - med_cpu = median(times_cpu) - - speedup_fused_vs_existing = med_existing / med_fused - speedup_fused_vs_cpu = med_cpu / med_fused - speedup_existing_vs_cpu = med_cpu / med_existing - - println(" Existing GPU: $(round(med_existing, digits=3)) ms") - println(" Fused GPU: $(round(med_fused, digits=3)) ms") - println(" CPU (DSP): $(round(med_cpu, digits=3)) ms") - println(" Fused vs Existing: $(round(speedup_fused_vs_existing, digits=2))x") - println(" Fused vs CPU: $(round(speedup_fused_vs_cpu, digits=2))x") - println(" Existing vs CPU: $(round(speedup_existing_vs_cpu, digits=2))x") - - push!(results, (signal_size, kernel_size, med_existing, med_fused, med_cpu)) -end - -println("\n" * "=" ^ 60) -println("SUMMARY TABLE") -println("=" ^ 60) -println("Signal | Kernel | Existing | Fused | CPU | Fused/Exist | Fused/CPU") -println("-" ^ 80) -for (ss, ks, exist, fused, cpu) in results - speedup1 = round(exist / fused, digits=2) - speedup2 = round(cpu / fused, digits=2) - println("$(lpad(ss, 9)) | $(lpad(ks, 6)) | $(lpad(round(exist, digits=2), 8)) | $(lpad(round(fused, digits=2), 7)) | $(lpad(round(cpu, digits=2), 7)) | $(lpad(speedup1, 11))x | $(lpad(speedup2, 9))x") -end diff --git a/perf/benchmark_inline_pad.jl b/perf/benchmark_inline_pad.jl deleted file mode 100644 index e6701f55b..000000000 --- a/perf/benchmark_inline_pad.jl +++ /dev/null @@ -1,200 +0,0 @@ -using Metal -using Metal.MPSGraphs: conv_fft_inline_pad, conv_fft_fused -using DSP -using Statistics - -println("=" ^ 70) -println("INLINE PADDING BENCHMARK: Fused vs Inline-Pad vs CPU") -println("=" ^ 70) - -configs_1d = [ - (10_000, 100), - (50_000, 100), - (100_000, 100), - (100_000, 500), - (500_000, 500), - (1_000_000, 500), -] - -println("\n1D CONVOLUTION") -println("-" ^ 70) -println("Signal | Kernel | Fused(ms) | Inline(ms) | Speedup | CPU(ms)") -println("-" ^ 70) - -for (signal_size, kernel_size) in configs_1d - signal_cpu = rand(Float32, signal_size) - kernel_cpu = rand(Float32, kernel_size) - signal_gpu = MtlVector(signal_cpu) - kernel_gpu = MtlVector(kernel_cpu) - - # Warmup - _ = conv_fft_fused(signal_gpu, kernel_gpu) - _ = conv_fft_inline_pad(signal_gpu, kernel_gpu) - Metal.synchronize() - - n_iters = 10 - - # Benchmark fused - times_fused = Float64[] - for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - _ = conv_fft_fused(signal_gpu, kernel_gpu) - Metal.synchronize() - end - Base.push!(times_fused, t * 1000) - end - - # Benchmark inline pad - times_inline = Float64[] - for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - _ = conv_fft_inline_pad(signal_gpu, kernel_gpu) - Metal.synchronize() - end - Base.push!(times_inline, t * 1000) - end - - # Benchmark CPU - cpu_iters = signal_size > 500_000 ? 3 : 5 - times_cpu = Float64[] - for _ in 1:cpu_iters - t = @elapsed begin - _ = DSP.conv(signal_cpu, kernel_cpu) - end - Base.push!(times_cpu, t * 1000) - end - - med_fused = median(times_fused) - med_inline = median(times_inline) - med_cpu = median(times_cpu) - speedup = med_fused / med_inline - - println("$(lpad(signal_size, 10)) | $(lpad(kernel_size, 6)) | $(lpad(round(med_fused, digits=2), 9)) | $(lpad(round(med_inline, digits=2), 10)) | $(lpad(round(speedup, digits=2), 7))x | $(lpad(round(med_cpu, digits=2), 6))") -end - -println("\n2D CONVOLUTION") -println("-" ^ 70) -println("Image | Kernel | Fused(ms) | Inline(ms) | Speedup | CPU(ms)") -println("-" ^ 70) - -configs_2d = [ - ((128, 128), (5, 5)), - ((256, 256), (5, 5)), - ((256, 256), (15, 15)), - ((512, 512), (15, 15)), - ((1024, 1024), (15, 15)), -] - -for (image_size, kernel_size) in configs_2d - image_cpu = rand(Float32, image_size...) - kernel_cpu = rand(Float32, kernel_size...) - image_gpu = MtlMatrix(image_cpu) - kernel_gpu = MtlMatrix(kernel_cpu) - - # Warmup - _ = conv_fft_fused(image_gpu, kernel_gpu) - _ = conv_fft_inline_pad(image_gpu, kernel_gpu) - Metal.synchronize() - - n_iters = 10 - - times_fused = Float64[] - for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - _ = conv_fft_fused(image_gpu, kernel_gpu) - Metal.synchronize() - end - Base.push!(times_fused, t * 1000) - end - - times_inline = Float64[] - for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - _ = conv_fft_inline_pad(image_gpu, kernel_gpu) - Metal.synchronize() - end - Base.push!(times_inline, t * 1000) - end - - cpu_iters = prod(image_size) > 500_000 ? 3 : 5 - times_cpu = Float64[] - for _ in 1:cpu_iters - t = @elapsed DSP.conv(image_cpu, kernel_cpu) - Base.push!(times_cpu, t * 1000) - end - - med_fused = median(times_fused) - med_inline = median(times_inline) - med_cpu = median(times_cpu) - speedup = med_fused / med_inline - - image_str = "$(image_size[1])x$(image_size[2])" - kernel_str = "$(kernel_size[1])x$(kernel_size[2])" - println("$(lpad(image_str, 10)) | $(lpad(kernel_str, 6)) | $(lpad(round(med_fused, digits=2), 9)) | $(lpad(round(med_inline, digits=2), 10)) | $(lpad(round(speedup, digits=2), 7))x | $(lpad(round(med_cpu, digits=2), 6))") -end - -println("\n3D CONVOLUTION") -println("-" ^ 70) -println("Volume | Kernel | Fused(ms) | Inline(ms) | Speedup | CPU(ms)") -println("-" ^ 70) - -configs_3d = [ - ((32, 32, 32), (3, 3, 3)), - ((64, 64, 64), (3, 3, 3)), - ((64, 64, 64), (5, 5, 5)), - ((128, 128, 128), (5, 5, 5)), -] - -for (vol_size, kernel_size) in configs_3d - vol_cpu = rand(Float32, vol_size...) - kernel_cpu = rand(Float32, kernel_size...) - vol_gpu = MtlArray(vol_cpu) - kernel_gpu = MtlArray(kernel_cpu) - - # Warmup - _ = conv_fft_fused(vol_gpu, kernel_gpu) - _ = conv_fft_inline_pad(vol_gpu, kernel_gpu) - Metal.synchronize() - - n_iters = 8 - - times_fused = Float64[] - for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - _ = conv_fft_fused(vol_gpu, kernel_gpu) - Metal.synchronize() - end - Base.push!(times_fused, t * 1000) - end - - times_inline = Float64[] - for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - _ = conv_fft_inline_pad(vol_gpu, kernel_gpu) - Metal.synchronize() - end - Base.push!(times_inline, t * 1000) - end - - cpu_iters = prod(vol_size) > 500_000 ? 2 : 3 - times_cpu = Float64[] - for _ in 1:cpu_iters - t = @elapsed DSP.conv(vol_cpu, kernel_cpu) - Base.push!(times_cpu, t * 1000) - end - - med_fused = median(times_fused) - med_inline = median(times_inline) - med_cpu = median(times_cpu) - speedup = med_fused / med_inline - - vol_str = "$(vol_size[1])x$(vol_size[2])x$(vol_size[3])" - kernel_str = "$(kernel_size[1])x$(kernel_size[2])x$(kernel_size[3])" - println("$(lpad(vol_str, 10)) | $(lpad(kernel_str, 6)) | $(lpad(round(med_fused, digits=2), 9)) | $(lpad(round(med_inline, digits=2), 10)) | $(lpad(round(speedup, digits=2), 7))x | $(lpad(round(med_cpu, digits=2), 6))") -end diff --git a/perf/benchmark_throughput.jl b/perf/benchmark_throughput.jl deleted file mode 100644 index c3d2c994d..000000000 --- a/perf/benchmark_throughput.jl +++ /dev/null @@ -1,129 +0,0 @@ -using Metal -using Metal.MPSGraphs: conv_fft, conv_fft_fused -using DSP -using Statistics - -println("=" ^ 70) -println("THROUGHPUT BENCHMARK: Async Pipeline Performance") -println("=" ^ 70) - -# Test configuration: moderate size where GPU wins -signal_size = 1_000_000 -kernel_size = 500 - -println("\nTest config: Signal=$(signal_size), Kernel=$(kernel_size)") -println("-" ^ 70) - -# Create test data -signal_cpu = rand(Float32, signal_size) -kernel_cpu = rand(Float32, kernel_size) -signal_gpu = MtlVector(signal_cpu) -kernel_gpu = MtlVector(kernel_cpu) - -# Warmup -for _ in 1:3 - _ = conv_fft_fused(signal_gpu, kernel_gpu) -end -Metal.synchronize() - -# Test different batch sizes for throughput -batch_sizes = [1, 5, 10, 20, 50] - -println("\n--- FUSED GPU: Sync vs Async Pipeline ---") -println("Batch | Sync (ms/op) | Async (ms/op) | Speedup") -println("-" ^ 50) - -for batch in batch_sizes - # Sync mode: wait after each operation - Metal.synchronize() - t_sync = @elapsed begin - for _ in 1:batch - _ = conv_fft_fused(signal_gpu, kernel_gpu) - Metal.synchronize() - end - end - ms_sync = (t_sync / batch) * 1000 - - # Async mode: queue all, sync once - Metal.synchronize() - t_async = @elapsed begin - for _ in 1:batch - _ = conv_fft_fused(signal_gpu, kernel_gpu) - end - Metal.synchronize() - end - ms_async = (t_async / batch) * 1000 - - speedup = ms_sync / ms_async - println("$(lpad(batch, 5)) | $(lpad(round(ms_sync, digits=3), 12)) | $(lpad(round(ms_async, digits=3), 13)) | $(round(speedup, digits=2))x") -end - -println("\n--- EXISTING GPU: Sync vs Async Pipeline ---") -println("Batch | Sync (ms/op) | Async (ms/op) | Speedup") -println("-" ^ 50) - -for batch in batch_sizes - # Sync mode - Metal.synchronize() - t_sync = @elapsed begin - for _ in 1:batch - _ = conv_fft(signal_gpu, kernel_gpu) - Metal.synchronize() - end - end - ms_sync = (t_sync / batch) * 1000 - - # Async mode - Metal.synchronize() - t_async = @elapsed begin - for _ in 1:batch - _ = conv_fft(signal_gpu, kernel_gpu) - end - Metal.synchronize() - end - ms_async = (t_async / batch) * 1000 - - speedup = ms_sync / ms_async - println("$(lpad(batch, 5)) | $(lpad(round(ms_sync, digits=3), 12)) | $(lpad(round(ms_async, digits=3), 13)) | $(round(speedup, digits=2))x") -end - -# Compare throughput: fused async vs existing async vs CPU -println("\n" * "=" ^ 70) -println("MAXIMUM THROUGHPUT COMPARISON (batch=50, async)") -println("=" ^ 70) - -batch = 50 - -# Fused async -Metal.synchronize() -t_fused = @elapsed begin - for _ in 1:batch - _ = conv_fft_fused(signal_gpu, kernel_gpu) - end - Metal.synchronize() -end -tput_fused = batch / t_fused - -# Existing async -Metal.synchronize() -t_existing = @elapsed begin - for _ in 1:batch - _ = conv_fft(signal_gpu, kernel_gpu) - end - Metal.synchronize() -end -tput_existing = batch / t_existing - -# CPU -t_cpu = @elapsed begin - for _ in 1:batch - _ = DSP.conv(signal_cpu, kernel_cpu) - end -end -tput_cpu = batch / t_cpu - -println("\nFused GPU: $(round(tput_fused, digits=1)) ops/sec ($(round(1000/tput_fused, digits=2)) ms/op)") -println("Existing GPU: $(round(tput_existing, digits=1)) ops/sec ($(round(1000/tput_existing, digits=2)) ms/op)") -println("CPU (DSP): $(round(tput_cpu, digits=1)) ops/sec ($(round(1000/tput_cpu, digits=2)) ms/op)") -println("\nFused vs CPU: $(round(tput_fused / tput_cpu, digits=2))x throughput") -println("Fused vs Existing: $(round(tput_fused / tput_existing, digits=2))x throughput") diff --git a/perf/bottleneck_analysis.jl b/perf/bottleneck_analysis.jl deleted file mode 100644 index 90b61f360..000000000 --- a/perf/bottleneck_analysis.jl +++ /dev/null @@ -1,495 +0,0 @@ -using Metal -using Metal.MPSGraphs: conv_fft, conv_fft_fused, conv_fft_inline_pad -using Metal.MPSGraphs: _get_cached_fused_conv_graph, FusedConvGraphKey, nextfastfft, _conv_output_size -using Metal.MPSGraphs: MPSGraphTensorData, MPSCommandBuffer, NSDictionary, encode!, commit!, wait_completed, default_exec_desc, nil -using Metal.MPSGraphs: _get_cached_buffers -using Statistics -using Printf - -println("=" ^ 80) -println("COMPREHENSIVE BOTTLENECK ANALYSIS") -println("=" ^ 80) - -# ============================================================================ -# SECTION 1: Component-level breakdown for FUSED implementation -# ============================================================================ - -function analyze_fused_components(signal_size, kernel_size; n_iters=50) - println("\n" * "-" ^ 80) - println("FUSED IMPLEMENTATION: Signal=$signal_size, Kernel=$kernel_size") - println("-" ^ 80) - - # Setup - signal = MtlVector(rand(Float32, signal_size)) - kernel = MtlVector(rand(Float32, kernel_size)) - - # Warmup - _ = conv_fft_fused(signal, kernel) - Metal.synchronize() - - # Get sizes - ns, nk = length(signal), length(kernel) - full_size = ns + nk - 1 - output_size = _conv_output_size(ns, nk, :full) - nfft = nextfastfft(full_size) - - println("FFT size: $nfft, Output size: $output_size") - - # Get cached graph and buffers - key = FusedConvGraphKey((nfft,), (nfft,), (output_size,), Float32) - cached = _get_cached_fused_conv_graph(key) - buffers = _get_cached_buffers((nfft,), Float32) - - results = Dict{String, Float64}() - - # 1. Cache lookup (graph) - times = Float64[] - for _ in 1:n_iters - t = @elapsed _get_cached_fused_conv_graph(key) - push!(times, t * 1e6) - end - results["Graph cache lookup"] = median(times) - - # 2. Buffer pool lookup - times = Float64[] - for _ in 1:n_iters - t = @elapsed _get_cached_buffers((nfft,), Float32) - push!(times, t * 1e6) - end - results["Buffer pool lookup"] = median(times) - - # 3. copyto! for signal - times = Float64[] - for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - copyto!(buffers.signal_padded, 1, signal, 1, ns) - Metal.synchronize() - end - push!(times, t * 1e6) - end - results["copyto! signal (sync)"] = median(times) - - # 4. copyto! for kernel - times = Float64[] - for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - copyto!(buffers.kernel_padded, 1, kernel, 1, nk) - Metal.synchronize() - end - push!(times, t * 1e6) - end - results["copyto! kernel (sync)"] = median(times) - - # 5. Zero-padding signal - times = Float64[] - for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - @view(buffers.signal_padded[(ns+1):nfft]) .= 0f0 - Metal.synchronize() - end - push!(times, t * 1e6) - end - results["Zero-pad signal (sync)"] = median(times) - - # 6. Zero-padding kernel - times = Float64[] - for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - @view(buffers.kernel_padded[(nk+1):nfft]) .= 0f0 - Metal.synchronize() - end - push!(times, t * 1e6) - end - results["Zero-pad kernel (sync)"] = median(times) - - # 7. MPSGraphTensorData creation - times = Float64[] - for _ in 1:n_iters - t = @elapsed begin - td1 = MPSGraphTensorData(buffers.signal_padded) - td2 = MPSGraphTensorData(buffers.kernel_padded) - td3 = MPSGraphTensorData(buffers.output) - end - push!(times, t * 1e6) - end - results["MPSGraphTensorData (3x)"] = median(times) - - # 8. NSDictionary creation - td1 = MPSGraphTensorData(buffers.signal_padded) - td2 = MPSGraphTensorData(buffers.kernel_padded) - td3 = MPSGraphTensorData(buffers.output) - times = Float64[] - for _ in 1:n_iters - t = @elapsed begin - feeds = NSDictionary(Dict(cached.signal_placeholder => td1, cached.kernel_placeholder => td2)) - results_dict = NSDictionary(Dict(cached.result => td3)) - end - push!(times, t * 1e6) - end - results["NSDictionary creation"] = median(times) - - # 9. MPSCommandBuffer creation - times = Float64[] - for _ in 1:n_iters - t = @elapsed MPSCommandBuffer(Metal.global_queue(Metal.current_device())) - push!(times, t * 1e6) - end - results["MPSCommandBuffer"] = median(times) - - # 10. encode! - times = Float64[] - for _ in 1:n_iters - td1 = MPSGraphTensorData(buffers.signal_padded) - td2 = MPSGraphTensorData(buffers.kernel_padded) - td3 = MPSGraphTensorData(buffers.output) - feeds = NSDictionary(Dict(cached.signal_placeholder => td1, cached.kernel_placeholder => td2)) - results_ns = NSDictionary(Dict(cached.result => td3)) - cmdbuf = MPSCommandBuffer(Metal.global_queue(Metal.current_device())) - t = @elapsed encode!(cmdbuf, cached.graph, feeds, results_ns, nil, default_exec_desc()) - push!(times, t * 1e6) - end - results["encode!"] = median(times) - - # 11. commit! - times = Float64[] - for _ in 1:n_iters - td1 = MPSGraphTensorData(buffers.signal_padded) - td2 = MPSGraphTensorData(buffers.kernel_padded) - td3 = MPSGraphTensorData(buffers.output) - feeds = NSDictionary(Dict(cached.signal_placeholder => td1, cached.kernel_placeholder => td2)) - results_ns = NSDictionary(Dict(cached.result => td3)) - cmdbuf = MPSCommandBuffer(Metal.global_queue(Metal.current_device())) - encode!(cmdbuf, cached.graph, feeds, results_ns, nil, default_exec_desc()) - t = @elapsed commit!(cmdbuf) - push!(times, t * 1e6) - end - results["commit!"] = median(times) - - # 12. wait_completed - times = Float64[] - for _ in 1:n_iters - td1 = MPSGraphTensorData(buffers.signal_padded) - td2 = MPSGraphTensorData(buffers.kernel_padded) - td3 = MPSGraphTensorData(buffers.output) - feeds = NSDictionary(Dict(cached.signal_placeholder => td1, cached.kernel_placeholder => td2)) - results_ns = NSDictionary(Dict(cached.result => td3)) - cmdbuf = MPSCommandBuffer(Metal.global_queue(Metal.current_device())) - encode!(cmdbuf, cached.graph, feeds, results_ns, nil, default_exec_desc()) - commit!(cmdbuf) - t = @elapsed wait_completed(cmdbuf) - push!(times, t * 1e6) - end - results["wait_completed"] = median(times) - - # 13. Output slice copy - output_arr = MtlVector{Float32}(undef, output_size) - times = Float64[] - for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - copyto!(output_arr, 1, buffers.output, 1, output_size) - Metal.synchronize() - end - push!(times, t * 1e6) - end - results["Output slice copy (sync)"] = median(times) - - # Full conv_fft_fused - times = Float64[] - for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - _ = conv_fft_fused(signal, kernel) - Metal.synchronize() - end - push!(times, t * 1e6) - end - results["TOTAL conv_fft_fused"] = median(times) - - # Print results - println("\nComponent breakdown (μs):") - total_components = 0.0 - for (name, time) in sort(collect(results), by=x->x[2], rev=true) - if name != "TOTAL conv_fft_fused" - total_components += time - pct = 100 * time / results["TOTAL conv_fft_fused"] - @printf(" %-30s %8.1f μs (%5.1f%%)\n", name, time, pct) - end - end - println() - @printf(" %-30s %8.1f μs\n", "Sum of components:", total_components) - @printf(" %-30s %8.1f μs\n", "Actual total:", results["TOTAL conv_fft_fused"]) - - return results -end - -# ============================================================================ -# SECTION 2: Compare all three implementations -# ============================================================================ - -function compare_implementations(signal_size, kernel_size; n_iters=30) - println("\n" * "=" ^ 80) - println("IMPLEMENTATION COMPARISON: Signal=$signal_size, Kernel=$kernel_size") - println("=" ^ 80) - - signal_cpu = rand(Float32, signal_size) - kernel_cpu = rand(Float32, kernel_size) - signal = MtlVector(signal_cpu) - kernel = MtlVector(kernel_cpu) - - # Warmup all - _ = conv_fft(signal, kernel) - _ = conv_fft_fused(signal, kernel) - _ = conv_fft_inline_pad(signal, kernel) - Metal.synchronize() - - results = Dict{String, Float64}() - - # conv_fft (original, now uses fused internally) - times = Float64[] - for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - _ = conv_fft(signal, kernel) - Metal.synchronize() - end - push!(times, t * 1000) - end - results["conv_fft"] = median(times) - - # conv_fft_fused (explicit) - times = Float64[] - for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - _ = conv_fft_fused(signal, kernel) - Metal.synchronize() - end - push!(times, t * 1000) - end - results["conv_fft_fused"] = median(times) - - # conv_fft_inline_pad - times = Float64[] - for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - _ = conv_fft_inline_pad(signal, kernel) - Metal.synchronize() - end - push!(times, t * 1000) - end - results["conv_fft_inline_pad"] = median(times) - - println("\nLatency comparison (ms):") - for (name, time) in sort(collect(results), by=x->x[2]) - @printf(" %-25s %8.3f ms\n", name, time) - end - - fastest = minimum(values(results)) - println("\nSpeedups vs fastest:") - for (name, time) in sort(collect(results), by=x->x[2]) - @printf(" %-25s %5.2fx\n", name, time / fastest) - end - - return results -end - -# ============================================================================ -# SECTION 3: Async pipeline analysis -# ============================================================================ - -function analyze_async_pipeline(signal_size, kernel_size; batch_sizes=[1, 5, 10, 20, 50]) - println("\n" * "=" ^ 80) - println("ASYNC PIPELINE ANALYSIS: Signal=$signal_size, Kernel=$kernel_size") - println("=" ^ 80) - - signal = MtlVector(rand(Float32, signal_size)) - kernel = MtlVector(rand(Float32, kernel_size)) - - # Warmup - for _ in 1:3 - _ = conv_fft_fused(signal, kernel) - end - Metal.synchronize() - - println("\nBatch | Sync (ms/op) | Async (ms/op) | Pipeline Speedup") - println("-" ^ 60) - - for batch in batch_sizes - # Sync mode: wait after each - Metal.synchronize() - t_sync = @elapsed begin - for _ in 1:batch - _ = conv_fft_fused(signal, kernel) - Metal.synchronize() - end - end - ms_sync = (t_sync / batch) * 1000 - - # Async mode: queue all, wait once - Metal.synchronize() - t_async = @elapsed begin - for _ in 1:batch - _ = conv_fft_fused(signal, kernel) - end - Metal.synchronize() - end - ms_async = (t_async / batch) * 1000 - - speedup = ms_sync / ms_async - @printf("%5d | %12.3f | %13.3f | %5.2fx\n", batch, ms_sync, ms_async, speedup) - end -end - -# ============================================================================ -# SECTION 4: Memory operation breakdown -# ============================================================================ - -function analyze_memory_operations(signal_size, kernel_size; n_iters=50) - println("\n" * "=" ^ 80) - println("MEMORY OPERATION DEEP DIVE: Signal=$signal_size, Kernel=$kernel_size") - println("=" ^ 80) - - ns, nk = signal_size, kernel_size - full_size = ns + nk - 1 - nfft = nextfastfft(full_size) - - println("Sizes: signal=$ns, kernel=$nk, FFT=$nfft, padding_signal=$(nfft-ns), padding_kernel=$(nfft-nk)") - - signal = MtlVector(rand(Float32, signal_size)) - kernel = MtlVector(rand(Float32, kernel_size)) - - # Pre-allocate - signal_padded = MtlVector{Float32}(undef, nfft) - kernel_padded = MtlVector{Float32}(undef, nfft) - - # Warmup - copyto!(signal_padded, 1, signal, 1, ns) - Metal.synchronize() - - results = Dict{String, Float64}() - - # Test 1: copyto! with sync - times = Float64[] - for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - copyto!(signal_padded, 1, signal, 1, ns) - Metal.synchronize() - end - push!(times, t * 1e6) - end - results["copyto! signal (with sync)"] = median(times) - - # Test 2: copyto! without sync (just queue time) - times = Float64[] - for _ in 1:n_iters - Metal.synchronize() - t = @elapsed copyto!(signal_padded, 1, signal, 1, ns) - push!(times, t * 1e6) - end - results["copyto! signal (queue only)"] = median(times) - - # Test 3: broadcast zero-fill with sync - times = Float64[] - for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - @view(signal_padded[(ns+1):nfft]) .= 0f0 - Metal.synchronize() - end - push!(times, t * 1e6) - end - results["Zero-fill broadcast (with sync)"] = median(times) - - # Test 4: fill! for zeros - times = Float64[] - for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - fill!(@view(signal_padded[(ns+1):nfft]), 0f0) - Metal.synchronize() - end - push!(times, t * 1e6) - end - results["fill! zeros (with sync)"] = median(times) - - # Test 5: Full array fill - times = Float64[] - for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - fill!(signal_padded, 0f0) - Metal.synchronize() - end - push!(times, t * 1e6) - end - results["fill! full array (with sync)"] = median(times) - - # Test 6: MtlArray allocation - times = Float64[] - for _ in 1:n_iters - GC.gc(false) - t = @elapsed MtlVector{Float32}(undef, nfft) - push!(times, t * 1e6) - end - results["MtlArray allocation"] = median(times) - - # Test 7: Combined copy + zero in one sync - times = Float64[] - for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - copyto!(signal_padded, 1, signal, 1, ns) - @view(signal_padded[(ns+1):nfft]) .= 0f0 - Metal.synchronize() - end - push!(times, t * 1e6) - end - results["copyto! + zero (one sync)"] = median(times) - - # Test 8: Metal.synchronize() alone - times = Float64[] - for _ in 1:n_iters - Metal.synchronize() - t = @elapsed Metal.synchronize() - push!(times, t * 1e6) - end - results["Metal.synchronize() alone"] = median(times) - - println("\nMemory operation timings (μs):") - for (name, time) in sort(collect(results), by=x->x[2], rev=true) - @printf(" %-35s %8.1f μs\n", name, time) - end - - return results -end - -# ============================================================================ -# RUN ANALYSIS -# ============================================================================ - -# Small signal (where GPU struggles) -analyze_fused_components(10_000, 100) -compare_implementations(10_000, 100) -analyze_memory_operations(10_000, 100) - -# Medium signal -analyze_fused_components(100_000, 500) -compare_implementations(100_000, 500) - -# Large signal (where GPU wins) -analyze_fused_components(1_000_000, 1000) -compare_implementations(1_000_000, 1000) - -# Async pipeline analysis -analyze_async_pipeline(100_000, 500) - -println("\n" * "=" ^ 80) -println("ANALYSIS COMPLETE") -println("=" ^ 80) diff --git a/perf/padding_alternatives.jl b/perf/padding_alternatives.jl deleted file mode 100644 index 9fdb81e09..000000000 --- a/perf/padding_alternatives.jl +++ /dev/null @@ -1,265 +0,0 @@ -using Metal -using Statistics -using Printf - -println("=" ^ 80) -println("PADDING ALTERNATIVES: Can we eliminate/reduce memory overhead?") -println("=" ^ 80) - -signal_size = 100_000 -kernel_size = 500 -full_size = signal_size + kernel_size - 1 -nfft = Metal.MPSGraphs.nextfastfft(full_size) - -println("\nSetup: signal=$signal_size, kernel=$kernel_size, FFT size=$nfft") -println("Padding needed: signal=$(nfft - signal_size), kernel=$(nfft - kernel_size)") - -signal = MtlVector(rand(Float32, signal_size)) -kernel = MtlVector(rand(Float32, kernel_size)) - -n_iters = 30 - -# ============================================================================ -# CURRENT APPROACH: Separate copy + zero-fill -# ============================================================================ - -println("\n" * "=" ^ 60) -println("METHOD 1: Current approach (copyto! + broadcast zero)") -println("=" ^ 60) - -signal_padded = MtlVector{Float32}(undef, nfft) -kernel_padded = MtlVector{Float32}(undef, nfft) - -# Warmup -copyto!(signal_padded, 1, signal, 1, signal_size) -@view(signal_padded[(signal_size+1):nfft]) .= 0f0 -Metal.synchronize() - -times = Float64[] -for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - copyto!(signal_padded, 1, signal, 1, signal_size) - @view(signal_padded[(signal_size+1):nfft]) .= 0f0 - copyto!(kernel_padded, 1, kernel, 1, kernel_size) - @view(kernel_padded[(kernel_size+1):nfft]) .= 0f0 - Metal.synchronize() - end - push!(times, t * 1e6) -end -current_time = median(times) -@printf(" Time: %.1f μs\n", current_time) - -# ============================================================================ -# ALTERNATIVE 1: Pre-fill with zeros, then copy -# ============================================================================ - -println("\n" * "=" ^ 60) -println("METHOD 2: Pre-fill zeros, then copy (fill! + copyto!)") -println("=" ^ 60) - -times = Float64[] -for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - fill!(signal_padded, 0f0) - fill!(kernel_padded, 0f0) - copyto!(signal_padded, 1, signal, 1, signal_size) - copyto!(kernel_padded, 1, kernel, 1, kernel_size) - Metal.synchronize() - end - push!(times, t * 1e6) -end -prefill_time = median(times) -@printf(" Time: %.1f μs (%.2fx vs current)\n", prefill_time, prefill_time / current_time) - -# ============================================================================ -# ALTERNATIVE 2: Use zeros() then copy -# ============================================================================ - -println("\n" * "=" ^ 60) -println("METHOD 3: Create with zeros() then copy") -println("=" ^ 60) - -times = Float64[] -for _ in 1:n_iters - Metal.synchronize() - GC.gc(false) - t = @elapsed begin - sp = Metal.zeros(Float32, nfft) - kp = Metal.zeros(Float32, nfft) - copyto!(sp, 1, signal, 1, signal_size) - copyto!(kp, 1, kernel, 1, kernel_size) - Metal.synchronize() - end - push!(times, t * 1e6) -end -zeros_time = median(times) -@printf(" Time: %.1f μs (%.2fx vs current)\n", zeros_time, zeros_time / current_time) - -# ============================================================================ -# ALTERNATIVE 3: Single kernel pad-copy (custom kernel) -# ============================================================================ - -println("\n" * "=" ^ 60) -println("METHOD 4: Custom Metal kernel for pad+copy") -println("=" ^ 60) - -# Define a custom kernel that copies and pads in one operation -function pad_copy_kernel(dest, src, src_len) - i = thread_position_in_grid_1d() - if i <= src_len - @inbounds dest[i] = src[i] - elseif i <= length(dest) - @inbounds dest[i] = 0f0 - end - return -end - -# Warmup -@metal threads=256 groups=cld(nfft, 256) pad_copy_kernel(signal_padded, signal, signal_size) -Metal.synchronize() - -times = Float64[] -for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - @metal threads=256 groups=cld(nfft, 256) pad_copy_kernel(signal_padded, signal, signal_size) - @metal threads=256 groups=cld(nfft, 256) pad_copy_kernel(kernel_padded, kernel, kernel_size) - Metal.synchronize() - end - push!(times, t * 1e6) -end -custom_kernel_time = median(times) -@printf(" Time: %.1f μs (%.2fx vs current)\n", custom_kernel_time, custom_kernel_time / current_time) - -# ============================================================================ -# ALTERNATIVE 4: Batch operations without intermediate sync -# ============================================================================ - -println("\n" * "=" ^ 60) -println("METHOD 5: Queue all ops, single sync at end") -println("=" ^ 60) - -times = Float64[] -for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - # Queue all operations without waiting - copyto!(signal_padded, 1, signal, 1, signal_size) - copyto!(kernel_padded, 1, kernel, 1, kernel_size) - @view(signal_padded[(signal_size+1):nfft]) .= 0f0 - @view(kernel_padded[(kernel_size+1):nfft]) .= 0f0 - # Single sync at end - Metal.synchronize() - end - push!(times, t * 1e6) -end -batch_time = median(times) -@printf(" Time: %.1f μs (%.2fx vs current)\n", batch_time, batch_time / current_time) - -# ============================================================================ -# ALTERNATIVE 5: No padding at all - what if data comes pre-padded? -# ============================================================================ - -println("\n" * "=" ^ 60) -println("METHOD 6: Pre-padded data (best case scenario)") -println("=" ^ 60) - -# Simulate pre-padded data -signal_prepadded = Metal.zeros(Float32, nfft) -copyto!(signal_prepadded, 1, signal, 1, signal_size) -kernel_prepadded = Metal.zeros(Float32, nfft) -copyto!(kernel_prepadded, 1, kernel, 1, kernel_size) -Metal.synchronize() - -times = Float64[] -for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - # Just copy pre-padded data (simulating zero-cost padding) - copyto!(signal_padded, signal_prepadded) - copyto!(kernel_padded, kernel_prepadded) - Metal.synchronize() - end - push!(times, t * 1e6) -end -prepadded_time = median(times) -@printf(" Time: %.1f μs (%.2fx vs current)\n", prepadded_time, prepadded_time / current_time) - -# ============================================================================ -# ALTERNATIVE 6: Skip padding entirely - use inline padding in MPSGraph -# ============================================================================ - -println("\n" * "=" ^ 60) -println("METHOD 7: MPSGraph inline padding (no Julia-side padding)") -println("=" ^ 60) - -# With inline padding, we pass the original unpadded arrays -# The MPSGraph handles padding internally via concat operations - -using Metal.MPSGraphs: conv_fft_fused, conv_fft_inline_pad - -# conv_fft_fused (needs pre-padded buffers, so we measure padding + execution) -times_fused = Float64[] -for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - _ = conv_fft_fused(signal, kernel) - Metal.synchronize() - end - push!(times_fused, t * 1000) # ms -end -fused_time = median(times_fused) - -# conv_fft_inline_pad (no Julia-side padding needed) -times_inline = Float64[] -for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - _ = conv_fft_inline_pad(signal, kernel) - Metal.synchronize() - end - push!(times_inline, t * 1000) # ms -end -inline_time = median(times_inline) - -@printf(" conv_fft_fused: %.3f ms (includes padding)\n", fused_time) -@printf(" conv_fft_inline_pad: %.3f ms (no Julia-side padding)\n", inline_time) -@printf(" Difference: %.3f ms saved\n", fused_time - inline_time) - -# ============================================================================ -# SUMMARY -# ============================================================================ - -println("\n" * "=" ^ 80) -println("SUMMARY: Padding time comparison") -println("=" ^ 80) - -results = [ - ("Current (copyto! + broadcast)", current_time), - ("Pre-fill zeros + copy", prefill_time), - ("Create zeros() + copy", zeros_time), - ("Custom Metal kernel", custom_kernel_time), - ("Batch ops (single sync)", batch_time), - ("Pre-padded data", prepadded_time), -] - -sort!(results, by=x->x[2]) - -println("\nRanked by speed:") -for (i, (name, time)) in enumerate(results) - speedup = current_time / time - @printf("%d. %-30s %7.1f μs (%5.2fx vs current)\n", i, name, time, speedup) -end - -println("\n" * "=" ^ 80) -println("KEY FINDINGS") -println("=" ^ 80) -println(""" -1. The padding overhead (~700 μs) comes from GPU kernel launch costs -2. Each GPU operation (copyto!, broadcast) has ~200 μs overhead -3. Custom kernels can potentially reduce this by combining operations -4. The inline-pad implementation in MPSGraph eliminates Julia-side padding -5. Pre-padded data would be fastest but requires user workflow changes -""") diff --git a/perf/simple_bottleneck.jl b/perf/simple_bottleneck.jl deleted file mode 100644 index 791a47533..000000000 --- a/perf/simple_bottleneck.jl +++ /dev/null @@ -1,362 +0,0 @@ -using Metal -using Metal.MPSGraphs: conv_fft, conv_fft_fused, conv_fft_inline_pad -using DSP -using Statistics -using Printf - -println("=" ^ 80) -println("BOTTLENECK PROFILING: Current Implementation Analysis") -println("=" ^ 80) - -# ============================================================================ -# Test 1: Implementation comparison across sizes -# ============================================================================ - -function benchmark_latency(f, args...; n_warmup=3, n_iters=20) - # Warmup - for _ in 1:n_warmup - f(args...) - end - Metal.synchronize() - - # Benchmark - times = Float64[] - for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - f(args...) - Metal.synchronize() - end - push!(times, t * 1000) # ms - end - return median(times), minimum(times), maximum(times) -end - -function benchmark_cpu(signal_cpu, kernel_cpu; n_iters=5) - times = Float64[] - for _ in 1:n_iters - t = @elapsed DSP.conv(signal_cpu, kernel_cpu) - push!(times, t * 1000) - end - return median(times) -end - -println("\n" * "=" ^ 80) -println("SECTION 1: LATENCY COMPARISON ACROSS SIZES") -println("=" ^ 80) - -configs = [ - (10_000, 100, "Small"), - (50_000, 250, "Medium-Small"), - (100_000, 500, "Medium"), - (500_000, 500, "Medium-Large"), - (1_000_000, 1000, "Large"), -] - -println("\n| Size | Signal | Kernel | conv_fft | conv_fused | conv_inline | CPU (ms) | Best GPU |") -println("|------------|----------|--------|----------|------------|-------------|----------|----------|") - -for (signal_size, kernel_size, label) in configs - signal_cpu = rand(Float32, signal_size) - kernel_cpu = rand(Float32, kernel_size) - signal = MtlVector(signal_cpu) - kernel = MtlVector(kernel_cpu) - - # Benchmark each - t_fft, _, _ = benchmark_latency(conv_fft, signal, kernel) - t_fused, _, _ = benchmark_latency(conv_fft_fused, signal, kernel) - t_inline, _, _ = benchmark_latency(conv_fft_inline_pad, signal, kernel) - t_cpu = benchmark_cpu(signal_cpu, kernel_cpu) - - # Find best GPU - times = [("conv_fft", t_fft), ("fused", t_fused), ("inline", t_inline)] - best_name, best_time = sort(times, by=x->x[2])[1] - - @printf("| %-10s | %8d | %6d | %8.3f | %10.3f | %11.3f | %8.2f | %-8s |\n", - label, signal_size, kernel_size, t_fft, t_fused, t_inline, t_cpu, best_name) -end - -# ============================================================================ -# Test 2: Deep dive into memory operations -# ============================================================================ - -println("\n" * "=" ^ 80) -println("SECTION 2: MEMORY OPERATION OVERHEAD ANALYSIS") -println("=" ^ 80) - -function analyze_memory_overhead(signal_size, kernel_size; n_iters=30) - println("\n--- Signal=$signal_size, Kernel=$kernel_size ---") - - signal = MtlVector(rand(Float32, signal_size)) - kernel = MtlVector(rand(Float32, kernel_size)) - - # Compute FFT size - full_size = signal_size + kernel_size - 1 - nfft = Metal.MPSGraphs.nextfastfft(full_size) - output_size = full_size - - println("FFT size: $nfft, Padding needed: $(nfft - signal_size) (signal), $(nfft - kernel_size) (kernel)") - - # Pre-allocate - signal_padded = MtlVector{Float32}(undef, nfft) - kernel_padded = MtlVector{Float32}(undef, nfft) - - # Test 1: copyto! queued (no sync) - times_copyto_queue = Float64[] - for _ in 1:n_iters - Metal.synchronize() - t = @elapsed copyto!(signal_padded, 1, signal, 1, signal_size) - push!(times_copyto_queue, t * 1e6) - end - med_copyto_queue = median(times_copyto_queue) - - # Test 2: copyto! with sync - times_copyto_sync = Float64[] - for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - copyto!(signal_padded, 1, signal, 1, signal_size) - Metal.synchronize() - end - push!(times_copyto_sync, t * 1e6) - end - med_copyto_sync = median(times_copyto_sync) - - # Test 3: broadcast zero-fill queued - times_zero_queue = Float64[] - for _ in 1:n_iters - Metal.synchronize() - t = @elapsed (@view(signal_padded[(signal_size+1):nfft]) .= 0f0) - push!(times_zero_queue, t * 1e6) - end - med_zero_queue = median(times_zero_queue) - - # Test 4: broadcast zero-fill with sync - times_zero_sync = Float64[] - for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - @view(signal_padded[(signal_size+1):nfft]) .= 0f0 - Metal.synchronize() - end - push!(times_zero_sync, t * 1e6) - end - med_zero_sync = median(times_zero_sync) - - # Test 5: Combined copy + zero with one sync - times_combined = Float64[] - for _ in 1:n_iters - Metal.synchronize() - t = @elapsed begin - copyto!(signal_padded, 1, signal, 1, signal_size) - @view(signal_padded[(signal_size+1):nfft]) .= 0f0 - copyto!(kernel_padded, 1, kernel, 1, kernel_size) - @view(kernel_padded[(kernel_size+1):nfft]) .= 0f0 - Metal.synchronize() - end - push!(times_combined, t * 1e6) - end - med_combined = median(times_combined) - - # Test 6: MtlArray allocation - times_alloc = Float64[] - for _ in 1:n_iters - GC.gc(false) - t = @elapsed MtlVector{Float32}(undef, nfft) - push!(times_alloc, t * 1e6) - end - med_alloc = median(times_alloc) - - # Test 7: Metal.synchronize() alone (baseline) - times_sync_alone = Float64[] - for _ in 1:n_iters - Metal.synchronize() - t = @elapsed Metal.synchronize() - push!(times_sync_alone, t * 1e6) - end - med_sync_alone = median(times_sync_alone) - - @printf(" copyto! (queue only): %7.1f μs\n", med_copyto_queue) - @printf(" copyto! (with sync): %7.1f μs\n", med_copyto_sync) - @printf(" Zero-fill (queue only): %7.1f μs\n", med_zero_queue) - @printf(" Zero-fill (with sync): %7.1f μs\n", med_zero_sync) - @printf(" Both inputs padded (sync): %7.1f μs\n", med_combined) - @printf(" MtlArray allocation: %7.1f μs\n", med_alloc) - @printf(" Metal.synchronize() alone: %7.1f μs\n", med_sync_alone) - - # Calculate actual GPU work time - gpu_work_time = med_copyto_sync - med_sync_alone - @printf(" => Actual GPU copy time: %7.1f μs\n", gpu_work_time) - - return med_combined -end - -analyze_memory_overhead(10_000, 100) -analyze_memory_overhead(100_000, 500) -analyze_memory_overhead(1_000_000, 1000) - -# ============================================================================ -# Test 3: Async pipeline benefits -# ============================================================================ - -println("\n" * "=" ^ 80) -println("SECTION 3: ASYNC PIPELINE BENEFITS") -println("=" ^ 80) - -function test_async_pipeline(signal_size, kernel_size; batch_sizes=[1, 5, 10, 20, 50]) - println("\n--- Signal=$signal_size, Kernel=$kernel_size ---") - - signal = MtlVector(rand(Float32, signal_size)) - kernel = MtlVector(rand(Float32, kernel_size)) - - # Warmup - for _ in 1:3 - _ = conv_fft_fused(signal, kernel) - end - Metal.synchronize() - - println("Batch | Sync (ms/op) | Async (ms/op) | Speedup | Throughput") - println("-" ^ 60) - - for batch in batch_sizes - # Sync mode - Metal.synchronize() - t_sync = @elapsed begin - for _ in 1:batch - _ = conv_fft_fused(signal, kernel) - Metal.synchronize() - end - end - ms_sync = (t_sync / batch) * 1000 - - # Async mode - Metal.synchronize() - t_async = @elapsed begin - for _ in 1:batch - _ = conv_fft_fused(signal, kernel) - end - Metal.synchronize() - end - ms_async = (t_async / batch) * 1000 - - speedup = ms_sync / ms_async - throughput = batch / t_async - - @printf("%5d | %12.3f | %13.3f | %6.2fx | %7.1f ops/s\n", - batch, ms_sync, ms_async, speedup, throughput) - end -end - -test_async_pipeline(100_000, 500) -test_async_pipeline(1_000_000, 1000) - -# ============================================================================ -# Test 4: GPU vs CPU crossover point -# ============================================================================ - -println("\n" * "=" ^ 80) -println("SECTION 4: GPU vs CPU CROSSOVER ANALYSIS") -println("=" ^ 80) - -sizes = [1_000, 2_500, 5_000, 7_500, 10_000, 15_000, 25_000, 50_000, 100_000] -kernel_size = 100 - -println("\nKernel size fixed at $kernel_size") -println("Signal Size | GPU (ms) | CPU (ms) | GPU/CPU | Winner") -println("-" ^ 60) - -for signal_size in sizes - signal_cpu = rand(Float32, signal_size) - kernel_cpu = rand(Float32, kernel_size) - signal = MtlVector(signal_cpu) - kernel = MtlVector(kernel_cpu) - - t_gpu, _, _ = benchmark_latency(conv_fft_fused, signal, kernel; n_iters=15) - t_cpu = benchmark_cpu(signal_cpu, kernel_cpu; n_iters=5) - - ratio = t_gpu / t_cpu - winner = t_gpu < t_cpu ? "GPU" : "CPU" - - @printf("%11d | %8.3f | %8.3f | %7.2fx | %s\n", - signal_size, t_gpu, t_cpu, ratio, winner) -end - -# ============================================================================ -# Test 5: 2D and 3D analysis -# ============================================================================ - -println("\n" * "=" ^ 80) -println("SECTION 5: 2D CONVOLUTION ANALYSIS") -println("=" ^ 80) - -configs_2d = [ - ((64, 64), (5, 5)), - ((128, 128), (5, 5)), - ((256, 256), (5, 5)), - ((256, 256), (15, 15)), - ((512, 512), (15, 15)), - ((1024, 1024), (15, 15)), -] - -println("\n| Image Size | Kernel | conv_fft | conv_fused | conv_inline | CPU (ms) | Speedup |") -println("|------------|--------|----------|------------|-------------|----------|---------|") - -for (img_size, kern_size) in configs_2d - img_cpu = rand(Float32, img_size...) - kern_cpu = rand(Float32, kern_size...) - img = MtlMatrix(img_cpu) - kern = MtlMatrix(kern_cpu) - - t_fft, _, _ = benchmark_latency(x -> conv_fft(x, kern; dims=(1,2)), img; n_iters=15) - t_fused, _, _ = benchmark_latency(conv_fft_fused, img, kern; n_iters=15) - t_inline, _, _ = benchmark_latency(conv_fft_inline_pad, img, kern; n_iters=15) - - cpu_iters = prod(img_size) > 500_000 ? 3 : 5 - t_cpu = benchmark_cpu(img_cpu, kern_cpu; n_iters=cpu_iters) - - best_gpu = min(t_fft, t_fused, t_inline) - speedup = t_cpu / best_gpu - - @printf("| %4dx%-5d | %2dx%-3d | %8.3f | %10.3f | %11.3f | %8.2f | %6.2fx |\n", - img_size[1], img_size[2], kern_size[1], kern_size[2], - t_fft, t_fused, t_inline, t_cpu, speedup) -end - -println("\n" * "=" ^ 80) -println("SECTION 6: 3D CONVOLUTION ANALYSIS") -println("=" ^ 80) - -configs_3d = [ - ((32, 32, 32), (3, 3, 3)), - ((64, 64, 64), (3, 3, 3)), - ((64, 64, 64), (5, 5, 5)), - ((128, 128, 64), (5, 5, 5)), -] - -println("\n| Volume Size | Kernel | conv_fft | conv_fused | conv_inline | CPU (ms) | Speedup |") -println("|----------------|--------|----------|------------|-------------|----------|---------|") - -for (vol_size, kern_size) in configs_3d - vol_cpu = rand(Float32, vol_size...) - kern_cpu = rand(Float32, kern_size...) - vol = MtlArray(vol_cpu) - kern = MtlArray(kern_cpu) - - t_fft, _, _ = benchmark_latency(x -> conv_fft(x, kern; dims=(1,2,3)), vol; n_iters=10) - t_fused, _, _ = benchmark_latency(conv_fft_fused, vol, kern; n_iters=10) - t_inline, _, _ = benchmark_latency(conv_fft_inline_pad, vol, kern; n_iters=10) - - cpu_iters = prod(vol_size) > 500_000 ? 2 : 3 - t_cpu = benchmark_cpu(vol_cpu, kern_cpu; n_iters=cpu_iters) - - best_gpu = min(t_fft, t_fused, t_inline) - speedup = t_cpu / best_gpu - - @printf("| %3dx%3dx%-5d | %dx%dx%-1d | %8.3f | %10.3f | %11.3f | %8.2f | %6.2fx |\n", - vol_size[1], vol_size[2], vol_size[3], kern_size[1], kern_size[2], kern_size[3], - t_fft, t_fused, t_inline, t_cpu, speedup) -end - -println("\n" * "=" ^ 80) -println("ANALYSIS COMPLETE") -println("=" ^ 80) From 2ea3be9bb368e1051c1915fafe21dbd728f2e4de Mon Sep 17 00:00:00 2001 From: Kaan Kesgin Date: Wed, 27 May 2026 20:05:09 +0200 Subject: [PATCH 05/12] Trim convolution engine to the single fused-FFT path Per the measured benchmarks (perf/CONVOLUTION_PERFORMANCE.md), the fused FFT path is fastest in every regime, so the extra internal paths cost only complexity: - Remove MPS direct convolution (conv_direct) + its conv2d MPSGraph wrappers (also removes the overlap with PR #745, which owns the NN/conv2d lane) - Remove the ConvFFTPlan / plan_conv_fft repeated-convolution plan (~2x slower) - Remove the dead inline-pad fused-graph variant (no callers) conv()/imfilter now route through the single fused-FFT engine; DSP.conv/DSP.xcorr stay the public interface. convolution.jl 1916 -> 1003 lines; operations.jl drops the conv2d wrappers (keeps concatTensors). Tests: engine 32/32, DSP 5/5. --- ext/MetalDSPExt.jl | 8 +- lib/mpsgraphs/convolution.jl | 958 +--------------------------------- lib/mpsgraphs/operations.jl | 83 --- test/mpsgraphs/convolution.jl | 165 +----- 4 files changed, 31 insertions(+), 1183 deletions(-) diff --git a/ext/MetalDSPExt.jl b/ext/MetalDSPExt.jl index e6321f1d4..3479553a7 100644 --- a/ext/MetalDSPExt.jl +++ b/ext/MetalDSPExt.jl @@ -17,14 +17,14 @@ const MtlConvNumber = Union{Float32, Float16, ComplexF32, ComplexF16} DSP.conv(u::MtlArray, v::MtlArray; algorithm = :auto) Full linear convolution of two `MtlArray`s on the GPU, computed via the FFT -convolution theorem (with an MPS direct-convolution fast path for small 2-D -kernels). Convolves over all dimensions, matching `DSP.conv` semantics. The -`algorithm` keyword accepts `:auto`, `:fft`, or `:direct`. +convolution theorem. Convolves over all dimensions, matching `DSP.conv` +semantics. The `algorithm` keyword (`:auto`/`:fft`) is accepted for +compatibility; the FFT path is always used. """ function DSP.conv( u::MtlArray{T, N}, v::MtlArray{T, N}; algorithm::Symbol = :auto ) where {T <: MtlConvNumber, N} - alg = algorithm in (:fft, :direct) ? algorithm : :auto + alg = algorithm === :fft ? :fft : :auto return Metal.MPSGraphs.conv(u, v; dims = ntuple(identity, N), mode = :full, algorithm = alg) end diff --git a/lib/mpsgraphs/convolution.jl b/lib/mpsgraphs/convolution.jl index a406a2a91..eedeec4eb 100644 --- a/lib/mpsgraphs/convolution.jl +++ b/lib/mpsgraphs/convolution.jl @@ -11,9 +11,8 @@ using AbstractFFTs -export conv, conv_fft, conv_fft!, conv_fft_fused, xcorr -export plan_conv_fft, ConvFFTPlan -export get_cached_conv_plan, clear_conv_plan_cache!, clear_fused_conv_cache! +export conv, conv_fft, conv_fft!, xcorr, imfilter +export clear_fused_conv_cache! # ============================================================================ # Helper Functions @@ -108,468 +107,6 @@ function _extract_conv_result( return result[ranges...] end -# ============================================================================ -# Convolution Plan (for repeated convolutions with same sizes) -# ============================================================================ - -""" - ConvFFTPlan{T, N} - -Pre-computed FFT convolution plan for efficient repeated convolutions. - -When you need to convolve many signals with the same kernel, or perform -multiple convolutions with arrays of the same shape, creating a plan -avoids redundant allocations and FFT setup. - -# Fields (internal) -- Pre-allocated padded signal/kernel buffers -- Pre-computed kernel FFT (when kernel is provided at plan time) -- Cached FFT size and output parameters - -# Example -```julia -# Create plan for 1D convolution -signal_size = 1000 -kernel_size = 100 -plan = plan_conv_fft(signal_size, kernel_size, Float32) - -# Use plan for multiple convolutions -for signal in signals - result = plan * signal # Uses pre-allocated buffers -end - -# Or with pre-computed kernel FFT -kernel = MtlVector(randn(Float32, 100)) -plan_with_kernel = plan_conv_fft(1000, kernel) -for signal in signals - result = plan_with_kernel * signal # Even faster - kernel FFT cached -end -``` -""" -struct ConvFFTPlan{T, N, IsReal} - signal_size::NTuple{N, Int} - kernel_size::NTuple{N, Int} - output_size::NTuple{N, Int} - full_size::NTuple{N, Int} - fft_size::NTuple{N, Int} - dims::Tuple{Vararg{Int}} - mode::Symbol - # Pre-allocated buffers - signal_padded::MtlArray{T, N} - kernel_padded::MtlArray{T, N} - # Pre-computed kernel FFT (if kernel was provided) - kernel_fft::Union{Nothing, MtlArray{<:Complex, N}} -end - -# Plan cache for automatic reuse -const _CONV_PLAN_CACHE = Dict{UInt64, ConvFFTPlan}() -const _CONV_PLAN_CACHE_LOCK = ReentrantLock() -const _CONV_PLAN_CACHE_MAX_SIZE = 32 - -""" - _conv_plan_cache_key(signal_size, kernel_size, T, dims, mode) - -Generate a unique key for caching convolution plans. -""" -function _conv_plan_cache_key( - signal_size::NTuple{N, Int}, kernel_size::NTuple{N, Int}, - ::Type{T}, dims::Tuple, mode::Symbol - ) where {N, T} - return hash((signal_size, kernel_size, T, dims, mode)) -end - -""" - plan_conv_fft(signal_size::Int, kernel_size::Int, T::Type; mode=:full) - -Create an FFT convolution plan for 1D arrays of the specified sizes and element type. - -# Arguments -- `signal_size`: Length of signals to convolve -- `kernel_size`: Length of kernels to convolve -- `T`: Element type (Float32, Float16, ComplexF32, or ComplexF16) -- `mode`: Output mode (`:full`, `:same`, or `:valid`) - -# Returns -A `ConvFFTPlan` that can be used with `*` or `mul!` for efficient convolution. - -# Example -```julia -plan = plan_conv_fft(1000, 100, Float32) -signal = MtlVector(randn(Float32, 1000)) -kernel = MtlVector(randn(Float32, 100)) -result = conv_fft(plan, signal, kernel) # Uses pre-allocated buffers -``` -""" -function plan_conv_fft( - signal_size::Int, kernel_size::Int, ::Type{T}; mode::Symbol = :full - ) where {T <: Union{Float32, Float16}} - return _create_conv_plan((signal_size,), (kernel_size,), T, (1,), mode, true) -end - -function plan_conv_fft( - signal_size::Int, kernel_size::Int, ::Type{Complex{T}}; mode::Symbol = :full - ) where {T <: Union{Float32, Float16}} - return _create_conv_plan((signal_size,), (kernel_size,), Complex{T}, (1,), mode, false) -end - -""" - plan_conv_fft(signal_size::NTuple{N,Int}, kernel_size::NTuple{N,Int}, T::Type; dims=1, mode=:full) - -Create an FFT convolution plan for N-dimensional arrays. -""" -function plan_conv_fft( - signal_size::NTuple{N, Int}, kernel_size::NTuple{N, Int}, ::Type{T}; - dims::Union{Int, Tuple{Vararg{Int}}} = 1, mode::Symbol = :full - ) where {N, T <: Union{Float32, Float16}} - dims_tuple = dims isa Int ? (dims,) : Tuple(dims) - return _create_conv_plan(signal_size, kernel_size, T, dims_tuple, mode, true) -end - -function plan_conv_fft( - signal_size::NTuple{N, Int}, kernel_size::NTuple{N, Int}, ::Type{Complex{T}}; - dims::Union{Int, Tuple{Vararg{Int}}} = 1, mode::Symbol = :full - ) where {N, T <: Union{Float32, Float16}} - dims_tuple = dims isa Int ? (dims,) : Tuple(dims) - return _create_conv_plan(signal_size, kernel_size, Complex{T}, dims_tuple, mode, false) -end - -""" - plan_conv_fft(signal_size, kernel::MtlArray; dims=1, mode=:full) - -Create an FFT convolution plan with a pre-computed kernel FFT. - -This is the most efficient option when convolving many signals with the same kernel. -The kernel's FFT is computed once at plan creation time. - -# Example -```julia -kernel = MtlVector(randn(Float32, 100)) -plan = plan_conv_fft(1000, kernel) - -# Each convolution now only requires one FFT (signal) instead of two -for signal in signals - result = conv_fft(plan, signal) -end -``` -""" -function plan_conv_fft( - signal_size::Int, kernel::MtlVector{T}; mode::Symbol = :full - ) where {T <: Union{Float32, Float16}} - plan = _create_conv_plan((signal_size,), (length(kernel),), T, (1,), mode, true) - return _precompute_kernel_fft!(plan, kernel) -end - -function plan_conv_fft( - signal_size::Int, kernel::MtlVector{Complex{T}}; mode::Symbol = :full - ) where {T <: Union{Float32, Float16}} - plan = _create_conv_plan((signal_size,), (length(kernel),), Complex{T}, (1,), mode, false) - return _precompute_kernel_fft!(plan, kernel) -end - -function plan_conv_fft( - signal_size::NTuple{N, Int}, kernel::MtlArray{T, N}; - dims::Union{Int, Tuple{Vararg{Int}}} = 1, mode::Symbol = :full - ) where {N, T <: Union{Float32, Float16}} - dims_tuple = dims isa Int ? (dims,) : Tuple(dims) - plan = _create_conv_plan(signal_size, size(kernel), T, dims_tuple, mode, true) - return _precompute_kernel_fft!(plan, kernel) -end - -function plan_conv_fft( - signal_size::NTuple{N, Int}, kernel::MtlArray{Complex{T}, N}; - dims::Union{Int, Tuple{Vararg{Int}}} = 1, mode::Symbol = :full - ) where {N, T <: Union{Float32, Float16}} - dims_tuple = dims isa Int ? (dims,) : Tuple(dims) - plan = _create_conv_plan(signal_size, size(kernel), Complex{T}, dims_tuple, mode, false) - return _precompute_kernel_fft!(plan, kernel) -end - -""" -Internal function to create a convolution plan. -""" -function _create_conv_plan( - signal_size::NTuple{N, Int}, kernel_size::NTuple{N, Int}, - ::Type{T}, dims::Tuple{Vararg{Int}}, mode::Symbol, is_real::Bool - ) where {N, T} - # Validate dimensions - for d in dims - 1 <= d <= N || - throw(ArgumentError("Invalid dimension $d for array with $N dimensions")) - end - - # Compute sizes - full_size = ntuple(N) do i - if i in dims - signal_size[i] + kernel_size[i] - 1 - else - signal_size[i] - end - end - - output_size = ntuple(N) do i - if i in dims - _conv_output_size(signal_size[i], kernel_size[i], mode) - else - signal_size[i] - end - end - - fft_size = ntuple(N) do i - if i in dims - nextfastfft(full_size[i]) - else - signal_size[i] - end - end - - # Allocate buffers - signal_padded = MtlArray{T}(undef, fft_size) - kernel_padded = MtlArray{T}(undef, fft_size) - - # Zero-fill once (will be overwritten in parts during convolution) - fill!(signal_padded, zero(T)) - fill!(kernel_padded, zero(T)) - - return ConvFFTPlan{T, N, is_real}( - signal_size, kernel_size, output_size, full_size, fft_size, - dims, mode, signal_padded, kernel_padded, nothing - ) -end - -""" -Internal function to pre-compute kernel FFT for a plan. -""" -function _precompute_kernel_fft!(plan::ConvFFTPlan{T, N, IsReal}, kernel::MtlArray{T, N}) where {T, N, IsReal} - # Copy kernel to padded buffer - kernel_ranges = ntuple(i -> 1:plan.kernel_size[i], N) - fill!(plan.kernel_padded, zero(T)) - plan.kernel_padded[kernel_ranges...] = kernel - - # Compute kernel FFT - if IsReal - kernel_fft = rfft(plan.kernel_padded, plan.dims) - else - kernel_fft = fft(plan.kernel_padded, plan.dims) - end - - # Store in a new plan (since structs are immutable, we create a new one) - # Note: This is a bit awkward, but avoids making the struct mutable - return ConvFFTPlan{T, N, IsReal}( - plan.signal_size, plan.kernel_size, plan.output_size, - plan.full_size, plan.fft_size, plan.dims, plan.mode, - plan.signal_padded, plan.kernel_padded, kernel_fft - ) -end - -""" - conv_fft(plan::ConvFFTPlan, signal, kernel) - -Perform convolution using a pre-computed plan. - -Uses pre-allocated buffers from the plan, avoiding allocations. -""" -function conv_fft( - plan::ConvFFTPlan{T, N, true}, signal::MtlArray{T, N}, kernel::MtlArray{T, N} - ) where {T <: Union{Float32, Float16}, N} - @assert size(signal) == plan.signal_size "Signal size $(size(signal)) doesn't match plan $(plan.signal_size)" - @assert size(kernel) == plan.kernel_size "Kernel size $(size(kernel)) doesn't match plan $(plan.kernel_size)" - - # Copy signal to padded buffer - signal_ranges = ntuple(i -> 1:plan.signal_size[i], N) - fill!(plan.signal_padded, zero(T)) - plan.signal_padded[signal_ranges...] = signal - - # FFT signal - S = rfft(plan.signal_padded, plan.dims) - - # Kernel FFT (use cached if available, otherwise compute) - K = if plan.kernel_fft !== nothing - plan.kernel_fft - else - kernel_ranges = ntuple(i -> 1:plan.kernel_size[i], N) - fill!(plan.kernel_padded, zero(T)) - plan.kernel_padded[kernel_ranges...] = kernel - rfft(plan.kernel_padded, plan.dims) - end - - # Multiply and inverse FFT - Y = S .* K - first_dim = minimum(plan.dims) - y = irfft(Y, plan.fft_size[first_dim], plan.dims) - - # Extract result - if N == 1 - return _extract_conv_result(y, plan.output_size[1], plan.full_size[1], plan.mode) - else - return _extract_conv_result(y, plan.output_size, plan.full_size, plan.mode, plan.dims) - end -end - -# Complex version -function conv_fft( - plan::ConvFFTPlan{Complex{T}, N, false}, signal::MtlArray{Complex{T}, N}, - kernel::MtlArray{Complex{T}, N} - ) where {T <: Union{Float32, Float16}, N} - @assert size(signal) == plan.signal_size "Signal size $(size(signal)) doesn't match plan $(plan.signal_size)" - @assert size(kernel) == plan.kernel_size "Kernel size $(size(kernel)) doesn't match plan $(plan.kernel_size)" - - # Copy signal to padded buffer - signal_ranges = ntuple(i -> 1:plan.signal_size[i], N) - fill!(plan.signal_padded, zero(Complex{T})) - plan.signal_padded[signal_ranges...] = signal - - # FFT signal - S = fft(plan.signal_padded, plan.dims) - - # Kernel FFT (use cached if available) - K = if plan.kernel_fft !== nothing - plan.kernel_fft - else - kernel_ranges = ntuple(i -> 1:plan.kernel_size[i], N) - fill!(plan.kernel_padded, zero(Complex{T})) - plan.kernel_padded[kernel_ranges...] = kernel - fft(plan.kernel_padded, plan.dims) - end - - # Multiply and inverse FFT - Y = S .* K - y = ifft(Y, plan.dims) - - # Extract result - if N == 1 - return _extract_conv_result(y, plan.output_size[1], plan.full_size[1], plan.mode) - else - return _extract_conv_result(y, plan.output_size, plan.full_size, plan.mode, plan.dims) - end -end - -""" - conv_fft(plan::ConvFFTPlan, signal) - -Perform convolution using a plan with pre-computed kernel FFT. - -This is the fastest option - only one FFT (for the signal) is needed. -""" -function conv_fft( - plan::ConvFFTPlan{T, N, true}, signal::MtlArray{T, N} - ) where {T <: Union{Float32, Float16}, N} - plan.kernel_fft === nothing && - throw(ArgumentError("Plan has no pre-computed kernel FFT. Use conv_fft(plan, signal, kernel) or create plan with kernel.")) - - @assert size(signal) == plan.signal_size "Signal size $(size(signal)) doesn't match plan $(plan.signal_size)" - - # Copy signal to padded buffer - signal_ranges = ntuple(i -> 1:plan.signal_size[i], N) - fill!(plan.signal_padded, zero(T)) - plan.signal_padded[signal_ranges...] = signal - - # FFT signal and multiply with cached kernel FFT - S = rfft(plan.signal_padded, plan.dims) - Y = S .* plan.kernel_fft - - # Inverse FFT - first_dim = minimum(plan.dims) - y = irfft(Y, plan.fft_size[first_dim], plan.dims) - - # Extract result - if N == 1 - return _extract_conv_result(y, plan.output_size[1], plan.full_size[1], plan.mode) - else - return _extract_conv_result(y, plan.output_size, plan.full_size, plan.mode, plan.dims) - end -end - -# Complex version with pre-computed kernel -function conv_fft( - plan::ConvFFTPlan{Complex{T}, N, false}, signal::MtlArray{Complex{T}, N} - ) where {T <: Union{Float32, Float16}, N} - plan.kernel_fft === nothing && - throw(ArgumentError("Plan has no pre-computed kernel FFT. Use conv_fft(plan, signal, kernel) or create plan with kernel.")) - - @assert size(signal) == plan.signal_size "Signal size $(size(signal)) doesn't match plan $(plan.signal_size)" - - # Copy signal to padded buffer - signal_ranges = ntuple(i -> 1:plan.signal_size[i], N) - fill!(plan.signal_padded, zero(Complex{T})) - plan.signal_padded[signal_ranges...] = signal - - # FFT signal and multiply with cached kernel FFT - S = fft(plan.signal_padded, plan.dims) - Y = S .* plan.kernel_fft - - # Inverse FFT - y = ifft(Y, plan.dims) - - # Extract result - if N == 1 - return _extract_conv_result(y, plan.output_size[1], plan.full_size[1], plan.mode) - else - return _extract_conv_result(y, plan.output_size, plan.full_size, plan.mode, plan.dims) - end -end - -""" - get_cached_conv_plan(signal_size, kernel_size, T; dims=1, mode=:full) - -Get or create a cached convolution plan for the given parameters. - -Plans are cached globally and reused for repeated convolutions with the same sizes. -This is useful when array sizes are known in advance and convolutions are repeated. - -# Thread Safety -Plan cache access is thread-safe using a lock. - -# Cache Size -The cache holds up to $_CONV_PLAN_CACHE_MAX_SIZE plans. When full, the oldest plan -is evicted (FIFO). -""" -function get_cached_conv_plan( - signal_size::NTuple{N, Int}, kernel_size::NTuple{N, Int}, ::Type{T}; - dims::Union{Int, Tuple{Vararg{Int}}} = 1, mode::Symbol = :full - ) where {N, T} - dims_tuple = dims isa Int ? (dims,) : Tuple(dims) - key = _conv_plan_cache_key(signal_size, kernel_size, T, dims_tuple, mode) - - lock(_CONV_PLAN_CACHE_LOCK) do - if haskey(_CONV_PLAN_CACHE, key) - return _CONV_PLAN_CACHE[key] - else - # Create new plan - is_real = T <: Real - plan = _create_conv_plan(signal_size, kernel_size, T, dims_tuple, mode, is_real) - - # Evict oldest if cache is full - if length(_CONV_PLAN_CACHE) >= _CONV_PLAN_CACHE_MAX_SIZE - # Simple FIFO eviction - delete first key - first_key = first(keys(_CONV_PLAN_CACHE)) - delete!(_CONV_PLAN_CACHE, first_key) - end - - _CONV_PLAN_CACHE[key] = plan - return plan - end - end -end - -# 1D convenience -function get_cached_conv_plan( - signal_size::Int, kernel_size::Int, ::Type{T}; mode::Symbol = :full - ) where {T} - return get_cached_conv_plan((signal_size,), (kernel_size,), T; dims = 1, mode = mode) -end - -""" - clear_conv_plan_cache!() - -Clear the global convolution plan cache, freeing GPU memory. -""" -function clear_conv_plan_cache!() - lock(_CONV_PLAN_CACHE_LOCK) do - empty!(_CONV_PLAN_CACHE) - end - return nothing -end - # ============================================================================ # Fused MPSGraph Convolution (Single Graph Execution) # ============================================================================ @@ -706,196 +243,6 @@ function _fast_pad_copy_contiguous!(dest::MtlArray{T, N}, src::MtlArray{T, N}) w return dest end -# ============================================================================ -# In-Place Padding Graphs (Zero-copy padding inside MPSGraph) -# ============================================================================ - -# Cache key for graphs with inline padding -struct InlinePadConvGraphKey - signal_sizes::Tuple{Vararg{Int}} # Original signal shape - kernel_sizes::Tuple{Vararg{Int}} # Original kernel shape - fft_sizes::Tuple{Vararg{Int}} # Padded FFT size - output_sizes::Tuple{Vararg{Int}} # Output shape - eltype::DataType -end - -# Cached graph with inline padding -struct CachedInlinePadConvGraph - graph::MPSGraph - signal_placeholder::MPSGraphTensor - kernel_placeholder::MPSGraphTensor - result::MPSGraphTensor -end - -const _inline_pad_conv_cache = Dict{InlinePadConvGraphKey, CachedInlinePadConvGraph}() -const _inline_pad_conv_cache_lock = ReentrantLock() - -""" -Helper to create a zeros tensor of a given shape inside MPSGraph. -Uses constantWithScalar + broadcastTensor. -""" -function _create_zeros_tensor(graph::MPSGraph, shape::NTuple{N, Int}, ::Type{T}) where {N, T} - zero_scalar = constantWithScalar(graph, T(0), T) - mps_shape = MPSShape([NSNumber(Int32(s)) for s in reverse(shape)]) - return broadcastTensor(graph, zero_scalar, mps_shape, "zeros_$(join(shape, 'x'))") -end - -""" -Build a fused N-D convolution graph with inline padding. -Accepts unpadded signal and kernel, pads inside the graph using concat. -""" -function _build_inline_pad_conv_graph_nd( - signal_sizes::NTuple{N, Int}, kernel_sizes::NTuple{N, Int}, - fft_sizes::NTuple{N, Int}, ::Type{T} - ) where {N, T <: Union{Float32, Float16}} - graph = MPSGraph() - - # Placeholders for UNPADDED inputs - signal_ph = placeholderTensor(graph, signal_sizes, T) - kernel_ph = placeholderTensor(graph, kernel_sizes, T) - - # Create zeros tensors for padding (for each dimension) - # Pad signal: concat(signal, zeros) along each dimension - signal_padded = signal_ph - for dim in 1:N - pad_size = fft_sizes[dim] - signal_sizes[dim] - if pad_size > 0 - # Create zeros for this dimension - # Shape: same as current signal_padded but with pad_size in this dimension - current_shape = ntuple(N) do i - if i == dim - pad_size - elseif i < dim - fft_sizes[i] # Already padded dimensions - else - signal_sizes[i] # Not yet padded dimensions - end - end - zeros_tensor = _create_zeros_tensor(graph, current_shape, T) - # Concat along this dimension (Metal uses reversed axis order) - metal_dim = N - dim # Convert Julia dim to Metal axis - tensors = NSArray([signal_padded, zeros_tensor]) - signal_padded = concatTensors(graph, tensors, metal_dim, "signal_pad_dim$(dim)") - end - end - - # Pad kernel similarly - kernel_padded = kernel_ph - for dim in 1:N - pad_size = fft_sizes[dim] - kernel_sizes[dim] - if pad_size > 0 - current_shape = ntuple(N) do i - if i == dim - pad_size - elseif i < dim - fft_sizes[i] - else - kernel_sizes[i] - end - end - zeros_tensor = _create_zeros_tensor(graph, current_shape, T) - metal_dim = N - dim - tensors = NSArray([kernel_padded, zeros_tensor]) - kernel_padded = concatTensors(graph, tensors, metal_dim, "kernel_pad_dim$(dim)") - end - end - - # Now proceed with FFT convolution on padded tensors - fft_desc_fwd = MPSGraphFFTDescriptor(inverse = false) - axes = NSArray([NSNumber(Int32(i)) for i in (N-1):-1:0]) - - signal_fft = realToHermiteanFFTWithTensor(graph, signal_padded, axes, fft_desc_fwd, "signal_rfft") - kernel_fft = realToHermiteanFFTWithTensor(graph, kernel_padded, axes, fft_desc_fwd, "kernel_rfft") - - product = multiplicationWithPrimaryTensor(graph, signal_fft, kernel_fft, "freq_multiply") - - total_size = prod(fft_sizes) - fft_desc_inv = MPSGraphFFTDescriptor(inverse = true) - fft_desc_inv.roundToOddHermitean = isodd(fft_sizes[1]) - result_unscaled = HermiteanToRealFFTWithTensor(graph, product, axes, fft_desc_inv, "irfft") - - scale_factor = constantWithScalar(graph, T(1) / T(total_size), T) - result_scaled = multiplicationWithPrimaryTensor(graph, result_unscaled, scale_factor, "scale") - - return CachedInlinePadConvGraph(graph, signal_ph, kernel_ph, result_scaled) -end - -""" -Get or create a cached inline-padding convolution graph. -""" -function _get_cached_inline_pad_conv_graph(key::InlinePadConvGraphKey) - cached = get(_inline_pad_conv_cache, key, nothing) - if cached !== nothing - return cached - end - lock(_inline_pad_conv_cache_lock) do - cached = get(_inline_pad_conv_cache, key, nothing) - if cached !== nothing - return cached - end - cached = _build_inline_pad_conv_graph_nd( - key.signal_sizes, key.kernel_sizes, key.fft_sizes, key.eltype) - _inline_pad_conv_cache[key] = cached - return cached - end -end - -""" - conv_fft_inline_pad(signal::MtlArray{T,N}, kernel::MtlArray{T,N}; mode=:full) - -Compute N-D convolution using a fused MPSGraph with inline padding. -Eliminates Julia-side copy/zero kernel launches by moving padding into the graph. -""" -function conv_fft_inline_pad( - signal::MtlArray{T, N}, kernel::MtlArray{T, N}; mode::Symbol = :full - ) where {T <: Union{Float32, Float16}, N} - signal_sizes = size(signal) - kernel_sizes = size(kernel) - - # Compute sizes - full_sizes = ntuple(i -> signal_sizes[i] + kernel_sizes[i] - 1, N) - output_sizes = ntuple(i -> _conv_output_size(signal_sizes[i], kernel_sizes[i], mode), N) - fft_sizes = ntuple(i -> nextfastfft(full_sizes[i]), N) - - # Get cached graph with inline padding - key = InlinePadConvGraphKey(signal_sizes, kernel_sizes, fft_sizes, output_sizes, T) - cached = _get_cached_inline_pad_conv_graph(key) - - # Get output buffer from pool - buffers = _get_cached_buffers(fft_sizes, T) - output = buffers.output - - # Execute graph - no Julia-side padding needed! - @autoreleasepool begin - feeds = Dict{MPSGraphTensor, MPSGraphTensorData}( - cached.signal_placeholder => MPSGraphTensorData(signal), - cached.kernel_placeholder => MPSGraphTensorData(kernel) - ) - - resultdict = Dict{MPSGraphTensor, MPSGraphTensorData}( - cached.result => MPSGraphTensorData(output) - ) - - cmdbuf = MPSCommandBuffer(Metal.global_queue(current_device())) - encode!(cmdbuf, cached.graph, NSDictionary(feeds), NSDictionary(resultdict), nil, default_exec_desc()) - commit!(cmdbuf) - wait_completed(cmdbuf) - - return _extract_conv_result_nd(output, output_sizes, full_sizes, mode) - end -end - -""" - clear_inline_pad_conv_cache!() - -Clear the inline padding convolution graph cache. -""" -function clear_inline_pad_conv_cache!() - lock(_inline_pad_conv_cache_lock) do - empty!(_inline_pad_conv_cache) - end - return nothing -end """ Build a fused N-D convolution graph: rfft(signal) * rfft(kernel) → irfft → scale @@ -1483,194 +830,14 @@ function conv_fft!( return output end -# ============================================================================ -# MPS Direct Convolution (for small kernels) -# ============================================================================ -# -# Uses MPSGraph's convolution2D operation for direct (non-FFT) convolution. -# This is optimized for small kernels (3×3, 5×5, 7×7) where FFT overhead dominates. -# -# Note: MPSGraph convolution expects 4D tensors in NHWC or NCHW format: -# - N = batch size -# - H = height -# - W = width -# - C = channels -# -# For signal processing, we treat 2D arrays as single-channel images with batch size 1. - -export conv_direct, imfilter - -""" - conv_direct(image::MtlMatrix, kernel::MtlMatrix; mode=:same, padding=:zeros) - -Compute 2D convolution using MPS direct convolution (optimized for small kernels). - -This function is optimized for small kernels (3×3, 5×5, 7×7) where it outperforms -FFT-based convolution. For large kernels, use `conv_fft` instead. - -# Arguments -- `image`: 2D input image (H×W) -- `kernel`: 2D convolution kernel (Kh×Kw) -- `mode`: Output size mode - - `:same` (default): Output has same size as input - - `:valid`: Only fully overlapping region - - `:full`: Full convolution output (not natively supported, falls back to FFT) -- `padding`: Padding type for `:same` mode - - `:zeros` (default): Zero padding - -# Returns -MtlMatrix with the convolution result. - -# Example -```julia -image = MtlMatrix(randn(Float32, 256, 256)) -kernel = MtlMatrix(Float32[ - 1 0 -1 - 2 0 -2 - 1 0 -1 -] ./ 8) # Sobel edge detector -edges = conv_direct(image, kernel) -``` - -# Notes -- For 3×3 kernels on 256×256 images, expect ~10-50x speedup over FFT -- Kernel is flipped internally to match mathematical convolution definition -- Currently supports Float32 and Float16 only -""" -function conv_direct( - image::MtlMatrix{T}, kernel::MtlMatrix{T}; - mode::Symbol = :same, padding::Symbol = :zeros - ) where {T <: Union{Float32, Float16}} - # For :full mode, fall back to FFT - if mode == :full - return conv_fft(image, kernel; dims = (1, 2), mode = :full) - end - - H, W = size(image) - Kh, Kw = size(kernel) - - # Validate kernel size (MPS works best with odd-sized kernels) - if Kh % 2 == 0 || Kw % 2 == 0 - @warn "Even-sized kernels may have unexpected centering. Odd sizes (3×3, 5×5, 7×7) recommended." maxlog = 1 - end - - # Flip kernel for mathematical convolution (MPS does correlation by default) - kernel_flipped = kernel[end:-1:1, end:-1:1] - - # Compute padding for :same mode - # Note: MPSGraph padding is (top, bottom) for Y, (left, right) for X - # In NHWC layout: H is the 2nd dim (padTop/padBottom), W is the 3rd dim (padLeft/padRight) - if mode == :same - # Symmetric padding to maintain size - pad_top = (Kh - 1) ÷ 2 - pad_bottom = Kh - 1 - pad_top - pad_left = (Kw - 1) ÷ 2 - pad_right = Kw - 1 - pad_left - elseif mode == :valid - pad_top = pad_bottom = pad_left = pad_right = 0 - else - throw(ArgumentError("Unknown mode: $mode. Use :same, :valid, or :full")) - end - - # Convert 2D arrays to 4D tensors for MPSGraph convolution - # Due to shape reversal in placeholderTensor (Julia shape is reversed for MPSGraph), - # we need to create 4D arrays where the Julia dimensions map correctly after reversal. - # - # We'll use NHWC layout with shapes that account for the reversal: - # - Julia shape (a, b, c, d) → MPSGraph shape (d, c, b, a) - # - # For NHWC image (N=1, H, W, C=1), MPSGraph expects shape (1, H, W, 1) - # So Julia must have shape (1, W, H, 1) which after reversal gives MPSGraph (1, H, W, 1) - # - # For HWIO kernel (Kh, Kw, Cin=1, Cout=1), MPSGraph expects shape (Kh, Kw, 1, 1) - # So Julia must have shape (1, 1, Kw, Kh) which after reversal gives MPSGraph (Kh, Kw, 1, 1) - - # Transpose the image (H, W) → (W, H) so that after reshape and reversal it matches - # Then add batch and channel dimensions - image_transposed = permutedims(image, (2, 1)) # (W, H) - image_4d = reshape(image_transposed, 1, W, H, 1) # Julia: (1, W, H, 1) → MPSGraph: (1, H, W, 1) - - # For kernel: transpose (Kh, Kw) → (Kw, Kh) then reshape - kernel_transposed = permutedims(kernel_flipped, (2, 1)) # (Kw, Kh) - kernel_4d = reshape(kernel_transposed, 1, 1, Kw, Kh) # Julia: (1, 1, Kw, Kh) → MPSGraph: (Kh, Kw, 1, 1) - - # Output size - if mode == :same - out_h, out_w = H, W - else # :valid - out_h = H - Kh + 1 - out_w = W - Kw + 1 - end - - # Create output array with reversed dimensions - # MPSGraph will produce (1, out_h, out_w, 1), which we specify as Julia (1, out_w, out_h, 1) - output = MtlArray{T}(undef, 1, out_w, out_h, 1) - - # Build and execute MPSGraph - @autoreleasepool begin - _conv2d_mpsgraph!(output, image_4d, kernel_4d, pad_top, pad_bottom, pad_left, pad_right) - end - - # Extract 2D result and transpose back to (H, W) - result_transposed = reshape(output, out_w, out_h) # (out_w, out_h) - return permutedims(result_transposed, (2, 1)) # (out_h, out_w) = (H, W) -end - -""" -Internal function to execute MPSGraph 2D convolution. -""" -function _conv2d_mpsgraph!( - output::MtlArray{T, 4}, image::MtlArray{T, 4}, kernel::MtlArray{T, 4}, - pad_top::Int, pad_bottom::Int, pad_left::Int, pad_right::Int - ) where {T} - graph = MPSGraph() - - # Create placeholders - placeImage = placeholderTensor(graph, size(image), T) - placeKernel = placeholderTensor(graph, size(kernel), T) - - feeds = Dict{MPSGraphTensor, MPSGraphTensorData}( - placeImage => MPSGraphTensorData(image), - placeKernel => MPSGraphTensorData(kernel) - ) - - # Create convolution descriptor - descriptor = MPSGraphConvolution2DOpDescriptor(; - strideX = 1, strideY = 1, - dilationX = 1, dilationY = 1, - paddingLeft = pad_left, paddingRight = pad_right, - paddingTop = pad_top, paddingBottom = pad_bottom, - paddingStyle = MPSGraphPaddingStyleExplicit, - dataLayout = MPSGraphTensorNamedDataLayoutNHWC, - weightsLayout = MPSGraphTensorNamedDataLayoutHWIO, - groups = 1 - ) - - # Perform convolution - convResult = convolution2DWithSourceTensor(graph, placeImage, placeKernel, descriptor, "conv2d") - - # Create result dictionary - resultdict = Dict{MPSGraphTensor, MPSGraphTensorData}( - convResult => MPSGraphTensorData(output) - ) - - # Execute - cmdbuf = MPSCommandBuffer(Metal.global_queue(device())) - encode!(cmdbuf, graph, NSDictionary(feeds), NSDictionary(resultdict), nil, default_exec_desc()) - commit!(cmdbuf) - wait_completed(cmdbuf) - - return output -end """ imfilter(image::MtlMatrix, kernel::MtlMatrix) -Apply a filter kernel to an image using direct convolution. +Apply a 2-D filter `kernel` to `image` via FFT-based convolution (`:same` mode). -This is a convenience function following the ImageFiltering.jl interface. -It automatically selects between MPS direct convolution (for small kernels) -and FFT convolution (for large kernels). +A convenience wrapper following the ImageFiltering.jl interface; the kernel is +centered on each pixel. # Arguments - `image`: 2D input image @@ -1706,42 +873,27 @@ edges = sqrt.(edges_x.^2 .+ edges_y.^2) ``` # Notes -- For kernels ≤ 11×11, uses MPS direct convolution -- For larger kernels, automatically falls back to FFT convolution -- The kernel is centered on each pixel (like ImageFiltering.jl's `imfilter`) +- Equivalent to `conv(image, kernel; dims=(1, 2), mode=:same)`. +- The kernel is centered on each pixel (like ImageFiltering.jl's `imfilter`). """ -# Threshold for switching between direct and FFT convolution -# MPS direct convolution is faster for small kernels -const _DIRECT_CONV_THRESHOLD = 11 - function imfilter(image::MtlMatrix{T}, kernel::MtlMatrix{T}) where {T <: Union{Float32, Float16}} - Kh, Kw = size(kernel) - - if Kh <= _DIRECT_CONV_THRESHOLD && Kw <= _DIRECT_CONV_THRESHOLD - return conv_direct(image, kernel; mode = :same) - else - return conv_fft(image, kernel; dims = (1, 2), mode = :same) - end + return conv_fft(image, kernel; dims = (1, 2), mode = :same) end # ============================================================================ -# Unified Convolution API (with automatic algorithm selection) +# Unified Convolution API (FFT-based) # ============================================================================ # -# The unified `conv()` function automatically selects the best algorithm: -# - For 2D arrays with small kernels: MPS direct convolution (faster) -# - For 2D arrays with large kernels: FFT convolution -# - For 1D arrays: FFT convolution (no MPS direct 1D support) -# - For N-D arrays: FFT convolution along specified dimensions +# The unified `conv()` dispatches to the FFT convolution engine for 1D, 2D, and +# N-D arrays. It is the internal entry point behind the public DSP.conv / +# DSP.xcorr interface (provided by the DSP.jl extension). """ conv(signal::MtlArray, kernel::MtlArray; mode=:full, dims=nothing, algorithm=:auto) -Compute linear convolution of `signal` and `kernel` with automatic algorithm selection. - -This is the recommended entry point for convolution operations. It automatically -selects between MPS direct convolution (optimized for small kernels) and FFT-based -convolution (better for large kernels or higher dimensions). +Compute the linear convolution of `signal` and `kernel` via the FFT convolution +theorem. This is the internal engine entry point; the public interface is +`DSP.conv` (provided by the DSP.jl extension). # Arguments - `signal`: Input signal (1D, 2D, or N-D MtlArray) @@ -1753,23 +905,11 @@ convolution (better for large kernels or higher dimensions). - `dims`: Dimensions along which to convolve - `nothing` (default): All dimensions for 1D/2D, dim 1 for N-D - Integer or tuple: Specific dimension(s) -- `algorithm`: Algorithm selection - - `:auto` (default): Automatically select best algorithm - - `:fft`: Force FFT-based convolution - - `:direct`: Force MPS direct convolution (2D only, small kernels) +- `algorithm`: accepted for `DSP.conv` compatibility; the FFT path is always used # Returns MtlArray with the convolution result. -# Algorithm Selection (when `algorithm=:auto`) - -For **2D matrices** with `:same` or `:valid` mode: -- Kernels ≤ 11×11: Uses MPS direct convolution (~8x faster for 3×3) -- Larger kernels: Uses FFT convolution - -For **1D vectors**, **N-D arrays**, or `:full` mode: -- Always uses FFT convolution - # Examples ```julia @@ -1780,45 +920,23 @@ signal = MtlVector(randn(Float32, 10000)) kernel = MtlVector(Float32[0.25, 0.5, 0.25]) # Simple smoothing smoothed = conv(signal, kernel; mode=:same) -# 2D image filtering (auto-selects direct convolution) +# 2D image filtering image = MtlMatrix(randn(Float32, 512, 512)) sobel_x = MtlMatrix(Float32[-1 0 1; -2 0 2; -1 0 1] ./ 8) edges = conv(image, sobel_x; mode=:same) - -# 2D with large kernel (auto-selects FFT) -large_kernel = MtlMatrix(randn(Float32, 33, 33)) -result = conv(image, large_kernel; mode=:same) - -# Force specific algorithm -result_fft = conv(image, sobel_x; mode=:same, algorithm=:fft) -result_direct = conv(image, sobel_x; mode=:same, algorithm=:direct) ``` -# Performance Tips - -1. For repeated convolutions with same sizes, use `plan_conv_fft()` or - `get_cached_conv_plan()` for even better performance with FFT. - -2. For small kernels (3×3, 5×5, 7×7), direct convolution is typically - 8-50x faster than FFT. - -3. For large kernels (>15×15), FFT becomes more efficient due to O(n log n) - vs O(n×m) complexity. - # See Also -- `conv_fft`: Force FFT-based convolution -- `conv_direct`: Force MPS direct convolution (2D only) -- `imfilter`: ImageFiltering.jl-compatible API for 2D filtering +- `imfilter`: ImageFiltering.jl-style 2D filtering - `xcorr`: Cross-correlation -- `plan_conv_fft`: Pre-computed FFT plan for repeated convolutions """ function conv( signal::MtlVector{T}, kernel::MtlVector{T}; mode::Symbol = :full, dims = nothing, algorithm::Symbol = :auto ) where {T <: Union{Float32, Float16}} - # 1D always uses FFT (no MPS direct 1D support) + # FFT-based engine; :direct is not available if algorithm == :direct - throw(ArgumentError("Direct convolution not supported for 1D arrays. Use :auto or :fft.")) + throw(ArgumentError("Direct convolution is not available (FFT-only engine). Use :auto or :fft.")) end return conv_fft(signal, kernel; mode = mode) end @@ -1834,45 +952,13 @@ function conv( return conv_fft(signal, kernel; mode = mode) end -# 2D with automatic algorithm selection +# 2D convolution (FFT-based) function conv( signal::MtlMatrix{T}, kernel::MtlMatrix{T}; mode::Symbol = :full, dims = nothing, algorithm::Symbol = :auto ) where {T <: Union{Float32, Float16}} - # Determine dims for FFT (default: both dimensions for 2D) conv_dims = dims === nothing ? (1, 2) : (dims isa Int ? (dims,) : Tuple(dims)) - - # Check if we should use direct convolution - Kh, Kw = size(kernel) - use_direct = false - - if algorithm == :auto - # Auto-select: use direct for small kernels with :same or :valid mode - # Direct convolution is only supported when convolving all dimensions - if conv_dims == (1, 2) && mode != :full - use_direct = Kh <= _DIRECT_CONV_THRESHOLD && Kw <= _DIRECT_CONV_THRESHOLD - end - elseif algorithm == :direct - # User requested direct convolution - if mode == :full - @warn "Direct convolution doesn't support :full mode. Falling back to FFT." maxlog = 1 - use_direct = false - elseif conv_dims != (1, 2) - throw(ArgumentError("Direct convolution requires convolving all dimensions (dims=(1,2) or nothing).")) - else - use_direct = true - end - elseif algorithm == :fft - use_direct = false - else - throw(ArgumentError("Unknown algorithm: $algorithm. Use :auto, :fft, or :direct.")) - end - - if use_direct - return conv_direct(signal, kernel; mode = mode) - else - return conv_fft(signal, kernel; dims = conv_dims, mode = mode) - end + return conv_fft(signal, kernel; dims = conv_dims, mode = mode) end # Complex 2D diff --git a/lib/mpsgraphs/operations.jl b/lib/mpsgraphs/operations.jl index e88d0a990..f8ea0d40d 100644 --- a/lib/mpsgraphs/operations.jl +++ b/lib/mpsgraphs/operations.jl @@ -83,86 +83,3 @@ function concatTensors(graph::MPSGraph, tensors::NSArray, dimension::Int, name = name:name::id{NSString}]::id{MPSGraphTensor} MPSGraphTensor(obj) end - -""" - convolution2DWithSourceTensor(graph, source, weights, descriptor, name="conv2d") - -2D convolution operation using MPSGraph. - -# Arguments -- `graph`: MPSGraph instance -- `source`: Input tensor in NHWC or NCHW format (depending on descriptor) -- `weights`: Convolution kernel/weights in OIHW or HWIO format (depending on descriptor) -- `descriptor`: MPSGraphConvolution2DOpDescriptor configuring stride, padding, dilation, etc. -- `name`: Operation name for debugging - -# Returns -MPSGraphTensor with the convolution result. -""" -function convolution2DWithSourceTensor( - graph::MPSGraph, source::MPSGraphTensor, weights::MPSGraphTensor, - descriptor::MPSGraphConvolution2DOpDescriptor, name = "conv2d" - ) - obj = @objc [graph::id{MPSGraph} convolution2DWithSourceTensor:source::id{MPSGraphTensor} - weightsTensor:weights::id{MPSGraphTensor} - descriptor:descriptor::id{MPSGraphConvolution2DOpDescriptor} - name:name::id{NSString}]::id{MPSGraphTensor} - MPSGraphTensor(obj) -end - -""" - MPSGraphConvolution2DOpDescriptor(; - strideX=1, strideY=1, - dilationX=1, dilationY=1, - paddingLeft=0, paddingRight=0, paddingTop=0, paddingBottom=0, - paddingStyle=MPSGraphPaddingStyleExplicit, - dataLayout=MPSGraphTensorNamedDataLayoutNHWC, - weightsLayout=MPSGraphTensorNamedDataLayoutHWIO, - groups=1 - ) - -Create a 2D convolution operation descriptor. - -# Arguments -- `strideX`, `strideY`: Stride in X and Y directions -- `dilationX`, `dilationY`: Dilation rate in X and Y directions -- `paddingLeft/Right/Top/Bottom`: Explicit padding values -- `paddingStyle`: One of: - - `MPSGraphPaddingStyleExplicit` (default) - use explicit padding values - - `MPSGraphPaddingStyleTF_VALID` - no padding - - `MPSGraphPaddingStyleTF_SAME` - pad to keep output same size as input -- `dataLayout`: Input/output tensor layout (NHWC or NCHW) -- `weightsLayout`: Kernel tensor layout (HWIO or OIHW) -- `groups`: Number of groups for grouped convolution -""" -function MPSGraphConvolution2DOpDescriptor(; - strideX::Integer = 1, strideY::Integer = 1, - dilationX::Integer = 1, dilationY::Integer = 1, - paddingLeft::Integer = 0, paddingRight::Integer = 0, - paddingTop::Integer = 0, paddingBottom::Integer = 0, - paddingStyle::MPSGraphPaddingStyle = MPSGraphPaddingStyleExplicit, - dataLayout::MPSGraphTensorNamedDataLayout = MPSGraphTensorNamedDataLayoutNHWC, - weightsLayout::MPSGraphTensorNamedDataLayout = MPSGraphTensorNamedDataLayoutHWIO, - groups::Integer = 1 - ) - # Create descriptor via alloc/init - desc = @objc [MPSGraphConvolution2DOpDescriptor alloc]::id{MPSGraphConvolution2DOpDescriptor} - desc = @objc [desc::id{MPSGraphConvolution2DOpDescriptor} init]::id{MPSGraphConvolution2DOpDescriptor} - descriptor = MPSGraphConvolution2DOpDescriptor(desc) - - # Set properties - descriptor.strideInX = UInt64(strideX) - descriptor.strideInY = UInt64(strideY) - descriptor.dilationRateInX = UInt64(dilationX) - descriptor.dilationRateInY = UInt64(dilationY) - descriptor.paddingLeft = UInt64(paddingLeft) - descriptor.paddingRight = UInt64(paddingRight) - descriptor.paddingTop = UInt64(paddingTop) - descriptor.paddingBottom = UInt64(paddingBottom) - descriptor.paddingStyle = paddingStyle - descriptor.dataLayout = dataLayout - descriptor.weightsLayout = weightsLayout - descriptor.groups = UInt64(groups) - - return descriptor -end diff --git a/test/mpsgraphs/convolution.jl b/test/mpsgraphs/convolution.jl index 07d424a1e..11a9ba61c 100644 --- a/test/mpsgraphs/convolution.jl +++ b/test/mpsgraphs/convolution.jl @@ -1,16 +1,8 @@ -# Tests for convolution operations in MPSGraphs -# -# This tests: -# - FFT-based convolution (conv_fft) -# - MPS direct convolution (conv_direct) -# - Unified conv() API with auto-selection -# - Convolution plan caching - -# The convolution engine is internal to MPSGraphs (public access is via DSP.conv). -# These tests exercise the engine directly, so import its symbols from the submodule. -using Metal.MPSGraphs: conv, conv_fft, conv_fft!, conv_fft_fused, xcorr, - plan_conv_fft, ConvFFTPlan, conv_direct, get_cached_conv_plan, - clear_conv_plan_cache!, imfilter +# Tests for the FFT-based convolution engine in MPSGraphs (conv_fft, the unified +# conv(), xcorr, imfilter). The engine is internal; the public interface is +# DSP.conv / DSP.xcorr (tested in test/dsp.jl). These tests exercise the engine +# directly, so import its symbols from the submodule. +using Metal.MPSGraphs: conv, conv_fft, conv_fft!, xcorr, imfilter # Simple reference convolution for verification (CPU) function ref_conv(u::Vector{T}, v::Vector{T}) where {T} @@ -131,128 +123,6 @@ if MPS.is_supported(device()) end end - # ============================================================================ - # Convolution Plan Tests - # ============================================================================ - - @testset "Convolution Plans" begin - @testset "Basic Plan Usage" begin - signal_size = 1000 - kernel_size = 100 - - plan = plan_conv_fft(signal_size, kernel_size, Float32; mode = :full) - - signal = MtlVector(rand(Float32, signal_size)) - kernel = MtlVector(rand(Float32, kernel_size)) - - result = conv_fft(plan, signal, kernel) - expected = conv_fft(signal, kernel; mode = :full) - - @test isapprox(Array(result), Array(expected), rtol = 1.0e-4) - end - - @testset "Plan with Pre-computed Kernel" begin - signal_size = 1000 - kernel = MtlVector(rand(Float32, 100)) - - plan = plan_conv_fft(signal_size, kernel; mode = :full) - - signal1 = MtlVector(rand(Float32, signal_size)) - signal2 = MtlVector(rand(Float32, signal_size)) - - result1 = conv_fft(plan, signal1) - result2 = conv_fft(plan, signal2) - - # Verify correctness - expected1 = conv_fft(signal1, kernel; mode = :full) - expected2 = conv_fft(signal2, kernel; mode = :full) - - @test isapprox(Array(result1), Array(expected1), rtol = 1.0e-4) - @test isapprox(Array(result2), Array(expected2), rtol = 1.0e-4) - end - - @testset "Cached Plan" begin - signal_size = (64, 64) - kernel_size = (5, 5) - - # Get same plan twice - plan1 = get_cached_conv_plan(signal_size, kernel_size, Float32; dims = (1, 2)) - plan2 = get_cached_conv_plan(signal_size, kernel_size, Float32; dims = (1, 2)) - - # Should be same object (cached) - @test plan1 === plan2 - - # Clean up - clear_conv_plan_cache!() - end - end - - # ============================================================================ - # MPS Direct Convolution Tests - # ============================================================================ - - @testset "MPS Direct Convolution" begin - @testset "Basic 2D Convolution" begin - @testset for T in [Float32, Float16] - image = rand(T, 64, 64) - kernel = rand(T, 3, 3) - - d_image = MtlMatrix(image) - d_kernel = MtlMatrix(kernel) - - # Direct convolution - result_direct = Array(conv_direct(d_image, d_kernel; mode = :same)) - - # FFT convolution (for comparison) - result_fft = Array(conv_fft(d_image, d_kernel; dims = (1, 2), mode = :same)) - - @test size(result_direct) == size(image) - @test isapprox(result_direct, result_fft, rtol = rtol(T)) - end - end - - @testset "Different Kernel Sizes" begin - image = rand(Float32, 128, 128) - d_image = MtlMatrix(image) - - for ks in [3, 5, 7, 9, 11] - kernel = rand(Float32, ks, ks) - d_kernel = MtlMatrix(kernel) - - result_direct = Array(conv_direct(d_image, d_kernel; mode = :same)) - result_fft = Array(conv_fft(d_image, d_kernel; dims = (1, 2), mode = :same)) - - @test isapprox(result_direct, result_fft, rtol = 1.0e-4) - end - end - - @testset "Valid Mode" begin - image = rand(Float32, 64, 64) - kernel = rand(Float32, 5, 5) - - d_image = MtlMatrix(image) - d_kernel = MtlMatrix(kernel) - - result = Array(conv_direct(d_image, d_kernel; mode = :valid)) - @test size(result) == (60, 60) - - # Compare with FFT - result_fft = Array(conv_fft(d_image, d_kernel; dims = (1, 2), mode = :valid)) - @test isapprox(result, result_fft, rtol = 1.0e-4) - end - - @testset "Full Mode Falls Back to FFT" begin - image = rand(Float32, 64, 64) - kernel = rand(Float32, 3, 3) - - d_image = MtlMatrix(image) - d_kernel = MtlMatrix(kernel) - - # Full mode should work (falls back to FFT internally) - result = Array(conv_direct(d_image, d_kernel; mode = :full)) - @test size(result) == (66, 66) - end - end # ============================================================================ # imfilter Tests @@ -329,31 +199,6 @@ if MPS.is_supported(device()) @test isapprox(Array(result_large), Array(expected_large), rtol = 1.0e-4) end - @testset "Algorithm Forcing" begin - image = MtlMatrix(rand(Float32, 64, 64)) - kernel = MtlMatrix(rand(Float32, 3, 3)) - - result_auto = conv(image, kernel; mode = :same, algorithm = :auto) - result_fft = conv(image, kernel; mode = :same, algorithm = :fft) - result_direct = conv(image, kernel; mode = :same, algorithm = :direct) - - # All should give numerically similar results - @test isapprox(Array(result_auto), Array(result_direct), rtol = 1.0e-6) - @test isapprox(Array(result_auto), Array(result_fft), rtol = 1.0e-4) - end - - @testset "Error Handling" begin - signal_1d = MtlVector(rand(Float32, 100)) - kernel_1d = MtlVector(rand(Float32, 10)) - - # 1D doesn't support direct - @test_throws ArgumentError conv(signal_1d, kernel_1d; algorithm = :direct) - - # Invalid algorithm - image = MtlMatrix(rand(Float32, 64, 64)) - kernel = MtlMatrix(rand(Float32, 3, 3)) - @test_throws ArgumentError conv(image, kernel; algorithm = :invalid) - end @testset "Complex Arrays" begin signal = MtlVector(rand(ComplexF32, 100)) From fbe99a97b0ba2a355c2fecd902908ef4f36a2553 Mon Sep 17 00:00:00 2001 From: Kaan Kesgin Date: Wed, 27 May 2026 20:17:07 +0200 Subject: [PATCH 06/12] Drop vestigial algorithm kwarg from internal conv() The engine is FFT-only, so the internal conv() no longer takes an algorithm keyword (and no longer throws on :direct). The public DSP.conv keeps algorithm for DSP API compatibility but ignores it. conv()/imfilter are now uniformly FFT-based. --- ext/MetalDSPExt.jl | 4 ++-- lib/mpsgraphs/convolution.jl | 34 ++++++++-------------------------- test/mpsgraphs/convolution.jl | 3 --- 3 files changed, 10 insertions(+), 31 deletions(-) diff --git a/ext/MetalDSPExt.jl b/ext/MetalDSPExt.jl index 3479553a7..b37511bfc 100644 --- a/ext/MetalDSPExt.jl +++ b/ext/MetalDSPExt.jl @@ -24,8 +24,8 @@ compatibility; the FFT path is always used. function DSP.conv( u::MtlArray{T, N}, v::MtlArray{T, N}; algorithm::Symbol = :auto ) where {T <: MtlConvNumber, N} - alg = algorithm === :fft ? :fft : :auto - return Metal.MPSGraphs.conv(u, v; dims = ntuple(identity, N), mode = :full, algorithm = alg) + # `algorithm` accepted for DSP.conv compatibility; the FFT engine is always used. + return Metal.MPSGraphs.conv(u, v; dims = ntuple(identity, N), mode = :full) end """ diff --git a/lib/mpsgraphs/convolution.jl b/lib/mpsgraphs/convolution.jl index eedeec4eb..cc7bd9fc1 100644 --- a/lib/mpsgraphs/convolution.jl +++ b/lib/mpsgraphs/convolution.jl @@ -889,7 +889,7 @@ end # DSP.xcorr interface (provided by the DSP.jl extension). """ - conv(signal::MtlArray, kernel::MtlArray; mode=:full, dims=nothing, algorithm=:auto) + conv(signal::MtlArray, kernel::MtlArray; mode=:full, dims=nothing) Compute the linear convolution of `signal` and `kernel` via the FFT convolution theorem. This is the internal engine entry point; the public interface is @@ -905,7 +905,6 @@ theorem. This is the internal engine entry point; the public interface is - `dims`: Dimensions along which to convolve - `nothing` (default): All dimensions for 1D/2D, dim 1 for N-D - Integer or tuple: Specific dimension(s) -- `algorithm`: accepted for `DSP.conv` compatibility; the FFT path is always used # Returns MtlArray with the convolution result. @@ -932,30 +931,23 @@ edges = conv(image, sobel_x; mode=:same) """ function conv( signal::MtlVector{T}, kernel::MtlVector{T}; - mode::Symbol = :full, dims = nothing, algorithm::Symbol = :auto + mode::Symbol = :full, dims = nothing ) where {T <: Union{Float32, Float16}} - # FFT-based engine; :direct is not available - if algorithm == :direct - throw(ArgumentError("Direct convolution is not available (FFT-only engine). Use :auto or :fft.")) - end return conv_fft(signal, kernel; mode = mode) end # Complex 1D function conv( signal::MtlVector{Complex{T}}, kernel::MtlVector{Complex{T}}; - mode::Symbol = :full, dims = nothing, algorithm::Symbol = :auto + mode::Symbol = :full, dims = nothing ) where {T <: Union{Float32, Float16}} - if algorithm == :direct - throw(ArgumentError("Direct convolution not supported for complex 1D arrays. Use :auto or :fft.")) - end return conv_fft(signal, kernel; mode = mode) end -# 2D convolution (FFT-based) +# 2D function conv( signal::MtlMatrix{T}, kernel::MtlMatrix{T}; - mode::Symbol = :full, dims = nothing, algorithm::Symbol = :auto + mode::Symbol = :full, dims = nothing ) where {T <: Union{Float32, Float16}} conv_dims = dims === nothing ? (1, 2) : (dims isa Int ? (dims,) : Tuple(dims)) return conv_fft(signal, kernel; dims = conv_dims, mode = mode) @@ -964,11 +956,8 @@ end # Complex 2D function conv( signal::MtlMatrix{Complex{T}}, kernel::MtlMatrix{Complex{T}}; - mode::Symbol = :full, dims = nothing, algorithm::Symbol = :auto + mode::Symbol = :full, dims = nothing ) where {T <: Union{Float32, Float16}} - if algorithm == :direct - throw(ArgumentError("Direct convolution not supported for complex arrays. Use :auto or :fft.")) - end conv_dims = dims === nothing ? (1, 2) : (dims isa Int ? (dims,) : Tuple(dims)) return conv_fft(signal, kernel; dims = conv_dims, mode = mode) end @@ -976,12 +965,8 @@ end # N-D generic (N > 2) function conv( signal::MtlArray{T, N}, kernel::MtlArray{T, N}; - mode::Symbol = :full, dims = nothing, algorithm::Symbol = :auto + mode::Symbol = :full, dims = nothing ) where {T <: Union{Float32, Float16}, N} - # N-D always uses FFT - if algorithm == :direct && N > 2 - throw(ArgumentError("Direct convolution only supported for 2D arrays. Use :auto or :fft.")) - end conv_dims = dims === nothing ? 1 : (dims isa Int ? (dims,) : Tuple(dims)) return conv_fft(signal, kernel; dims = conv_dims, mode = mode) end @@ -989,11 +974,8 @@ end # Complex N-D function conv( signal::MtlArray{Complex{T}, N}, kernel::MtlArray{Complex{T}, N}; - mode::Symbol = :full, dims = nothing, algorithm::Symbol = :auto + mode::Symbol = :full, dims = nothing ) where {T <: Union{Float32, Float16}, N} - if algorithm == :direct - throw(ArgumentError("Direct convolution not supported for complex arrays. Use :auto or :fft.")) - end conv_dims = dims === nothing ? 1 : (dims isa Int ? (dims,) : Tuple(dims)) return conv_fft(signal, kernel; dims = conv_dims, mode = mode) end diff --git a/test/mpsgraphs/convolution.jl b/test/mpsgraphs/convolution.jl index 11a9ba61c..b8b8e99df 100644 --- a/test/mpsgraphs/convolution.jl +++ b/test/mpsgraphs/convolution.jl @@ -207,9 +207,6 @@ if MPS.is_supported(device()) result = conv(signal, kernel; mode = :full) expected = conv_fft(signal, kernel; mode = :full) @test isapprox(Array(result), Array(expected), rtol = 1.0e-4) - - # Complex doesn't support direct - @test_throws ArgumentError conv(signal, kernel; algorithm = :direct) end end From 5213a7e7661e0716464621428935d96464efe7e4 Mon Sep 17 00:00:00 2001 From: Kaan Kesgin Date: Wed, 27 May 2026 20:49:57 +0200 Subject: [PATCH 07/12] Add CPU-vs-GPU convolution benchmarks to perf notes DSP.conv GPU (MetalDSPExt) vs CPU (FFTW): GPU has ~0.5ms fixed overhead and wins at scale (1-D >= 1e6: 3x; 2-D >= 512x512: 1.8-3.1x); CPU is faster for small inputs. --- perf/CONVOLUTION_PERFORMANCE.md | 28 ++++++++++++++++++++++++++++ 1 file changed, 28 insertions(+) diff --git a/perf/CONVOLUTION_PERFORMANCE.md b/perf/CONVOLUTION_PERFORMANCE.md index e18f28800..135ceb62e 100644 --- a/perf/CONVOLUTION_PERFORMANCE.md +++ b/perf/CONVOLUTION_PERFORMANCE.md @@ -6,6 +6,34 @@ Benchmarks backing the design of the GPU convolution support (the `DSP.conv` / 5 runs × 30 calls (ms/call). These numbers drive the "single-implementation" decision and are intended as source material for the PR description. +## CPU vs GPU (`DSP.conv`) + +GPU `DSP.conv` (via `MetalDSPExt`) vs CPU `DSP.conv` (FFTW-backed), `Float32`. +GPU times are compute-only (`Metal.synchronize()`, data already on device). The +GPU has a ~0.4–0.8 ms fixed dispatch/sync overhead, so it wins only at scale. + +1-D (kernel 127): + +| signal | CPU (ms) | GPU (ms) | speedup | +|-------:|---------:|---------:|--------:| +| 10³ | 0.037 | 0.412 | 0.1× | +| 10⁴ | 0.061 | 0.422 | 0.1× | +| 10⁵ | 0.288 | 0.491 | 0.6× | +| 10⁶ | 2.469 | 0.815 | **3.0×** | + +2-D (kernel 15×15): + +| image | CPU (ms) | GPU (ms) | speedup | +|------:|---------:|---------:|--------:| +| 128² | 0.128 | 0.598 | 0.2× | +| 256² | 0.427 | 0.592 | 0.7× | +| 512² | 1.457 | 0.830 | **1.8×** | +| 1024² | 5.135 | 1.633 | **3.1×** | + +Guidance: the GPU path pays off for large signals/images (1-D ≳ 10⁶, 2-D ≳ 512²) +or when data already lives on the GPU; for small inputs CPU FFTW is faster. A +one-off `CPU→GPU→CPU` round trip shifts the crossover further right. + ## Single 2-D image, `mode = :same` | image | kernel | fused FFT | MPS-direct | fused speedup | From 6917475601314b02adfa6e8da9885b7b0caa500c Mon Sep 17 00:00:00 2001 From: Kaan Kesgin Date: Thu, 28 May 2026 09:08:31 +0200 Subject: [PATCH 08/12] Bound fused-conv caches; clear errors for unsupported inputs (a) _fused_conv_graph_cache and _fused_conv_buffer_pool use a bounded FIFO (cap 32), so they no longer grow without limit across many distinct sizes. (b) DSP.conv/DSP.xcorr raise a clear ArgumentError for unsupported element types (integers etc.) instead of a silent slow CPU fallback, and DSP.xcorr now accepts + validates padmode/scaling. (Float64 is already rejected at MtlArray construction by Metal.) Tests: engine 31/31, DSP 10/10. --- ext/MetalDSPExt.jl | 26 ++++++++++++++++++++------ lib/mpsgraphs/convolution.jl | 22 ++++++++++++++++++++-- test/dsp.jl | 19 +++++++++++++++++++ 3 files changed, 59 insertions(+), 8 deletions(-) diff --git a/ext/MetalDSPExt.jl b/ext/MetalDSPExt.jl index b37511bfc..760e8add7 100644 --- a/ext/MetalDSPExt.jl +++ b/ext/MetalDSPExt.jl @@ -13,6 +13,16 @@ import DSP # Element types the MPSGraph FFT/convolution engine supports. const MtlConvNumber = Union{Float32, Float16, ComplexF32, ComplexF16} +# Clear error for unsupported element types (e.g. Float64). The MtlArray methods +# below intercept DSP.conv/xcorr before DSP's generic CPU path, which would hit +# disallowed scalar indexing on the device. +@noinline function _conv_unsupported(u, v) + throw(ArgumentError( + "Metal convolution supports Float32/Float16 and their Complex types; " * + "got $(eltype(u)) and $(eltype(v)). Convert first, e.g. with `Float32.(x)`." + )) +end + """ DSP.conv(u::MtlArray, v::MtlArray; algorithm = :auto) @@ -23,21 +33,25 @@ compatibility; the FFT path is always used. """ function DSP.conv( u::MtlArray{T, N}, v::MtlArray{T, N}; algorithm::Symbol = :auto - ) where {T <: MtlConvNumber, N} + ) where {T <: Number, N} + T <: MtlConvNumber || _conv_unsupported(u, v) # `algorithm` accepted for DSP.conv compatibility; the FFT engine is always used. return Metal.MPSGraphs.conv(u, v; dims = ntuple(identity, N), mode = :full) end """ - DSP.xcorr(u::MtlVector, v::MtlVector; padmode = :none) + DSP.xcorr(u::MtlVector, v::MtlVector; padmode = :none, scaling = :none) Cross-correlation of two GPU vectors, conjugating `v` (the DSP/MATLAB -convention). Only `padmode = :none` (the full correlation) is supported. +convention). Only `padmode = :none` (the full correlation) and `scaling = :none` +are supported. """ function DSP.xcorr( - u::MtlVector{T}, v::MtlVector{T}; padmode::Symbol = :none - ) where {T <: MtlConvNumber} - padmode === :none || throw(ArgumentError("MetalDSPExt only supports padmode = :none")) + u::MtlVector{T}, v::MtlVector{T}; padmode::Symbol = :none, scaling::Symbol = :none + ) where {T <: Number} + T <: MtlConvNumber || _conv_unsupported(u, v) + padmode === :none || throw(ArgumentError("MetalDSPExt supports only padmode = :none")) + scaling === :none || throw(ArgumentError("MetalDSPExt supports only scaling = :none")) return Metal.MPSGraphs.xcorr(u, v; mode = :full) end diff --git a/lib/mpsgraphs/convolution.jl b/lib/mpsgraphs/convolution.jl index cc7bd9fc1..668bd07b0 100644 --- a/lib/mpsgraphs/convolution.jl +++ b/lib/mpsgraphs/convolution.jl @@ -127,9 +127,22 @@ struct CachedFusedConvGraph result::MPSGraphTensor end -# Thread-safe cache for fused convolution graphs +# Thread-safe cache for fused convolution graphs (bounded FIFO; see _record_and_evict!) +const _FUSED_CONV_CACHE_MAX_ENTRIES = 32 const _fused_conv_graph_cache = Dict{FusedConvGraphKey, CachedFusedConvGraph}() const _fused_conv_graph_cache_lock = ReentrantLock() +const _fused_conv_graph_cache_order = FusedConvGraphKey[] + +# Drop the oldest cached entry once a cache exceeds the size cap (call under the +# cache's lock). Keeps the graph cache and buffer pool from growing unboundedly +# when many distinct convolution sizes are used. +function _record_and_evict!(cache::AbstractDict, order::AbstractVector, key) + push!(order, key) + if length(order) > _FUSED_CONV_CACHE_MAX_ENTRIES + delete!(cache, popfirst!(order)) + end + return nothing +end # ============================================================================ # Buffer Pooling for Fused Convolution @@ -148,9 +161,10 @@ mutable struct CachedFusedConvBuffers{T, N} output::MtlArray{T, N} end -# Thread-safe buffer pool +# Thread-safe buffer pool (bounded FIFO) const _fused_conv_buffer_pool = Dict{BufferPoolKey, CachedFusedConvBuffers}() const _fused_conv_buffer_pool_lock = ReentrantLock() +const _fused_conv_buffer_pool_order = BufferPoolKey[] """ Get or create cached buffers for fused convolution. @@ -173,6 +187,7 @@ function _get_cached_buffers(fft_sizes::NTuple{N, Int}, ::Type{T}) where {N, T} output = MtlArray{T, N}(undef, fft_sizes) cached = CachedFusedConvBuffers{T, N}(signal_padded, kernel_padded, output) _fused_conv_buffer_pool[key] = cached + _record_and_evict!(_fused_conv_buffer_pool, _fused_conv_buffer_pool_order, key) return cached end end @@ -185,6 +200,7 @@ Clear the fused convolution buffer pool, freeing GPU memory. function clear_fused_conv_buffer_pool!() lock(_fused_conv_buffer_pool_lock) do empty!(_fused_conv_buffer_pool) + empty!(_fused_conv_buffer_pool_order) end return nothing end @@ -301,6 +317,7 @@ function _get_cached_fused_conv_graph(key::FusedConvGraphKey) # Build N-D fused graph cached = _build_fused_conv_graph_nd(key.signal_fft_size, key.eltype) _fused_conv_graph_cache[key] = cached + _record_and_evict!(_fused_conv_graph_cache, _fused_conv_graph_cache_order, key) return cached end end @@ -420,6 +437,7 @@ Clear the fused convolution graph cache and buffer pool, freeing GPU memory. function clear_fused_conv_cache!() lock(_fused_conv_graph_cache_lock) do empty!(_fused_conv_graph_cache) + empty!(_fused_conv_graph_cache_order) end clear_fused_conv_buffer_pool!() return nothing diff --git a/test/dsp.jl b/test/dsp.jl index 19e0d35ff..a9fd47707 100644 --- a/test/dsp.jl +++ b/test/dsp.jl @@ -35,3 +35,22 @@ end v = rand(Float32, 20) @test Array(xcorr(MtlArray(u), MtlArray(v))) ≈ xcorr(u, v) rtol = 1.0f-3 end + +@testset "unsupported inputs" begin + # Unsupported element types error clearly instead of a silent slow CPU fallback. + @test_throws ArgumentError conv(MtlArray(rand(Int32, 16)), MtlArray(rand(Int32, 4))) + # xcorr supports only padmode = :none and scaling = :none. + u = MtlVector(rand(Float32, 16)) + v = MtlVector(rand(Float32, 8)) + @test_throws ArgumentError xcorr(u, v; padmode = :longest) + @test_throws ArgumentError xcorr(u, v; scaling = :biased) +end + +@testset "graph cache is bounded" begin + Metal.MPSGraphs.clear_fused_conv_cache!() + for n in 100:140 # 41 distinct sizes, cap is 32 + conv(MtlVector(rand(Float32, n)), MtlVector(rand(Float32, 5))) + end + @test length(Metal.MPSGraphs._fused_conv_graph_cache) <= 32 + @test length(Metal.MPSGraphs._fused_conv_buffer_pool) <= 32 +end From c2909b8d3f72c35044103439e5a94299af27ef7a Mon Sep 17 00:00:00 2001 From: Kaan Kesgin Date: Thu, 28 May 2026 11:09:14 +0200 Subject: [PATCH 09/12] Pipeline convolutions: drop the per-call host wait in conv_fft_fused conv_fft_fused waited on its command buffer every call, serializing all convolutions. The graph, the padding, and the result-extraction copy all run on the same in-order queue, and the extraction copies out of the pooled output buffer, so the per-call host wait is unnecessary -- the caller's Array() / synchronize() provides the needed sync. Removing it lets successive convolutions pipeline on the GPU. Throughput (batched): 1.2-2.5x faster (e.g. 2D 512x512, 15x15: 0.83 -> 0.33 ms). Correctness unchanged: engine 31/31, DSP 10/10. --- lib/mpsgraphs/convolution.jl | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/lib/mpsgraphs/convolution.jl b/lib/mpsgraphs/convolution.jl index 668bd07b0..f473b15ab 100644 --- a/lib/mpsgraphs/convolution.jl +++ b/lib/mpsgraphs/convolution.jl @@ -381,9 +381,13 @@ function conv_fft_fused( cmdbuf = MPSCommandBuffer(Metal.global_queue(current_device())) encode!(cmdbuf, cached.graph, NSDictionary(feeds), NSDictionary(resultdict), nil, default_exec_desc()) commit!(cmdbuf) - wait_completed(cmdbuf) + # No per-call host wait: the graph runs on the same in-order queue as the + # padding above and the extraction copy below, so results are correct once + # the caller synchronizes (e.g. via `Array`). Skipping the wait lets + # successive convolutions pipeline on the GPU instead of serializing. - # Extract appropriate region based on mode (must copy since output buffer is reused) + # Extract the requested region. This copies, so the pooled `output` buffer + # can be safely overwritten by the next call enqueued on the same queue. return _extract_conv_result_nd(output, output_sizes, full_sizes, mode) end end From 75a83a3f7bef5d8cc035c4d73bb33d1f84556c2f Mon Sep 17 00:00:00 2001 From: Kaan Kesgin Date: Thu, 28 May 2026 11:15:09 +0200 Subject: [PATCH 10/12] Update CPU-vs-GPU benchmarks after pipelining optimization Crossover moved left (GPU now wins from ~1e5 / 256^2) and large-size speedups roughly doubled (512^2: 1.8x -> 3.9x; 1e6: 3.0x -> 4.7x). Notes that profiling ruled out host-side graph-setup caching (feeds build is ~5us). --- perf/CONVOLUTION_PERFORMANCE.md | 36 ++++++++++++++++++++------------- 1 file changed, 22 insertions(+), 14 deletions(-) diff --git a/perf/CONVOLUTION_PERFORMANCE.md b/perf/CONVOLUTION_PERFORMANCE.md index 135ceb62e..86142af9b 100644 --- a/perf/CONVOLUTION_PERFORMANCE.md +++ b/perf/CONVOLUTION_PERFORMANCE.md @@ -9,30 +9,38 @@ decision and are intended as source material for the PR description. ## CPU vs GPU (`DSP.conv`) GPU `DSP.conv` (via `MetalDSPExt`) vs CPU `DSP.conv` (FFTW-backed), `Float32`. -GPU times are compute-only (`Metal.synchronize()`, data already on device). The -GPU has a ~0.4–0.8 ms fixed dispatch/sync overhead, so it wins only at scale. +GPU times are batched throughput (one `Metal.synchronize()`, data already on +device). Convolutions pipeline on the GPU — no per-call host wait — so the +effective per-call overhead is ~0.2 ms. 1-D (kernel 127): | signal | CPU (ms) | GPU (ms) | speedup | |-------:|---------:|---------:|--------:| -| 10³ | 0.037 | 0.412 | 0.1× | -| 10⁴ | 0.061 | 0.422 | 0.1× | -| 10⁵ | 0.288 | 0.491 | 0.6× | -| 10⁶ | 2.469 | 0.815 | **3.0×** | +| 10³ | 0.039 | 0.191 | 0.2× | +| 10⁴ | 0.059 | 0.202 | 0.3× | +| 10⁵ | 0.292 | 0.202 | **1.4×** | +| 10⁶ | 2.436 | 0.514 | **4.7×** | 2-D (kernel 15×15): | image | CPU (ms) | GPU (ms) | speedup | |------:|---------:|---------:|--------:| -| 128² | 0.128 | 0.598 | 0.2× | -| 256² | 0.427 | 0.592 | 0.7× | -| 512² | 1.457 | 0.830 | **1.8×** | -| 1024² | 5.135 | 1.633 | **3.1×** | - -Guidance: the GPU path pays off for large signals/images (1-D ≳ 10⁶, 2-D ≳ 512²) -or when data already lives on the GPU; for small inputs CPU FFTW is faster. A -one-off `CPU→GPU→CPU` round trip shifts the crossover further right. +| 128² | 0.116 | 0.378 | 0.3× | +| 256² | 0.444 | 0.334 | **1.3×** | +| 512² | 1.398 | 0.359 | **3.9×** | +| 1024² | 5.155 | 1.381 | **3.7×** | + +Guidance: the GPU path wins from ~10⁵ (1-D) / ~256² (2-D) upward, reaching ~4–5× +at large sizes; for smaller inputs CPU FFTW is faster, and a one-off +`CPU→GPU→CPU` round trip shifts the crossover right. + +Note: an earlier per-call host wait serialized convolutions (crossover ~10⁶ / +512²). Removing it so calls pipeline (commit `c2909b8d`) moved the crossover left +and roughly doubled the large-size speedup. Profiling showed the residual +per-call cost is GPU FFT compute plus kernel-launch overhead, not host-side graph +setup (feeds/`MPSGraphTensorData`/`NSDictionary` build is ~5 µs), so further +host-side caching was not worthwhile. ## Single 2-D image, `mode = :same` From 4b7a591f18906a7fbc1b9766191fab5ef13f5ecd Mon Sep 17 00:00:00 2001 From: Kaan Kesgin Date: Thu, 28 May 2026 11:50:14 +0200 Subject: [PATCH 11/12] Polish for review: drop banner comments, dead code, internal docstrings; Runic - Replace the '# ====' banner comments with plain section comments to match the sibling MPSGraphs files (matmul.jl, fft.jl, ... use none) - Remove dead code: _fast_pad_copy_contiguous! and its kernel _pad_copy_kernel_nd! (no callers; the N-D path uses copyto! + _zero_padding_regions_fused!) - Convert internal-helper docstrings to brief comments (matmul.jl style) - Apply Runic formatting (convolution.jl, ext/MetalDSPExt.jl) - Drop the explanatory comment from src/Metal.jl (no net change there) - Move convolution performance notes out of perf/ (folded into the PR description) No functional or performance change: engine 31/31, DSP 10/10; GPU timings match the prior commit when measured back-to-back. --- ext/MetalDSPExt.jl | 10 +-- lib/mpsgraphs/convolution.jl | 110 ++++++-------------------------- perf/CONVOLUTION_PERFORMANCE.md | 101 ----------------------------- src/Metal.jl | 2 - 4 files changed, 27 insertions(+), 196 deletions(-) delete mode 100644 perf/CONVOLUTION_PERFORMANCE.md diff --git a/ext/MetalDSPExt.jl b/ext/MetalDSPExt.jl index 760e8add7..12833f41c 100644 --- a/ext/MetalDSPExt.jl +++ b/ext/MetalDSPExt.jl @@ -17,10 +17,12 @@ const MtlConvNumber = Union{Float32, Float16, ComplexF32, ComplexF16} # below intercept DSP.conv/xcorr before DSP's generic CPU path, which would hit # disallowed scalar indexing on the device. @noinline function _conv_unsupported(u, v) - throw(ArgumentError( - "Metal convolution supports Float32/Float16 and their Complex types; " * - "got $(eltype(u)) and $(eltype(v)). Convert first, e.g. with `Float32.(x)`." - )) + throw( + ArgumentError( + "Metal convolution supports Float32/Float16 and their Complex types; " * + "got $(eltype(u)) and $(eltype(v)). Convert first, e.g. with `Float32.(x)`." + ) + ) end """ diff --git a/lib/mpsgraphs/convolution.jl b/lib/mpsgraphs/convolution.jl index f473b15ab..d93a6fce7 100644 --- a/lib/mpsgraphs/convolution.jl +++ b/lib/mpsgraphs/convolution.jl @@ -14,9 +14,7 @@ using AbstractFFTs export conv, conv_fft, conv_fft!, xcorr, imfilter export clear_fused_conv_cache! -# ============================================================================ # Helper Functions -# ============================================================================ """ nextfastfft(n::Integer) @@ -36,13 +34,10 @@ function nextfastfft(n::Integer) m == 1 && return n n += 1 end + return end -""" - _conv_output_size(signal_size, kernel_size, mode) - -Compute the output size for convolution based on mode. -""" +# Output size of a 1-D convolution dimension, per mode. function _conv_output_size(signal_size::Int, kernel_size::Int, mode::Symbol) full_size = signal_size + kernel_size - 1 if mode == :full @@ -56,11 +51,7 @@ function _conv_output_size(signal_size::Int, kernel_size::Int, mode::Symbol) end end -""" - _extract_conv_result(result, output_size, full_size, mode) - -Extract the appropriate portion of the convolution result based on mode. -""" +# Extract the requested portion of a 1-D convolution result, per mode. function _extract_conv_result( result::MtlArray{T, 1}, output_size::Int, full_size::Int, mode::Symbol ) where {T} @@ -107,9 +98,7 @@ function _extract_conv_result( return result[ranges...] end -# ============================================================================ # Fused MPSGraph Convolution (Single Graph Execution) -# ============================================================================ # Cache key for fused convolution graphs struct FusedConvGraphKey @@ -144,9 +133,7 @@ function _record_and_evict!(cache::AbstractDict, order::AbstractVector, key) return nothing end -# ============================================================================ # Buffer Pooling for Fused Convolution -# ============================================================================ # Key for buffer pool: (fft_sizes, eltype) struct BufferPoolKey @@ -166,17 +153,14 @@ const _fused_conv_buffer_pool = Dict{BufferPoolKey, CachedFusedConvBuffers}() const _fused_conv_buffer_pool_lock = ReentrantLock() const _fused_conv_buffer_pool_order = BufferPoolKey[] -""" -Get or create cached buffers for fused convolution. -Returns pre-allocated padded signal, kernel, and output buffers. -""" +# Get or create pooled padded signal/kernel/output buffers for a given FFT size. function _get_cached_buffers(fft_sizes::NTuple{N, Int}, ::Type{T}) where {N, T} key = BufferPoolKey(fft_sizes, T) cached = get(_fused_conv_buffer_pool, key, nothing) if cached !== nothing return cached end - lock(_fused_conv_buffer_pool_lock) do + return lock(_fused_conv_buffer_pool_lock) do cached = get(_fused_conv_buffer_pool, key, nothing) if cached !== nothing return cached @@ -205,9 +189,7 @@ function clear_fused_conv_buffer_pool!() return nothing end -# ============================================================================ # Fast Padding Kernel (Single kernel for copy + zero-pad) -# ============================================================================ # Custom Metal kernel that copies source data and zero-pads in one operation # This is ~4.5x faster than separate copyto! + broadcast zero operations @@ -221,50 +203,20 @@ function _pad_copy_kernel_1d!(dest, src, src_len) return end -""" -Copy source array to destination with zero-padding using a single GPU kernel. -Much faster than separate copyto! + broadcast operations (~4.5x speedup). -""" -function _fast_pad_copy!(dest::MtlVector{T}, src::MtlVector{T}) where T - src_len = length(src) - dest_len = length(dest) - threads = min(256, dest_len) - groups = cld(dest_len, threads) - @metal threads=threads groups=groups _pad_copy_kernel_1d!(dest, src, src_len) - return dest -end - -# N-D version: pad along all dimensions (linearized) -function _pad_copy_kernel_nd!(dest, src, src_linear_len) - i = thread_position_in_grid_1d() - if i <= src_linear_len - @inbounds dest[i] = src[i] - elseif i <= length(dest) - @inbounds dest[i] = zero(eltype(dest)) - end - return -end - -""" -Fast N-D padding: copies source to destination buffer with zero-padding. -For N-D arrays, this only works correctly when source fits contiguously at the start. -For general N-D padding with different sizes per dimension, use _fast_pad_copy_nd!. -""" -function _fast_pad_copy_contiguous!(dest::MtlArray{T, N}, src::MtlArray{T, N}) where {T, N} +# Copy a 1-D source into a zero-padded destination with a single GPU kernel +# (faster than separate copyto! + broadcast fill). +function _fast_pad_copy!(dest::MtlVector{T}, src::MtlVector{T}) where {T} src_len = length(src) dest_len = length(dest) threads = min(256, dest_len) groups = cld(dest_len, threads) - @metal threads=threads groups=groups _pad_copy_kernel_nd!(dest, src, src_len) + @metal threads = threads groups = groups _pad_copy_kernel_1d!(dest, src, src_len) return dest end -""" -Build a fused N-D convolution graph: rfft(signal) * rfft(kernel) → irfft → scale -All operations in a single MPSGraph for minimal command submission overhead. -Works for any dimensionality (1D, 2D, 3D, etc). -""" +# Build the fused N-D convolution graph (rfft * rfft -> irfft -> scale) as a single +# MPSGraph, for any dimensionality. function _build_fused_conv_graph_nd(fft_sizes::NTuple{N, Int}, ::Type{T}) where {N, T <: Union{Float32, Float16}} graph = MPSGraph() @@ -277,7 +229,7 @@ function _build_fused_conv_graph_nd(fft_sizes::NTuple{N, Int}, ::Type{T}) where # Metal uses reversed axis ordering: axis 0 in Metal = last axis in Julia # For N-D, we transform all dimensions - axes = NSArray([NSNumber(Int32(i)) for i in (N-1):-1:0]) + axes = NSArray([NSNumber(Int32(i)) for i in (N - 1):-1:0]) # Forward rfft on both inputs signal_fft = realToHermiteanFFTWithTensor(graph, signal_ph, axes, fft_desc_fwd, "signal_rfft") @@ -301,15 +253,13 @@ function _build_fused_conv_graph_nd(fft_sizes::NTuple{N, Int}, ::Type{T}) where return CachedFusedConvGraph(graph, signal_ph, kernel_ph, result_scaled) end -""" -Get or create a cached fused convolution graph. -""" +# Get or create the cached fused convolution graph for a given FFT size. function _get_cached_fused_conv_graph(key::FusedConvGraphKey) cached = get(_fused_conv_graph_cache, key, nothing) if cached !== nothing return cached end - lock(_fused_conv_graph_cache_lock) do + return lock(_fused_conv_graph_cache_lock) do cached = get(_fused_conv_graph_cache, key, nothing) if cached !== nothing return cached @@ -407,6 +357,7 @@ function _zero_padding_regions_fused!(arr::MtlArray{T, N}, data_sizes::NTuple{N, @view(arr[ranges...]) .= zero(T) end end + return end # Helper for extracting result based on mode @@ -450,16 +401,9 @@ end # Export the new function # (added to exports at top of file) -# ============================================================================ # Efficient Padding Helpers -# ============================================================================ - -""" - _zero_padding_regions!(arr, data_sizes, padded_sizes, dims) -Zero only the padding regions of an N-D array, avoiding unnecessary writes. -For each dimension in `dims`, zeros elements from (data_size+1) to padded_size. -""" +# Zero only the padding regions of an N-D array (avoids rewriting the data block). function _zero_padding_regions!( arr::MtlArray{T, N}, data_sizes::NTuple{N, Int}, padded_sizes::NTuple{N, Int}, dims::Tuple @@ -471,7 +415,7 @@ function _zero_padding_regions!( ranges = ntuple(N) do i if i == d # Padding region in this dimension - (data_sizes[i]+1):padded_sizes[i] + (data_sizes[i] + 1):padded_sizes[i] elseif i < d # For earlier dimensions, include entire padded size # (to cover corner regions that previous strips may have missed) @@ -488,9 +432,7 @@ function _zero_padding_regions!( return nothing end -# ============================================================================ # 1D FFT Convolution (Real Inputs - Optimized) -# ============================================================================ """ conv_fft(signal::MtlVector, kernel::MtlVector; mode=:full) @@ -528,12 +470,10 @@ function conv_fft( ) where {T <: Union{Float32, Float16}} # Delegate to fused implementation for better performance # (single MPSGraph execution instead of 4+ separate operations) - return conv_fft_fused(signal, kernel; mode=mode) + return conv_fft_fused(signal, kernel; mode = mode) end -# ============================================================================ # 1D FFT Convolution (Complex Inputs) -# ============================================================================ """ conv_fft(signal::MtlVector{Complex{T}}, kernel::MtlVector{Complex{T}}; mode=:full) @@ -563,10 +503,10 @@ function conv_fft( # Zero only the padding regions if ns < nfft - @view(signal_padded[(ns+1):nfft]) .= zero(Complex{T}) + @view(signal_padded[(ns + 1):nfft]) .= zero(Complex{T}) end if nk < nfft - @view(kernel_padded[(nk+1):nfft]) .= zero(Complex{T}) + @view(kernel_padded[(nk + 1):nfft]) .= zero(Complex{T}) end # FFT @@ -583,9 +523,7 @@ function conv_fft( return _extract_conv_result(y, output_size, full_size, mode) end -# ============================================================================ # N-D FFT Convolution (along specified dimensions) -# ============================================================================ """ conv_fft(signal::MtlArray, kernel::MtlArray; dims=1, mode=:full) @@ -626,7 +564,7 @@ function conv_fft( # Use fused implementation when convolving along ALL dimensions (faster single-graph execution) if length(dims_tuple) == N && Set(dims_tuple) == Set(1:N) - return conv_fft_fused(signal, kernel; mode=mode) + return conv_fft_fused(signal, kernel; mode = mode) end # Compute output sizes for each convolved dimension @@ -754,9 +692,7 @@ function conv_fft( return _extract_conv_result(y, output_sizes, full_sizes, mode, dims_tuple) end -# ============================================================================ # Cross-correlation -# ============================================================================ # GPU-friendly reverse along specified dimensions # Uses broadcasting to avoid scalar indexing @@ -831,9 +767,7 @@ function xcorr( end end -# ============================================================================ # In-place convolution (output pre-allocated) -# ============================================================================ """ conv_fft!(output, signal, kernel; dims=1, mode=:full) @@ -902,9 +836,7 @@ function imfilter(image::MtlMatrix{T}, kernel::MtlMatrix{T}) where {T <: Union{F return conv_fft(image, kernel; dims = (1, 2), mode = :same) end -# ============================================================================ # Unified Convolution API (FFT-based) -# ============================================================================ # # The unified `conv()` dispatches to the FFT convolution engine for 1D, 2D, and # N-D arrays. It is the internal entry point behind the public DSP.conv / diff --git a/perf/CONVOLUTION_PERFORMANCE.md b/perf/CONVOLUTION_PERFORMANCE.md deleted file mode 100644 index 86142af9b..000000000 --- a/perf/CONVOLUTION_PERFORMANCE.md +++ /dev/null @@ -1,101 +0,0 @@ -# Convolution performance notes - -Benchmarks backing the design of the GPU convolution support (the `DSP.conv` / -`DSP.xcorr` extension and its internal FFT engine). Measured on Apple Silicon -(Metal), `Float32`, with warmup + `Metal.synchronize()`, reported as the best of -5 runs × 30 calls (ms/call). These numbers drive the "single-implementation" -decision and are intended as source material for the PR description. - -## CPU vs GPU (`DSP.conv`) - -GPU `DSP.conv` (via `MetalDSPExt`) vs CPU `DSP.conv` (FFTW-backed), `Float32`. -GPU times are batched throughput (one `Metal.synchronize()`, data already on -device). Convolutions pipeline on the GPU — no per-call host wait — so the -effective per-call overhead is ~0.2 ms. - -1-D (kernel 127): - -| signal | CPU (ms) | GPU (ms) | speedup | -|-------:|---------:|---------:|--------:| -| 10³ | 0.039 | 0.191 | 0.2× | -| 10⁴ | 0.059 | 0.202 | 0.3× | -| 10⁵ | 0.292 | 0.202 | **1.4×** | -| 10⁶ | 2.436 | 0.514 | **4.7×** | - -2-D (kernel 15×15): - -| image | CPU (ms) | GPU (ms) | speedup | -|------:|---------:|---------:|--------:| -| 128² | 0.116 | 0.378 | 0.3× | -| 256² | 0.444 | 0.334 | **1.3×** | -| 512² | 1.398 | 0.359 | **3.9×** | -| 1024² | 5.155 | 1.381 | **3.7×** | - -Guidance: the GPU path wins from ~10⁵ (1-D) / ~256² (2-D) upward, reaching ~4–5× -at large sizes; for smaller inputs CPU FFTW is faster, and a one-off -`CPU→GPU→CPU` round trip shifts the crossover right. - -Note: an earlier per-call host wait serialized convolutions (crossover ~10⁶ / -512²). Removing it so calls pipeline (commit `c2909b8d`) moved the crossover left -and roughly doubled the large-size speedup. Profiling showed the residual -per-call cost is GPU FFT compute plus kernel-launch overhead, not host-side graph -setup (feeds/`MPSGraphTensorData`/`NSDictionary` build is ~5 µs), so further -host-side caching was not worthwhile. - -## Single 2-D image, `mode = :same` - -| image | kernel | fused FFT | MPS-direct | fused speedup | -|------:|-------:|----------:|-----------:|--------------:| -| 16² | 3² | 0.53 | 2.30 | 4.3× | -| 32² | 3² | 0.58 | 2.79 | 4.8× | -| 64² | 5² | 0.53 | 3.38 | 6.3× | -| 128² | 5² | 0.59 | 5.76 | 9.8× | -| 256² | 5² | 0.58 | 5.18 | 9.0× | -| 512² | 3²–63² | 0.6–0.8 | 2.4–11.1 | 2–18× | - -The fused FFT path is ~0.6 ms and essentially **kernel-size independent**, while -MPS-direct (`convolution2DWithSourceTensor`) is slower at every size measured — -including the small kernels the old `:auto` heuristic routed to it. - -## Repeated convolution: plan vs one-shot (512², `mode = :full`) - -| kernel | one-shot fused | `ConvFFTPlan` (reuse) | ratio | -|-------:|---------------:|----------------------:|------:| -| 7² | 0.65 | 1.23 | 0.52× | -| 15² | 0.62 | 1.19 | 0.52× | -| 31² | 0.66 | 1.19 | 0.56× | - -The `ConvFFTPlan` repeated-convolution path is **~2× slower** than simply calling -the one-shot fused convolution. The fused graph and its transforms are already -cached by the FFT graph cache, so the plan's separate rfft/irfft + buffer reuse -adds overhead without a payoff. - -## Design decisions - -- **Single implementation = fused FFT.** It is the fastest path in every regime - measured, so collapsing to it (per reviewer guidance on the FFT PR — "one - implementation, like CUDA.jl") costs **no performance**. It also removes the - `:auto` heuristic that mis-routed small kernels to the slower direct path, so - the trim is a net speed-up for that case. -- **Drop MPS-direct** from the signal-convolution engine: always slower here. -- **Drop `ConvFFTPlan` / `plan_conv_fft`:** slower than one-shot — negative value. -- **`imfilter`** routes to the fused FFT path (its old direct branch is removed). - -## Coordination with PR #745 and NNlib - -PR #745 (`MPSGraphs.graph_conv!`) implements **NN-style** convolution -(`convolution2DWithSourceTensor` with stride/dilation/padding/groups) for the -NNlib/Flux lane (issue #210). This work covers **signal-processing** convolution -(`DSP.conv` / `DSP.xcorr`, FFT-based). They are complementary: - -- Dropping MPS-direct here also drops our `convolution2DWithSourceTensor` / - `MPSGraphConvolution2DOpDescriptor` wrappers, **removing the overlap** with #745. -- We do **not** add `NNlib.conv` here — it would duplicate #745. The NN path - belongs in #745 (or a follow-up NNlib extension), not the DSP/signal PR. - -## Caveats - -Numbers are for single-array (all-dims) `Float32` signal convolution on one GPU. -NN-style batched/channelled convolution (many small filters, NCHW layout) is a -different workload where the MPS conv2d primitive is appropriate — that is #745's -domain, not this PR's. diff --git a/src/Metal.jl b/src/Metal.jl index 1ddadebc8..2d31a2686 100644 --- a/src/Metal.jl +++ b/src/Metal.jl @@ -60,8 +60,6 @@ include("../lib/mps/MPS.jl") export MPS include("../lib/mpsgraphs/MPSGraphs.jl") export MPSGraphs -# The convolution engine lives in MPSGraphs and is exposed publicly through the -# DSP.jl package extension (DSP.conv / DSP.xcorr), not as bespoke Metal exports. # LinearAlgebra include("linalg.jl") From 5c48317f59ea4c932e95a7a0ac00aafe781a787b Mon Sep 17 00:00:00 2001 From: Kaan Kesgin Date: Thu, 28 May 2026 12:18:58 +0200 Subject: [PATCH 12/12] Trim for review: defer imfilter, collapse conv() dispatch, tighten docs - Defer imfilter to a follow-up (a convenience on top of conv; narrows this PR to the core DSP.conv/xcorr) -- removes the function, its export, and its testset - Collapse the six conv() methods to two (real/complex over MtlArray{T,N}) with a compile-time N==1 branch -- behavior-identical - Tighten the conv() docstring; convert the test's banner comments to plain ones convolution.jl 938 -> 824 lines. The hot path (conv_fft_fused, graph builder, padding) is byte-identical, so no performance change. Tests: engine 28/28, DSP 10/10. --- lib/mpsgraphs/convolution.jl | 132 +++------------------------------- test/mpsgraphs/convolution.jl | 54 ++------------ 2 files changed, 13 insertions(+), 173 deletions(-) diff --git a/lib/mpsgraphs/convolution.jl b/lib/mpsgraphs/convolution.jl index d93a6fce7..f01e71205 100644 --- a/lib/mpsgraphs/convolution.jl +++ b/lib/mpsgraphs/convolution.jl @@ -11,7 +11,7 @@ using AbstractFFTs -export conv, conv_fft, conv_fft!, xcorr, imfilter +export conv, conv_fft, conv_fft!, xcorr export clear_fused_conv_cache! # Helper Functions @@ -787,55 +787,6 @@ function conv_fft!( end -""" - imfilter(image::MtlMatrix, kernel::MtlMatrix) - -Apply a 2-D filter `kernel` to `image` via FFT-based convolution (`:same` mode). - -A convenience wrapper following the ImageFiltering.jl interface; the kernel is -centered on each pixel. - -# Arguments -- `image`: 2D input image -- `kernel`: 2D filter kernel - -# Returns -Filtered image with same size as input (`:same` mode). - -# Example -```julia -using Metal - -# Create test image -image = MtlMatrix(randn(Float32, 512, 512)) - -# Gaussian blur (5×5 approximation) -gaussian = MtlMatrix(Float32[ - 1 4 6 4 1 - 4 16 24 16 4 - 6 24 36 24 6 - 4 16 24 16 4 - 1 4 6 4 1 -] ./ 256) - -blurred = imfilter(image, gaussian) - -# Sobel edge detection -sobel_x = MtlMatrix(Float32[-1 0 1; -2 0 2; -1 0 1] ./ 8) -sobel_y = MtlMatrix(Float32[-1 -2 -1; 0 0 0; 1 2 1] ./ 8) -edges_x = imfilter(image, sobel_x) -edges_y = imfilter(image, sobel_y) -edges = sqrt.(edges_x.^2 .+ edges_y.^2) -``` - -# Notes -- Equivalent to `conv(image, kernel; dims=(1, 2), mode=:same)`. -- The kernel is centered on each pixel (like ImageFiltering.jl's `imfilter`). -""" -function imfilter(image::MtlMatrix{T}, kernel::MtlMatrix{T}) where {T <: Union{Float32, Float16}} - return conv_fft(image, kernel; dims = (1, 2), mode = :same) -end - # Unified Convolution API (FFT-based) # # The unified `conv()` dispatches to the FFT convolution engine for 1D, 2D, and @@ -845,92 +796,27 @@ end """ conv(signal::MtlArray, kernel::MtlArray; mode=:full, dims=nothing) -Compute the linear convolution of `signal` and `kernel` via the FFT convolution -theorem. This is the internal engine entry point; the public interface is -`DSP.conv` (provided by the DSP.jl extension). - -# Arguments -- `signal`: Input signal (1D, 2D, or N-D MtlArray) -- `kernel`: Convolution kernel (same dimensions as signal) -- `mode`: Output size mode - - `:full` (default): Full convolution output - - `:same`: Output has same size as signal (centered) - - `:valid`: Only fully overlapping region -- `dims`: Dimensions along which to convolve - - `nothing` (default): All dimensions for 1D/2D, dim 1 for N-D - - Integer or tuple: Specific dimension(s) - -# Returns -MtlArray with the convolution result. - -# Examples - -```julia -using Metal, Metal.MPSGraphs +Linear convolution of `signal` and `kernel` via the FFT convolution theorem. +This is the internal engine; the public interface is `DSP.conv` (DSP.jl extension). -# 1D signal processing -signal = MtlVector(randn(Float32, 10000)) -kernel = MtlVector(Float32[0.25, 0.5, 0.25]) # Simple smoothing -smoothed = conv(signal, kernel; mode=:same) - -# 2D image filtering -image = MtlMatrix(randn(Float32, 512, 512)) -sobel_x = MtlMatrix(Float32[-1 0 1; -2 0 2; -1 0 1] ./ 8) -edges = conv(image, sobel_x; mode=:same) -``` - -# See Also -- `imfilter`: ImageFiltering.jl-style 2D filtering -- `xcorr`: Cross-correlation +`mode` is `:full` (default), `:same`, or `:valid`. `dims` selects the convolved +dimensions (default: all dimensions for 1-D/2-D, dimension 1 for N-D). """ -function conv( - signal::MtlVector{T}, kernel::MtlVector{T}; - mode::Symbol = :full, dims = nothing - ) where {T <: Union{Float32, Float16}} - return conv_fft(signal, kernel; mode = mode) -end - -# Complex 1D -function conv( - signal::MtlVector{Complex{T}}, kernel::MtlVector{Complex{T}}; - mode::Symbol = :full, dims = nothing - ) where {T <: Union{Float32, Float16}} - return conv_fft(signal, kernel; mode = mode) -end - -# 2D -function conv( - signal::MtlMatrix{T}, kernel::MtlMatrix{T}; - mode::Symbol = :full, dims = nothing - ) where {T <: Union{Float32, Float16}} - conv_dims = dims === nothing ? (1, 2) : (dims isa Int ? (dims,) : Tuple(dims)) - return conv_fft(signal, kernel; dims = conv_dims, mode = mode) -end - -# Complex 2D -function conv( - signal::MtlMatrix{Complex{T}}, kernel::MtlMatrix{Complex{T}}; - mode::Symbol = :full, dims = nothing - ) where {T <: Union{Float32, Float16}} - conv_dims = dims === nothing ? (1, 2) : (dims isa Int ? (dims,) : Tuple(dims)) - return conv_fft(signal, kernel; dims = conv_dims, mode = mode) -end - -# N-D generic (N > 2) function conv( signal::MtlArray{T, N}, kernel::MtlArray{T, N}; mode::Symbol = :full, dims = nothing ) where {T <: Union{Float32, Float16}, N} - conv_dims = dims === nothing ? 1 : (dims isa Int ? (dims,) : Tuple(dims)) + N == 1 && return conv_fft(signal, kernel; mode = mode) + conv_dims = dims === nothing ? (N == 2 ? (1, 2) : 1) : (dims isa Int ? (dims,) : Tuple(dims)) return conv_fft(signal, kernel; dims = conv_dims, mode = mode) end -# Complex N-D function conv( signal::MtlArray{Complex{T}, N}, kernel::MtlArray{Complex{T}, N}; mode::Symbol = :full, dims = nothing ) where {T <: Union{Float32, Float16}, N} - conv_dims = dims === nothing ? 1 : (dims isa Int ? (dims,) : Tuple(dims)) + N == 1 && return conv_fft(signal, kernel; mode = mode) + conv_dims = dims === nothing ? (N == 2 ? (1, 2) : 1) : (dims isa Int ? (dims,) : Tuple(dims)) return conv_fft(signal, kernel; dims = conv_dims, mode = mode) end diff --git a/test/mpsgraphs/convolution.jl b/test/mpsgraphs/convolution.jl index b8b8e99df..25c357247 100644 --- a/test/mpsgraphs/convolution.jl +++ b/test/mpsgraphs/convolution.jl @@ -1,8 +1,8 @@ # Tests for the FFT-based convolution engine in MPSGraphs (conv_fft, the unified -# conv(), xcorr, imfilter). The engine is internal; the public interface is +# conv(), xcorr). The engine is internal; the public interface is # DSP.conv / DSP.xcorr (tested in test/dsp.jl). These tests exercise the engine # directly, so import its symbols from the submodule. -using Metal.MPSGraphs: conv, conv_fft, conv_fft!, xcorr, imfilter +using Metal.MPSGraphs: conv, conv_fft, conv_fft!, xcorr # Simple reference convolution for verification (CPU) function ref_conv(u::Vector{T}, v::Vector{T}) where {T} @@ -26,9 +26,7 @@ rtol(::Type{ComplexF32}) = 1.0e-4 if MPS.is_supported(device()) - # ============================================================================ # FFT Convolution Tests - # ============================================================================ @testset "FFT Convolution" begin @testset "1D Real Convolution" begin @@ -104,9 +102,7 @@ if MPS.is_supported(device()) end end - # ============================================================================ # Cross-correlation Tests - # ============================================================================ @testset "Cross-correlation" begin @testset for T in [Float32, Float16] @@ -124,45 +120,7 @@ if MPS.is_supported(device()) end - # ============================================================================ - # imfilter Tests - # ============================================================================ - - @testset "imfilter" begin - @testset "Small Kernel (uses direct)" begin - image = rand(Float32, 256, 256) - kernel = Float32[ - -1 0 1 - -2 0 2 - -1 0 1 - ] ./ 8 # Sobel - - d_image = MtlMatrix(image) - d_kernel = MtlMatrix(kernel) - - result = Array(imfilter(d_image, d_kernel)) - @test size(result) == size(image) - - # Verify correctness against FFT - result_fft = Array(conv_fft(d_image, d_kernel; dims = (1, 2), mode = :same)) - @test isapprox(result, result_fft, rtol = 1.0e-4) - end - - @testset "Large Kernel (uses FFT)" begin - image = rand(Float32, 256, 256) - kernel = rand(Float32, 15, 15) # Larger than threshold - - d_image = MtlMatrix(image) - d_kernel = MtlMatrix(kernel) - - result = Array(imfilter(d_image, d_kernel)) - @test size(result) == size(image) - end - end - - # ============================================================================ # Unified conv() API Tests - # ============================================================================ @testset "Unified conv() API" begin @testset "1D Convolution" begin @@ -178,20 +136,18 @@ if MPS.is_supported(device()) @test length(result_valid) == 991 end - @testset "2D Auto-Selection" begin + @testset "2D matches conv_fft" begin image = MtlMatrix(rand(Float32, 128, 128)) small_kernel = MtlMatrix(rand(Float32, 3, 3)) large_kernel = MtlMatrix(rand(Float32, 15, 15)) - # Small kernel should auto-select direct result_small = conv(image, small_kernel; mode = :same) @test size(result_small) == size(image) - # Large kernel should auto-select FFT result_large = conv(image, large_kernel; mode = :same) @test size(result_large) == size(image) - # Both should give same result as explicit FFT + # conv() should match an explicit conv_fft call expected_small = conv_fft(image, small_kernel; dims = (1, 2), mode = :same) expected_large = conv_fft(image, large_kernel; dims = (1, 2), mode = :same) @@ -210,9 +166,7 @@ if MPS.is_supported(device()) end end - # ============================================================================ # Edge Cases - # ============================================================================ @testset "Edge Cases" begin @testset "Single Element Kernel" begin