[Metal] Add implicit matmul pathway for mx.conv3d (#3147)

Co-authored-by: Gleb Sterkin <[email protected]>
Co-authored-by: Angelos Katharopoulos <[email protected]>
This commit is contained in:
Gleb Sterkin
2026-02-23 17:52:50 -08:00
committed by GitHub
co-authored by Gleb Sterkin Angelos Katharopoulos
parent d4c81062ad
commit 1d8d693d08
13 changed files with 1238 additions and 65 deletions
+152
View File
@@ -0,0 +1,152 @@
import math
import time
import mlx.core as mx
import numpy as np
import torch
N_warmup = 2
N_iter_bench = 10
N_iter_func = 10
def bench(f, a, b, b_prime):
for i in range(N_warmup):
f(a, b, b_prime)
torch.mps.synchronize()
s = time.perf_counter_ns()
for i in range(N_iter_bench):
f(a, b, b_prime)
e = time.perf_counter_ns()
return (e - s) * 1e-9
def make_mx_conv_3D(strides=(1, 1, 1), padding=(0, 0, 0), groups=1):
def mx_conv_3D(a, b, b_prime):
y = a
for i in range(N_iter_func):
y = mx.conv3d(y, b, stride=strides, padding=padding, groups=groups)
y = mx.conv3d(y, b_prime, stride=strides, padding=padding, groups=groups)
mx.eval(y)
return y
return mx_conv_3D
def make_pt_conv_3D(strides=(1, 1, 1), padding=(0, 0, 0), groups=1):
@torch.no_grad()
def pt_conv_3D(a, b, b_prime):
y = a
for i in range(N_iter_func):
y = torch.conv3d(y, b, stride=strides, padding=padding, groups=groups)
y = torch.conv3d(y, b_prime, stride=strides, padding=padding, groups=groups)
torch.mps.synchronize()
return y
return pt_conv_3D
def bench_shape(N, D, H, W, C, kD, kH, kW, O, strides, padding, groups, np_dtype):
scale = 1.0 / math.sqrt(kD * kH * kW * C)
a_np = np.random.uniform(0, 0.5, (N, D, H, W, C))
b_np = np.random.uniform(-scale, scale, (O, kD, kH, kW, int(C / groups)))
b_prime_np = np.random.uniform(-scale, scale, (C, kD, kH, kW, int(O / groups)))
a_np, b_np, b_prime_np = map(lambda x: x.astype(np_dtype), (a_np, b_np, b_prime_np))
a_mx, b_mx, b_prime_mx = map(lambda x: mx.array(x), (a_np, b_np, b_prime_np))
a_pt, b_pt, b_prime_pt = map(
lambda x: torch.from_numpy(x.transpose(0, 4, 1, 2, 3)).to("mps"),
(a_np, b_np, b_prime_np),
)
torch.mps.synchronize()
f_mx = make_mx_conv_3D(strides, padding, groups)
f_pt = make_pt_conv_3D(strides, padding, groups)
time_torch = bench(f_pt, a_pt, b_pt, b_prime_pt)
time_mlx = bench(f_mx, a_mx, b_mx, b_prime_mx)
# Measure MLX memory
mx.clear_cache()
mx.reset_peak_memory()
y = mx.conv3d(a_mx, b_mx, stride=strides, padding=padding, groups=groups)
mx.eval(y)
mlx_peak_mb = mx.get_peak_memory() / 1024**2
mlx_active_mb = mx.get_active_memory() / 1024**2
del y
# Measure PyTorch MPS memory
torch.mps.synchronize()
torch.mps.empty_cache()
y = torch.conv3d(a_pt, b_pt, stride=strides, padding=padding, groups=groups)
torch.mps.synchronize()
pt_current_mb = torch.mps.current_allocated_memory() / 1024**2
pt_driver_mb = torch.mps.driver_allocated_memory() / 1024**2
del y
out_mx = mx.conv3d(a_mx, b_mx, stride=strides, padding=padding, groups=groups)
out_pt = torch.conv3d(
a_pt.to("cpu"), b_pt.to("cpu"), stride=strides, padding=padding, groups=groups
)
out_pt = torch.permute(out_pt, (0, 2, 3, 4, 1))
out_pt = out_pt.numpy(force=True)
atol = 2e-5 if np_dtype == np.float32 else 5e-4
if not np.allclose(out_pt, out_mx, atol=atol):
print(
f"Failed at {(N, D, H, W, C)}, {(O, kD, kH, kW, C)} "
f"[strides = {strides}, padding = {padding}, groups = {groups}] "
f"with max(|a - b|) = {np.max(np.abs(out_pt - out_mx))}"
)
return time_mlx, time_torch, mlx_peak_mb, mlx_active_mb, pt_current_mb, pt_driver_mb
if __name__ == "__main__":
dtypes = ("float16", "float32")
shapes = (
# (C % 16 == 0)
(4, 16, 16, 16, 32, 3, 3, 3, 32, (1, 1, 1), (1, 1, 1), 1),
(4, 16, 16, 16, 64, 3, 3, 3, 64, (1, 1, 1), (1, 1, 1), 1),
(4, 16, 16, 16, 128, 3, 3, 3, 128, (1, 1, 1), (1, 1, 1), 1),
(4, 32, 32, 32, 64, 3, 3, 3, 64, (1, 1, 1), (1, 1, 1), 1),
(4, 32, 32, 32, 128, 3, 3, 3, 128, (1, 1, 1), (1, 1, 1), 1),
# Larger spatial dims
(2, 64, 64, 64, 32, 3, 3, 3, 64, (1, 1, 1), (1, 1, 1), 1),
(1, 64, 64, 64, 64, 3, 3, 3, 128, (1, 1, 1), (1, 1, 1), 1),
# Strided
(4, 32, 32, 32, 64, 3, 3, 3, 128, (2, 2, 2), (1, 1, 1), 1),
# Asymmetric kernels
(4, 32, 32, 32, 64, 3, 1, 1, 128, (1, 1, 1), (1, 0, 0), 1),
(4, 32, 32, 32, 64, 1, 3, 3, 128, (1, 1, 1), (0, 1, 1), 1),
# (C % 16 != 0)
(4, 16, 16, 16, 21, 3, 3, 3, 21, (1, 1, 1), (1, 1, 1), 1),
(4, 16, 16, 16, 55, 3, 3, 3, 55, (1, 1, 1), (1, 1, 1), 1),
(4, 32, 32, 32, 55, 3, 3, 3, 55, (1, 1, 1), (1, 1, 1), 1),
(4, 16, 16, 16, 3, 3, 3, 3, 32, (1, 1, 1), (1, 1, 1), 1),
)
for dtype in dtypes:
print(f"\n{'=' * 120}" f"\n dtype: {dtype}" f"\n{'=' * 120}")
print(
f"{'(N, D, H, W, C)':<26s} {'( O, kD, kH, kW, C)':<24s} "
f"{'stride':<12s} {'pads':<12s} {'groups':>6s} "
f"{'diff%':>7s} "
f"{'MLX peak':>9s} {'MLX act':>8s} {'PT cur':>8s} {'PT drv':>8s}"
)
for N, D, H, W, C, kD, kH, kW, O, strides, padding, groups in shapes:
np_dtype = getattr(np, dtype)
time_mlx, time_torch, mlx_peak, mlx_act, pt_cur, pt_drv = bench_shape(
N, D, H, W, C, kD, kH, kW, O, strides, padding, groups, np_dtype
)
diff = time_torch / time_mlx - 1.0
print(
f"({N}, {D:3d}, {H:3d}, {W:3d}, {C:3d}), ({O:3d}, {kD:2d}, {kH:2d}, {kW:2d}, {C:3d}), "
f"{strides}, {padding}, {groups:6d}, "
f"{100. * diff:+6.1f}% "
f"{mlx_peak:8.1f} {mlx_act:7.1f} {pt_cur:7.1f} {pt_drv:7.1f}"
)
+1
View File
@@ -71,6 +71,7 @@ if(MLX_METAL_JIT)
kernels/steel/conv/loaders/loader_channel_l.h
kernels/steel/conv/loaders/loader_channel_n.h)
make_jit_source(steel/conv/kernels/steel_conv)
make_jit_source(steel/conv/kernels/steel_conv_3d)
make_jit_source(steel/conv/kernels/steel_conv_general kernels/steel/defines.h
kernels/steel/conv/loaders/loader_general.h)
+286 -49
View File
@@ -4,6 +4,7 @@
#include <numeric>
#include "mlx/backend/gpu/copy.h"
#include "mlx/backend/gpu/slicing.h"
#include "mlx/backend/metal/device.h"
#include "mlx/backend/metal/kernels.h"
#include "mlx/backend/metal/kernels/defines.h"
@@ -19,13 +20,23 @@ namespace mlx::core {
namespace {
inline array
ensure_row_contiguous(const array& x, metal::Device& d, const Stream& s) {
if (x.flags().row_contiguous) {
return x;
}
auto result = contiguous_copy_gpu(x, s);
d.add_temporary(result, s.index);
return result;
}
template <int N>
void explicit_gemm_conv_ND_gpu(
const Stream& s,
metal::Device& d,
const array& in,
const array& wt,
array out,
array& out,
const MLXConvParams<N>& conv_params) {
// Get gemm shapes
int implicit_M = out.size() / conv_params.O;
@@ -96,7 +107,7 @@ void explicit_gemm_conv_group_ND_gpu(
metal::Device& d,
const array& in,
const array& wt,
array out,
array& out,
const MLXConvParams<N>& conv_params) {
const int groups = conv_params.groups;
const int C_per_group = conv_params.C / conv_params.groups;
@@ -182,7 +193,7 @@ void implicit_gemm_conv_2D_gpu(
metal::Device& d,
const array& in,
const array& wt,
array out,
array& out,
const MLXConvParams<2>& conv_params) {
const int groups = conv_params.groups;
const int C_per_group = conv_params.C / conv_params.groups;
@@ -315,7 +326,7 @@ void implicit_gemm_conv_2D_general_gpu(
metal::Device& d,
const array& in,
const array& wt,
array out,
array& out,
const MLXConvParams<2>& conv_params) {
// Deduce implicit gemm size
int implicit_M = conv_params.N * conv_params.oS[0] * conv_params.oS[1];
@@ -489,12 +500,223 @@ void implicit_gemm_conv_2D_general_gpu(
compute_encoder.dispatch_threadgroups(grid_dims, group_dims);
}
void implicit_gemm_conv_3D_gpu(
const Stream& s,
metal::Device& d,
const array& in,
const array& wt,
array& out,
const MLXConvParams<3>& conv_params) {
const int groups = conv_params.groups;
const int C_per_group = conv_params.C / conv_params.groups;
const int O_per_group = conv_params.O / conv_params.groups;
// Deduce implicit gemm size
const int implicit_M =
conv_params.N * conv_params.oS[0] * conv_params.oS[1] * conv_params.oS[2];
const int implicit_N = O_per_group;
const int implicit_K =
conv_params.wS[0] * conv_params.wS[1] * conv_params.wS[2] * C_per_group;
// Determine block and warp tiles
int wm = 2, wn = 2;
int bm = implicit_M >= 8192 && C_per_group >= 64 ? 64 : 32;
int bn = (bm == 64 || implicit_N >= 64) ? 64 : 32;
int bk = 16;
if (implicit_N <= 16) {
bn = 8;
wm = 4;
wn = 1;
}
int tn = (implicit_N + bn - 1) / bn;
int tm = (implicit_M + bm - 1) / bm;
int swizzle_log = 0;
bool small_filter =
(conv_params.wS[0] <= 16 && conv_params.wS[1] <= 16 &&
conv_params.wS[2] <= 16);
int channel_k_iters = ((C_per_group + bk - 1) / bk);
int gemm_k_iters = conv_params.wS[0] * conv_params.wS[1] * conv_params.wS[2] *
channel_k_iters;
// Fix host side helper params
int sign = (conv_params.flip ? -1 : 1);
int ijw = conv_params.in_strides[3] * conv_params.kdil[2];
int ijh = conv_params.in_strides[2] * conv_params.kdil[1];
int ijd = conv_params.in_strides[1] * conv_params.kdil[0];
int inp_jump_w = sign * ijw;
int inp_jump_h = sign * (ijh - (conv_params.wS[2] - 1) * ijw);
int inp_jump_d = sign *
(ijd - (conv_params.wS[1] - 1) * ijh - (conv_params.wS[2] - 1) * ijw);
int inp_jump_c = bk - sign * (conv_params.wS[0] - 1) * ijd -
sign * (conv_params.wS[1] - 1) * ijh -
sign * (conv_params.wS[2] - 1) * ijw;
// Build implicit gemm params
ImplicitGemmConv3DParams gemm_params{
/* const int M = */ implicit_M,
/* const int N = */ implicit_N,
/* const int K = */ implicit_K,
/* const int gemm_k_iterations = */ gemm_k_iters,
/* const int inp_jump_w = */ inp_jump_w,
/* const int inp_jump_h = */ inp_jump_h,
/* const int inp_jump_d = */ inp_jump_d,
/* const int inp_jump_c = */ inp_jump_c,
/* const int tiles_n = */ tn,
/* const int tiles_m = */ tm,
/* const int swizzle_log = */ swizzle_log};
// Determine kernel
std::string kname;
kname.reserve(64);
concatenate(
kname,
"implicit_gemm_conv_3d_",
type_to_name(out),
"_bm",
bm,
"_bn",
bn,
"_bk",
bk,
"_wm",
wm,
"_wn",
wn,
"_filter_",
small_filter ? 's' : 'l');
// Encode and dispatch kernel
auto& compute_encoder = d.get_command_encoder(s.index);
auto kernel =
get_steel_conv_3d_kernel(d, kname, out, bm, bn, bk, wm, wn, small_filter);
compute_encoder.set_compute_pipeline_state(kernel);
// Deduce grid launch dimensions
int tile = 1 << swizzle_log;
size_t grid_dim_y = (tm + tile - 1) / tile;
size_t grid_dim_x = tn * tile;
MTL::Size group_dims = MTL::Size(32, wn, wm);
MTL::Size grid_dims = MTL::Size(grid_dim_x, grid_dim_y, groups);
// Encode arrays
compute_encoder.set_input_array(in, 0);
compute_encoder.set_input_array(wt, 1);
compute_encoder.set_output_array(out, 2);
// Encode params
compute_encoder.set_bytes(conv_params, 3);
compute_encoder.set_bytes(gemm_params, 4);
// Launch kernel
compute_encoder.dispatch_threadgroups(grid_dims, group_dims);
}
void pad_and_slice_conv_3D_gpu(
const Stream& s,
metal::Device& d,
const array& in_pre,
const array& wt_pre,
array& out,
const MLXConvParams<3>& conv_params) {
// For now assume conv_params.groups == 1
int extra_c = ((conv_params.C + 15) / 16) * 16 - conv_params.C;
int extra_o = ((conv_params.O + 15) / 16) * 16 - conv_params.O;
// Pad function
auto pad_array = [&](const array& x, int pad_ax_first, int pad_ax_last) {
if (pad_ax_first == 0 && pad_ax_last == 0) {
return ensure_row_contiguous(x, d, s);
}
auto xshape = x.shape();
xshape.front() += pad_ax_first;
xshape.back() += pad_ax_last;
array x_copy(xshape, x.dtype(), nullptr, {});
array zero(0, x.dtype());
pad_gpu(x, zero, x_copy, {0, -1}, {0, 0}, s);
d.add_temporary(x_copy, s.index);
return x_copy;
};
// Allocate space for the intermediate output. Don't save it as a temporary
// since it will be sliced to the output so they share the buffer.
auto oshape = out.shape();
oshape.back() += extra_o;
array intermediate(oshape, out.dtype(), nullptr, {});
intermediate.set_data(allocator::malloc(intermediate.nbytes()));
// Actually pad and conv
array in = pad_array(in_pre, 0, extra_c);
array wt = pad_array(wt_pre, extra_o, extra_c);
auto new_params =
MLXConvParams<3>::with_padded_channels(conv_params, extra_o, extra_c);
implicit_gemm_conv_3D_gpu(s, d, in, wt, intermediate, new_params);
// Slice out
out.copy_shared_buffer(
intermediate, intermediate.strides(), {0}, intermediate.data_size());
}
void dispatch_conv_3D_gpu(
const Stream& s,
metal::Device& d,
const array& in_pre,
const array& wt_pre,
array& out,
const MLXConvParams<3>& conv_params,
std::vector<array>& copies) {
bool is_idil_one = conv_params.idil[0] == 1 && conv_params.idil[1] == 1 &&
conv_params.idil[2] == 1;
const int C_per_group = conv_params.C / conv_params.groups;
const int O_per_group = conv_params.O / conv_params.groups;
bool mod16_channels =
C_per_group % 16 == 0 && (O_per_group <= 16 || O_per_group % 16 == 0);
// Check if we can do implicit gemm but the channels are not divisible by 16
// so we can pad and slice.
//
// We check it first because it doesn't need contiguous inputs and it needs
// different output allocation.
if (is_idil_one && !mod16_channels && conv_params.groups == 1) {
return pad_and_slice_conv_3D_gpu(s, d, in_pre, wt_pre, out, conv_params);
}
// Allocate the output and ensure contiguous inputs
out.set_data(allocator::malloc(out.nbytes()));
auto in = ensure_row_contiguous(in_pre, d, s);
auto wt = ensure_row_contiguous(wt_pre, d, s);
// Perform the implicit gemm
if (is_idil_one && mod16_channels) {
return implicit_gemm_conv_3D_gpu(s, d, in, wt, out, conv_params);
}
// Explicit gemms where we unfold and do a matmul
// (separate one for groups > 1)
if (conv_params.groups > 1) {
return explicit_gemm_conv_group_ND_gpu(s, d, in, wt, out, conv_params);
}
return explicit_gemm_conv_ND_gpu(s, d, in, wt, out, conv_params);
}
void winograd_conv_2D_gpu(
const Stream& s,
metal::Device& d,
const array& in,
const array& wt,
array out,
array& out,
const MLXConvParams<2>& conv_params,
std::vector<array>& copies_w) {
Shape padded_shape = {
@@ -688,7 +910,7 @@ void depthwise_conv_2D_gpu(
metal::Device& d,
const array& in,
const array& wt,
array out,
array& out,
const MLXConvParams<2>& conv_params) {
std::string base_name;
base_name.reserve(32);
@@ -750,7 +972,7 @@ void dispatch_conv_2D_gpu(
metal::Device& d,
const array& in,
const array& wt,
array out,
array& out,
const MLXConvParams<2>& conv_params,
std::vector<array>& copies) {
bool is_stride_one = conv_params.str[0] == 1 && conv_params.str[1] == 1;
@@ -811,8 +1033,8 @@ void depthwise_conv_1D_gpu(
const Stream& s,
metal::Device& d,
const array& in,
array wt,
array out) {
const array& wt,
array& out) {
bool large = in.size() > INT32_MAX || in.data_size() > INT32_MAX;
std::string base_name;
base_name.reserve(32);
@@ -822,10 +1044,6 @@ void depthwise_conv_1D_gpu(
large ? "_large" : "",
type_to_name(out));
if (!wt.flags().row_contiguous) {
wt = contiguous_copy_gpu(wt, s);
d.add_temporary(wt, s.index);
}
auto& compute_encoder = d.get_command_encoder(s.index);
auto kernel = d.get_kernel(base_name);
compute_encoder.set_compute_pipeline_state(kernel);
@@ -860,9 +1078,9 @@ void depthwise_conv_1D_gpu(
void conv_1D_gpu(
const Stream& s,
metal::Device& d,
const array& in,
const array& wt,
array out,
const array& in_pre,
const array& wt_pre,
array& out,
const std::vector<int>& padding,
const std::vector<int>& wt_strides,
const std::vector<int>& wt_dilation,
@@ -870,6 +1088,11 @@ void conv_1D_gpu(
int groups,
bool flip,
std::vector<array>& copies) {
// Allocate space and ensure weights are contiguous
out.set_data(allocator::malloc(out.nbytes()));
auto in = ensure_row_contiguous(in_pre, d, s);
auto wt = ensure_row_contiguous(wt_pre, d, s);
bool is_idil_one = in_dilation[0] == 1;
int C = in.shape(2);
int O = wt.shape(0);
@@ -942,9 +1165,9 @@ void conv_1D_gpu(
void conv_2D_gpu(
const Stream& s,
metal::Device& d,
const array& in,
const array& wt,
array out,
const array& in_pre,
const array& wt_pre,
array& out,
const std::vector<int>& padding,
const std::vector<int>& wt_strides,
const std::vector<int>& wt_dilation,
@@ -952,6 +1175,11 @@ void conv_2D_gpu(
const int groups,
bool flip,
std::vector<array>& copies) {
// Allocate space and ensure weights are contiguous
out.set_data(allocator::malloc(out.nbytes()));
auto in = ensure_row_contiguous(in_pre, d, s);
auto wt = ensure_row_contiguous(wt_pre, d, s);
// Make conv params
MLXConvParams<2> conv_params{
/* const int N = */ static_cast<int>(in.shape(0)),
@@ -989,8 +1217,20 @@ void conv_3D_gpu(
const std::vector<int>& wt_strides,
const std::vector<int>& wt_dilation,
const std::vector<int>& in_dilation,
int groups,
bool flip,
std::vector<array>& copies) {
// We will use the contiguous strides for the conv params because that is
// what the rest of the code expects.
constexpr int NDIM = 3;
int64_t in_arr_strides[NDIM + 2];
int64_t wt_arr_strides[NDIM + 2];
in_arr_strides[NDIM + 1] = wt_arr_strides[NDIM + 1] = 1;
for (int i = NDIM; i >= 0; i--) {
in_arr_strides[i] = in_arr_strides[i + 1] * in.shape(i + 1);
wt_arr_strides[i] = wt_arr_strides[i + 1] * wt.shape(i + 1);
}
// Make conv params
MLXConvParams<3> conv_params{
/* const int N = */ static_cast<int>(in.shape(0)),
@@ -1015,48 +1255,42 @@ void conv_3D_gpu(
/* const int idil[NDIM] = */
{in_dilation[0], in_dilation[1], in_dilation[2]},
/* const size_t in_strides[NDIM + 2] = */
{in.strides()[0],
in.strides()[1],
in.strides()[2],
in.strides()[3],
in.strides()[4]},
{in_arr_strides[0],
in_arr_strides[1],
in_arr_strides[2],
in_arr_strides[3],
in_arr_strides[4]},
/* const size_t wt_strides[NDIM + 2] = */
{wt.strides()[0],
wt.strides()[1],
wt.strides()[2],
wt.strides()[3],
wt.strides()[4]},
{wt_arr_strides[0],
wt_arr_strides[1],
wt_arr_strides[2],
wt_arr_strides[3],
wt_arr_strides[4]},
/* const size_t out_strides[NDIM + 2] = */
{out.strides()[0],
out.strides()[1],
out.strides()[2],
out.strides()[3],
out.strides()[4]},
/* const int groups = */ 1,
{out.strides(0),
out.strides(1),
out.strides(2),
out.strides(3),
out.strides(4)},
/* const int groups = */ groups,
/* const bool flip = */ flip,
};
return explicit_gemm_conv_ND_gpu(s, d, in, wt, out, conv_params);
return dispatch_conv_3D_gpu(s, d, in, wt, out, conv_params, copies);
}
} // namespace
void Convolution::eval_gpu(const std::vector<array>& inputs, array& out) {
out.set_data(allocator::malloc(out.nbytes()));
auto& s = stream();
auto& d = metal::device(s.device);
// Ensure contiguity
// Intermediates that are put here will be added to the command encoder as
// temporaries.
std::vector<array> copies;
auto in = inputs[0];
auto wt = inputs[1];
if (!in.flags().row_contiguous) {
in = contiguous_copy_gpu(in, s);
copies.push_back(in);
}
if (!wt.flags().row_contiguous) {
wt = contiguous_copy_gpu(wt, s);
copies.push_back(wt);
}
// Some shortcuts for brevity
const array& in = inputs[0];
const array& wt = inputs[1];
// 3D conv
if (out.ndim() == 5) {
@@ -1070,6 +1304,7 @@ void Convolution::eval_gpu(const std::vector<array>& inputs, array& out) {
kernel_strides_,
kernel_dilation_,
input_dilation_,
groups_,
flip_,
copies);
}
@@ -1112,7 +1347,9 @@ void Convolution::eval_gpu(const std::vector<array>& inputs, array& out) {
}
// Record copies
d.add_temporaries(std::move(copies), s.index);
if (!copies.empty()) {
d.add_temporaries(std::move(copies), s.index);
}
}
} // namespace mlx::core
+1
View File
@@ -41,6 +41,7 @@ const char* steel_gemm_gather();
const char* steel_gemm_segmented();
const char* conv();
const char* steel_conv();
const char* steel_conv_3d();
const char* steel_conv_general();
const char* gemv_masked();
const char* steel_attention();
+29
View File
@@ -770,6 +770,35 @@ MTL::ComputePipelineState* get_steel_conv_kernel(
return d.get_kernel(kernel_name, lib);
}
MTL::ComputePipelineState* get_steel_conv_3d_kernel(
metal::Device& d,
const std::string& kernel_name,
const array& out,
int bm,
int bn,
int bk,
int wm,
int wn,
bool small_filter) {
const auto& lib_name = kernel_name;
auto lib = d.get_library(lib_name, [&]() {
std::ostringstream kernel_source;
kernel_source << metal::utils() << metal::conv() << metal::steel_conv_3d()
<< get_template_definition(
lib_name,
"implicit_gemm_conv_3d",
get_type_string(out.dtype()),
bm,
bn,
bk,
wm,
wn,
small_filter);
return kernel_source.str();
});
return d.get_kernel(kernel_name, lib);
}
MTL::ComputePipelineState* get_steel_conv_general_kernel(
metal::Device& d,
const std::string& kernel_name,
+11
View File
@@ -201,6 +201,17 @@ MTL::ComputePipelineState* get_steel_conv_kernel(
int n_channel_specialization,
bool small_filter);
MTL::ComputePipelineState* get_steel_conv_3d_kernel(
metal::Device& d,
const std::string& kernel_name,
const array& out,
int bm,
int bn,
int bk,
int wm,
int wn,
bool small_filter);
MTL::ComputePipelineState* get_gemv_masked_kernel(
metal::Device& d,
const std::string& kernel_name,
+2
View File
@@ -67,6 +67,7 @@ set(STEEL_HEADERS
steel/conv/loaders/loader_channel_n.h
steel/conv/loaders/loader_general.h
steel/conv/kernels/steel_conv.h
steel/conv/kernels/steel_conv_3d.h
steel/conv/kernels/steel_conv_general.h
steel/gemm/gemm.h
steel/gemm/mma.h
@@ -143,6 +144,7 @@ if(NOT MLX_METAL_JIT)
build_kernel(ternary ternary.h ternary_ops.h)
build_kernel(unary unary.h unary_ops.h)
build_kernel(steel/conv/kernels/steel_conv ${STEEL_HEADERS})
build_kernel(steel/conv/kernels/steel_conv_3d ${STEEL_HEADERS})
build_kernel(steel/conv/kernels/steel_conv_general ${STEEL_HEADERS})
build_kernel(steel/gemm/kernels/steel_gemm_fused ${STEEL_HEADERS})
build_kernel(steel/gemm/kernels/steel_gemm_gather ${STEEL_HEADERS})
@@ -0,0 +1,135 @@
// Copyright © 2024 Apple Inc.
#include <metal_stdlib>
using namespace metal;
template <
typename T,
int BM,
int BN,
int BK,
int WM,
int WN,
bool SMALL_FILTER = false>
[[kernel, max_total_threads_per_threadgroup(WM * WN * 32)]] void
implicit_gemm_conv_3d(
const device T* A [[buffer(0)]],
const device T* B [[buffer(1)]],
device T* C [[buffer(2)]],
const constant MLXConvParams<3>* params [[buffer(3)]],
const constant ImplicitGemmConv3DParams* gemm_params [[buffer(4)]],
uint3 tid [[threadgroup_position_in_grid]],
uint3 lid [[thread_position_in_threadgroup]],
uint simd_gid [[simdgroup_index_in_threadgroup]],
uint simd_lid [[thread_index_in_simdgroup]]) {
using namespace mlx::steel;
(void)lid;
constexpr bool transpose_a = false;
constexpr bool transpose_b = true;
constexpr short tgp_padding_a = 16 / sizeof(T);
constexpr short tgp_padding_b = 16 / sizeof(T);
constexpr short shape_a_cols = (transpose_a ? BM : BK) + tgp_padding_a;
constexpr short shape_b_cols = (transpose_b ? BK : BN) + tgp_padding_b;
constexpr short shape_a_rows = (transpose_a ? BK : BM);
constexpr short shape_b_rows = (transpose_b ? BN : BK);
constexpr short tgp_mem_size_a = shape_a_cols * shape_a_rows;
constexpr short tgp_mem_size_b = shape_b_cols * shape_b_rows;
constexpr short tgp_size = WM * WN * 32;
// Input loader
using loader_a_t = typename metal::conditional_t<
// If the filter is small we can precompute masks for bounds checking
SMALL_FILTER,
Conv3DInputBlockLoaderSmallFilter<T, BM, BN, BK, tgp_size, tgp_padding_a>,
Conv3DInputBlockLoaderLargeFilter<
T,
BM,
BN,
BK,
tgp_size,
tgp_padding_a>>;
// Weight loader
using loader_b_t =
Conv3DWeightBlockLoader<T, BM, BN, BK, tgp_size, tgp_padding_b>;
using mma_t = BlockMMA<
T,
T,
BM,
BN,
BK,
WM,
WN,
transpose_a,
transpose_b,
shape_a_cols,
shape_b_cols>;
threadgroup T As[tgp_mem_size_a];
threadgroup T Bs[tgp_mem_size_b];
const int tid_y = ((tid.y) << gemm_params->swizzle_log) +
((tid.x) & ((1 << gemm_params->swizzle_log) - 1));
const int tid_x = (tid.x) >> gemm_params->swizzle_log;
if (gemm_params->tiles_n <= tid_x || gemm_params->tiles_m <= tid_y) {
return;
}
const int c_row = tid_y * BM;
const int c_col = tid_x * BN;
const int K = gemm_params->K;
const int N = gemm_params->N;
const int C_per_group = params->C / params->groups;
// Groups
A += tid.z * C_per_group;
B += tid.z * N * K;
C += tid.z * N;
B += c_col * K;
C += c_row * (N * params->groups) + c_col;
const int2 offsets_a(0, c_row);
const int2 offsets_b(0, c_col);
// Prepare threadgroup loading operations
loader_a_t loader_a(
A, As, offsets_a, params, gemm_params, simd_gid, simd_lid);
loader_b_t loader_b(
B, Bs, offsets_b, params, gemm_params, simd_gid, simd_lid);
// Prepare threadgroup mma operation
mma_t mma_op(simd_gid, simd_lid);
int gemm_k_iterations = gemm_params->gemm_k_iterations;
for (int k = 0; k < gemm_k_iterations; k++) {
threadgroup_barrier(mem_flags::mem_threadgroup);
// Load elements into threadgroup
loader_a.load_unsafe();
loader_b.load_unsafe();
threadgroup_barrier(mem_flags::mem_threadgroup);
// Multiply and accumulate threadgroup elements
mma_op.mma(As, Bs);
// Prepare for next iteration
loader_a.next();
loader_b.next();
}
threadgroup_barrier(mem_flags::mem_none);
// Store results to device memory
short tgp_bm = min(BM, gemm_params->M - c_row);
short tgp_bn = min(BN, gemm_params->N - c_col);
const int ldc = N * params->groups;
mma_op.store_result_safe(C, ldc, short2(tgp_bn, tgp_bm));
}
@@ -0,0 +1,48 @@
// Copyright © 2024 Apple Inc.
#include <metal_stdlib>
// clang-format off
#include "mlx/backend/metal/kernels/utils.h"
#include "mlx/backend/metal/kernels/steel/gemm/mma.h"
#include "mlx/backend/metal/kernels/steel/conv/conv.h"
#include "mlx/backend/metal/kernels/steel/conv/params.h"
#include "mlx/backend/metal/kernels/steel/conv/kernels/steel_conv_3d.h"
#define instantiate_implicit_conv_3d( \
name, \
itype, \
bm, \
bn, \
bk, \
wm, \
wn, \
fn, \
f) \
instantiate_kernel( \
"implicit_gemm_conv_3d_" #name "_bm" #bm "_bn" #bn \
"_bk" #bk "_wm" #wm "_wn" #wn "_filter_" #fn, \
implicit_gemm_conv_3d, \
itype, \
bm, \
bn, \
bk, \
wm, \
wn, \
f)
#define instantiate_implicit_conv_3d_filter(name, itype, bm, bn, bk, wm, wn) \
instantiate_implicit_conv_3d(name, itype, bm, bn, bk, wm, wn, s, true) \
instantiate_implicit_conv_3d(name, itype, bm, bn, bk, wm, wn, l, false)
#define instantiate_implicit_3d_blocks(name, itype) \
instantiate_implicit_conv_3d_filter(name, itype, 32, 8, 16, 4, 1) \
instantiate_implicit_conv_3d_filter(name, itype, 64, 8, 16, 4, 1) \
instantiate_implicit_conv_3d_filter(name, itype, 32, 32, 16, 2, 2) \
instantiate_implicit_conv_3d_filter(name, itype, 32, 64, 16, 2, 2) \
instantiate_implicit_conv_3d_filter(name, itype, 64, 32, 16, 2, 2) \
instantiate_implicit_conv_3d_filter(name, itype, 64, 64, 16, 2, 2)
instantiate_implicit_3d_blocks(float32, float);
instantiate_implicit_3d_blocks(float16, half);
instantiate_implicit_3d_blocks(bfloat16, bfloat16_t); // clang-format on
@@ -447,5 +447,509 @@ struct Conv2DWeightBlockLoader {
}
};
template <
typename T,
short BM,
short BN,
short BK,
short tgp_size,
short tgp_padding = 0>
struct Conv3DInputBlockLoaderLargeFilter {
// Destination dimensions
STEEL_CONST short BROWS = BM;
STEEL_CONST short BCOLS = BK;
// Read dimensions
STEEL_CONST short dst_ld = BCOLS + tgp_padding;
STEEL_CONST short vec_size = tgp_size / (BROWS * BCOLS) >= 8 ? 8 : 4;
// Thread read shape
STEEL_CONST short TCOLS = BCOLS / vec_size;
STEEL_CONST short TROWS = tgp_size / TCOLS;
// Rows / strided reads within the block
STEEL_CONST short n_rows = BROWS / TROWS;
// Thread location indices
const short thread_idx;
const short bi;
const short bj;
// threadgroup and device memory
threadgroup T* dst;
const constant MLXConvParams<3>* params;
const constant ImplicitGemmConv3DParams* gemm_params;
short weight_d;
short weight_h;
short weight_w;
short kdil_d;
short kdil_h;
short kdil_w;
const device T* src[n_rows];
int read_n[n_rows];
int read_id[n_rows];
int read_ih[n_rows];
int read_iw[n_rows];
/* Constructor */
METAL_FUNC Conv3DInputBlockLoaderLargeFilter(
const device T* src_,
threadgroup T* dst_,
const int2 offsets,
const constant MLXConvParams<3>* params_,
const constant ImplicitGemmConv3DParams* gemm_params_,
uint simd_group_id [[simdgroup_index_in_threadgroup]],
uint simd_lane_id [[thread_index_in_simdgroup]])
: thread_idx(simd_group_id * 32 + simd_lane_id),
bi(thread_idx / TCOLS),
bj(vec_size * (thread_idx % TCOLS)),
dst(dst_ + bi * dst_ld + bj),
params(params_),
gemm_params(gemm_params_),
weight_d(0),
weight_h(0),
weight_w(0),
kdil_d(params_->flip ? -params_->kdil[0] : params_->kdil[0]),
kdil_h(params_->flip ? -params_->kdil[1] : params_->kdil[1]),
kdil_w(params_->flip ? -params_->kdil[2] : params_->kdil[2]) {
int out_n_pixels = params->oS[0] * params->oS[1] * params->oS[2];
STEEL_PRAGMA_UNROLL
for (short i = 0; i < n_rows; ++i) {
int offset_ndhw = offsets.y + bi + i * TROWS;
int n = offset_ndhw / out_n_pixels;
int dhw = offset_ndhw % out_n_pixels;
int od = dhw / (params->oS[1] * params->oS[2]);
int hw = dhw % (params->oS[1] * params->oS[2]);
int oh = hw / params->oS[2];
int ow = hw % params->oS[2];
int id = od * params->str[0] - params->pad[0];
int ih = oh * params->str[1] - params->pad[1];
int iw = ow * params->str[2] - params->pad[2];
read_n[i] = n;
if (params->flip) {
read_id[i] = id + (params->wS[0] - 1) * params->kdil[0];
read_ih[i] = ih + (params->wS[1] - 1) * params->kdil[1];
read_iw[i] = iw + (params->wS[2] - 1) * params->kdil[2];
} else {
read_id[i] = id;
read_ih[i] = ih;
read_iw[i] = iw;
}
// Adjust for flip
if (params->flip) {
id += (params->wS[0] - 1) * params->kdil[0];
ih += (params->wS[1] - 1) * params->kdil[1];
iw += (params->wS[2] - 1) * params->kdil[2];
}
// Read from input if in bounds
src[i] = src_ + n * params->in_strides[0] + id * params->in_strides[1] +
ih * params->in_strides[2] + iw * params->in_strides[3] + bj;
}
}
/* Load from device memory into threadgroup memory - without bound checking */
METAL_FUNC void load_unsafe() const {
STEEL_PRAGMA_UNROLL
for (short i = 0, is = 0; i < n_rows; ++i, is += TROWS) {
// Find bounds
int n = read_n[i];
int id = read_id[i] + weight_d * kdil_d;
int ih = read_ih[i] + weight_h * kdil_h;
int iw = read_iw[i] + weight_w * kdil_w;
// Read from input if in bounds
if ((n < params->N) && (id >= 0 && id < params->iS[0]) &&
(ih >= 0 && ih < params->iS[1]) && (iw >= 0 && iw < params->iS[2])) {
STEEL_PRAGMA_UNROLL
for (short j = 0; j < vec_size; ++j) {
dst[is * dst_ld + j] = src[i][j];
}
}
// Zero pad otherwise
else {
STEEL_PRAGMA_UNROLL
for (short j = 0; j < vec_size; ++j) {
dst[is * dst_ld + j] = T(0);
}
}
}
}
/* Iteration helper */
METAL_FUNC void next() {
if (++weight_w < params->wS[2]) {
STEEL_PRAGMA_UNROLL
for (short i = 0; i < n_rows; i++) {
src[i] += gemm_params->inp_jump_w;
}
return;
}
weight_w = 0;
if (++weight_h < params->wS[1]) {
STEEL_PRAGMA_UNROLL
for (short i = 0; i < n_rows; i++) {
src[i] += gemm_params->inp_jump_h;
}
return;
}
weight_h = 0;
if (++weight_d < params->wS[0]) {
STEEL_PRAGMA_UNROLL
for (short i = 0; i < n_rows; i++) {
src[i] += gemm_params->inp_jump_d;
}
return;
}
weight_d = 0;
STEEL_PRAGMA_UNROLL
for (short i = 0; i < n_rows; i++) {
src[i] += gemm_params->inp_jump_c;
}
}
};
template <
typename T,
short BM,
short BN,
short BK,
short tgp_size,
short tgp_padding = 0>
struct Conv3DInputBlockLoaderSmallFilter {
// Destination dimensions
STEEL_CONST short BROWS = BM;
STEEL_CONST short BCOLS = BK;
// Read dimensions
STEEL_CONST short dst_ld = BCOLS + tgp_padding;
STEEL_CONST short vec_size = tgp_size / (BROWS * BCOLS) >= 8 ? 8 : 4;
// Thread read shape
STEEL_CONST short TCOLS = BCOLS / vec_size;
STEEL_CONST short TROWS = tgp_size / TCOLS;
// Rows / strided reads within the block
STEEL_CONST short n_rows = BROWS / TROWS;
using mask_t = short;
// Thread location indices
const short thread_idx;
const short bi;
const short bj;
// threadgroup and device memory
threadgroup T* dst;
const constant MLXConvParams<3>* params;
const constant ImplicitGemmConv3DParams* gemm_params;
short weight_d;
short weight_h;
short weight_w;
const device T* src[n_rows];
mask_t mask_d[n_rows];
mask_t mask_h[n_rows];
mask_t mask_w[n_rows];
/* Constructor */
METAL_FUNC Conv3DInputBlockLoaderSmallFilter(
const device T* src_,
threadgroup T* dst_,
const int2 offsets,
const constant MLXConvParams<3>* params_,
const constant ImplicitGemmConv3DParams* gemm_params_,
uint simd_group_id [[simdgroup_index_in_threadgroup]],
uint simd_lane_id [[thread_index_in_simdgroup]])
: thread_idx(simd_group_id * 32 + simd_lane_id),
bi(thread_idx / TCOLS),
bj(vec_size * (thread_idx % TCOLS)),
dst(dst_ + bi * dst_ld + bj),
params(params_),
gemm_params(gemm_params_),
weight_d(0),
weight_h(0),
weight_w(0) {
int out_n_pixels = params->oS[0] * params->oS[1] * params->oS[2];
int read_n[n_rows];
int read_id[n_rows];
int read_ih[n_rows];
int read_iw[n_rows];
STEEL_PRAGMA_UNROLL
for (short i = 0; i < n_rows; ++i) {
int offset_ndhw = offsets.y + bi + i * TROWS;
int n = offset_ndhw / out_n_pixels;
int dhw = offset_ndhw % out_n_pixels;
int od = dhw / (params->oS[1] * params->oS[2]);
int hw = dhw % (params->oS[1] * params->oS[2]);
int oh = hw / params->oS[2];
int ow = hw % params->oS[2];
int id = od * params->str[0] - params->pad[0];
int ih = oh * params->str[1] - params->pad[1];
int iw = ow * params->str[2] - params->pad[2];
read_n[i] = n;
read_id[i] = id;
read_ih[i] = ih;
read_iw[i] = iw;
// Adjust for flip
if (params->flip) {
id += (params->wS[0] - 1) * params->kdil[0];
ih += (params->wS[1] - 1) * params->kdil[1];
iw += (params->wS[2] - 1) * params->kdil[2];
}
// Read from input if in bounds
src[i] = src_ + n * params->in_strides[0] + id * params->in_strides[1] +
ih * params->in_strides[2] + iw * params->in_strides[3] + bj;
}
STEEL_PRAGMA_UNROLL
for (short i = 0; i < n_rows; ++i) {
mask_d[i] = 0;
mask_h[i] = 0;
mask_w[i] = 0;
}
for (short kd = 0; kd < params->wS[0]; kd++) {
short flip_d = params->flip ? params->wS[0] - kd - 1 : kd;
STEEL_PRAGMA_UNROLL
for (short i = 0; i < n_rows; ++i) {
int n = read_n[i];
int id = read_id[i] + flip_d * params->kdil[0];
bool in_bounds = n < params->N && id >= 0 && id < params->iS[0];
mask_d[i] |= (in_bounds << kd);
}
}
for (short kh = 0; kh < params->wS[1]; kh++) {
short flip_h = params->flip ? params->wS[1] - kh - 1 : kh;
STEEL_PRAGMA_UNROLL
for (short i = 0; i < n_rows; ++i) {
int ih = read_ih[i] + flip_h * params->kdil[1];
bool in_bounds = ih >= 0 && ih < params->iS[1];
mask_h[i] |= (in_bounds << kh);
}
}
for (short kw = 0; kw < params->wS[2]; kw++) {
short flip_w = params->flip ? params->wS[2] - kw - 1 : kw;
STEEL_PRAGMA_UNROLL
for (short i = 0; i < n_rows; ++i) {
int iw = read_iw[i] + flip_w * params->kdil[2];
bool in_bounds = iw >= 0 && iw < params->iS[2];
mask_w[i] |= (in_bounds << kw);
}
}
}
/* Load from device memory into threadgroup memory - without bound checking */
METAL_FUNC void load_unsafe() const {
mask_t d_mask = mask_t(1) << weight_d;
mask_t h_mask = mask_t(1) << weight_h;
mask_t w_mask = mask_t(1) << weight_w;
STEEL_PRAGMA_UNROLL
for (short i = 0, is = 0; i < n_rows; ++i, is += TROWS) {
// Read from input if in bounds
if ((mask_d[i] & d_mask) && (mask_h[i] & h_mask) &&
(mask_w[i] & w_mask)) {
STEEL_PRAGMA_UNROLL
for (short j = 0; j < vec_size; ++j) {
dst[is * dst_ld + j] = src[i][j];
}
}
// Zero pad otherwise
else {
STEEL_PRAGMA_UNROLL
for (short j = 0; j < vec_size; ++j) {
dst[is * dst_ld + j] = T(0);
}
}
}
}
/* Iteration helper */
METAL_FUNC void next() {
if (++weight_w < params->wS[2]) {
STEEL_PRAGMA_UNROLL
for (short i = 0; i < n_rows; i++) {
src[i] += gemm_params->inp_jump_w;
}
return;
}
weight_w = 0;
if (++weight_h < params->wS[1]) {
STEEL_PRAGMA_UNROLL
for (short i = 0; i < n_rows; i++) {
src[i] += gemm_params->inp_jump_h;
}
return;
}
weight_h = 0;
if (++weight_d < params->wS[0]) {
STEEL_PRAGMA_UNROLL
for (short i = 0; i < n_rows; i++) {
src[i] += gemm_params->inp_jump_d;
}
return;
}
weight_d = 0;
STEEL_PRAGMA_UNROLL
for (short i = 0; i < n_rows; i++) {
src[i] += gemm_params->inp_jump_c;
}
}
};
template <
typename T,
short BM,
short BN,
short BK,
short tgp_size,
short tgp_padding = 0>
struct Conv3DWeightBlockLoader {
// Destination dimensions
STEEL_CONST short BROWS = BN;
STEEL_CONST short BCOLS = BK;
// Read dimensions
STEEL_CONST short dst_ld = BCOLS + tgp_padding;
STEEL_CONST short vec_size =
(BN == 8) ? 1 : (tgp_size / (BROWS * BCOLS) >= 8 ? 8 : 4);
// Thread read shape
STEEL_CONST short TCOLS = BCOLS / vec_size;
STEEL_CONST short TROWS = tgp_size / TCOLS;
// Rows / strided reads within the block
STEEL_CONST short n_rows = BROWS / TROWS;
// Leading dimension for src
const int src_ld;
// Thread location indices
const short thread_idx;
const short bi;
const short bj;
// threadgroup and device memory
threadgroup T* dst;
const device T* src;
const constant MLXConvParams<3>* params;
int weight_dhw;
int weight_step;
const int read_n;
const bool do_read;
/* Constructor */
METAL_FUNC Conv3DWeightBlockLoader(
const device T* src_,
threadgroup T* dst_,
const int2 offsets,
const constant MLXConvParams<3>* params_,
const constant ImplicitGemmConv3DParams* gemm_params_,
uint simd_group_id [[simdgroup_index_in_threadgroup]],
uint simd_lane_id [[thread_index_in_simdgroup]])
: src_ld(params_->wt_strides[0]),
thread_idx(simd_group_id * 32 + simd_lane_id),
bi(thread_idx / TCOLS),
bj(vec_size * (thread_idx % TCOLS)),
dst(dst_ + bi * dst_ld + bj),
src(src_ + bi * src_ld + bj),
params(params_),
weight_dhw(0),
weight_step(params->C / params->groups),
read_n(offsets.y + bi),
do_read(read_n + n_rows * TROWS <= gemm_params_->N) {}
/* Load from device memory into threadgroup memory - without bound checking */
METAL_FUNC void load_unsafe() const {
if (BN != 8 || do_read) {
STEEL_PRAGMA_UNROLL
for (short i = 0; i < BN; i += TROWS) {
STEEL_PRAGMA_UNROLL
for (short j = 0; j < vec_size; j++) {
dst[i * dst_ld + j] = src[i * src_ld + j];
}
}
} else {
for (short i = 0; i < BN; i += TROWS) {
if ((read_n + i) < params->O) {
STEEL_PRAGMA_UNROLL
for (short j = 0; j < vec_size; j++) {
dst[i * dst_ld + j] = src[i * src_ld + j];
}
} else {
STEEL_PRAGMA_UNROLL
for (short j = 0; j < vec_size; j++) {
dst[i * dst_ld + j] = T(0);
}
}
}
}
}
/* Iteration helper */
METAL_FUNC void next() {
if (++weight_dhw < (params->wS[0] * params->wS[1] * params->wS[2])) {
src += weight_step;
return;
}
weight_dhw = 0;
src +=
BK - (params->wS[0] * params->wS[1] * params->wS[2] - 1) * weight_step;
}
};
} // namespace steel
} // namespace mlx
+56 -15
View File
@@ -4,21 +4,45 @@
template <int NDIM>
struct MLXConvParams {
const int N; // Batch size
const int C; // In channels
const int O; // Out channels
const int iS[NDIM]; // Input spatial dim
const int wS[NDIM]; // Weight spatial dim
const int oS[NDIM]; // Output spatial dim
const int str[NDIM]; // Kernel strides
const int pad[NDIM]; // Input padding
const int kdil[NDIM]; // Kernel dilation
const int idil[NDIM]; // Input dilation
const int64_t in_strides[NDIM + 2]; // In strides
const int64_t wt_strides[NDIM + 2]; // Wt strides
const int64_t out_strides[NDIM + 2]; // Out strides
const int groups; // Input channel groups
const bool flip;
int N; // Batch size
int C; // In channels
int O; // Out channels
int iS[NDIM]; // Input spatial dim
int wS[NDIM]; // Weight spatial dim
int oS[NDIM]; // Output spatial dim
int str[NDIM]; // Kernel strides
int pad[NDIM]; // Input padding
int kdil[NDIM]; // Kernel dilation
int idil[NDIM]; // Input dilation
int64_t in_strides[NDIM + 2]; // In strides
int64_t wt_strides[NDIM + 2]; // Wt strides
int64_t out_strides[NDIM + 2]; // Out strides
int groups; // Input channel groups
bool flip;
static MLXConvParams<NDIM>
with_padded_channels(MLXConvParams<NDIM> other, int pad_out, int pad_in) {
MLXConvParams<NDIM> params = other;
// Update strides
for (int i = 0; i < NDIM + 1; i++) {
params.in_strides[i] =
(params.in_strides[i] / params.C) * (params.C + pad_in);
params.wt_strides[i] =
(params.wt_strides[i] / params.C) * (params.C + pad_in);
params.out_strides[i] =
(params.out_strides[i] / params.O) * (params.O + pad_out);
}
params.in_strides[NDIM + 1] = 1;
params.wt_strides[NDIM + 1] = 1;
params.out_strides[NDIM + 1] = 1;
// Update channels
params.C += pad_in;
params.O += pad_out;
return params;
};
};
namespace mlx {
@@ -40,6 +64,23 @@ struct ImplicitGemmConv2DParams {
const int swizzle_log;
};
struct ImplicitGemmConv3DParams {
const int M;
const int N;
const int K;
const int gemm_k_iterations;
const int inp_jump_w;
const int inp_jump_h;
const int inp_jump_d;
const int inp_jump_c;
const int tiles_n;
const int tiles_m;
const int swizzle_log;
};
struct Conv2DGeneralJumpParams {
const int f_wgt_jump_h;
const int f_wgt_jump_w;
+13
View File
@@ -257,6 +257,19 @@ MTL::ComputePipelineState* get_steel_conv_kernel(
return d.get_kernel(kernel_name);
}
MTL::ComputePipelineState* get_steel_conv_3d_kernel(
metal::Device& d,
const std::string& kernel_name,
const array&,
int,
int,
int,
int,
int,
bool) {
return d.get_kernel(kernel_name);
}
MTL::ComputePipelineState* get_steel_conv_general_kernel(
metal::Device& d,
const std::string& kernel_name,
-1
View File
@@ -550,7 +550,6 @@ class TestConv(mlx_tests.MLXTestCase):
(1, 1, 6),
(4, 16, 32),
):
continue
for idim, kdim, stride, padding in (
((1, 1, 1), (1, 1, 1), (1, 1, 1), (0, 0, 0)),
((3, 3, 3), (3, 1, 1), (1, 1, 1), (0, 0, 0)),