Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 9 additions & 5 deletions external/ggml/src/ggml-metal/ggml-metal-device.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -725,7 +725,7 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv_ext(ggml_
char name[256];

const ggml_type tsrc0 = op->src[0]->type;
const ggml_type tsrc1 = op->src[1]->type;
const ggml_type tsrc1 = op->src[1]->type;
const int ne12 = op->src[1]->ne[2];
const int r2 = ne12 / op->src[0]->ne[2];
const int r3 = op->src[1]->ne[3] / op->src[0]->ne[3];
Expand Down Expand Up @@ -753,13 +753,16 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv_ext(ggml_
return res;
}

ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mm(ggml_metal_library_t lib, const ggml_tensor * op) {
ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mm(ggml_metal_library_t lib, const ggml_tensor * op) {
char base[256];
char name[256];

const ggml_type tsrc0 = op->src[0]->type;
const ggml_type tsrc1 = op->src[1]->type;

const bool full_f32 = tsrc0 == GGML_TYPE_F32 && tsrc1 == GGML_TYPE_F32 &&
ggml_get_op_params_i32(op, 0) == GGML_PREC_F32;

const bool bc_inp = op->src[0]->ne[0] % 32 != 0;

constexpr int NRA = SZ_SIMDGROUP * N_MM_BLOCK_Y * N_MM_SIMD_GROUP_Y;
Expand All @@ -777,7 +780,8 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mm(ggml_meta
const int16_t r2 = (int16_t) (ne12 / op->src[0]->ne[2]);
const int16_t r3 = (int16_t) (ne13 / op->src[0]->ne[3]);

snprintf(base, 256, "kernel_mul_mm_%s_%s", ggml_type_name(tsrc0), ggml_type_name(tsrc1));
snprintf(base, 256, "kernel_mul_mm_%s_%s%s", ggml_type_name(tsrc0), ggml_type_name(tsrc1),
full_f32 ? "_prec_f32" : "");
snprintf(name, 256, "%s_bci=%d_bco=%d_ne12=%d_ne13=%d_r2=%d_r3=%d",
base, bc_inp, bc_out, ne12, ne13, r2, r3);

Expand All @@ -801,13 +805,13 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mm(ggml_meta
res.nr0 = NRA;
res.nr1 = NRB;

const size_t smem_a = NRA * N_MM_NK_TOTAL * sizeof(ggml_fp16_t);
const size_t smem_a = NRA * N_MM_NK_TOTAL * (full_f32 ? sizeof(float) : sizeof(ggml_fp16_t));
res.smem = smem_a;
} else {
res.nr0 = 64;
res.nr1 = 32;

res.smem = bc_out ? 8192 : (4096 + 2048);
res.smem = full_f32 ? (8192 + 4096) : (bc_out ? 8192 : (4096 + 2048));
}

res.nsg = N_MM_SIMD_GROUP_X * N_MM_SIMD_GROUP_Y;
Expand Down
10 changes: 9 additions & 1 deletion external/ggml/src/ggml-metal/ggml-metal-ops.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -2371,7 +2371,15 @@ int ggml_metal_op_mul_mat(ggml_metal_op_t ctx, int idx) {
!ggml_is_transposed(op->src[1]) &&
// for now the matrix-matrix multiplication kernel only works on A14+/M1+ SoCs
// AMD GPU and older A-chips will reuse matrix-vector multiplication kernel
props_dev->has_simdgroup_mm && ne00 >= 64 && ne11 > ne11_mm_min) {
// Short F32 contractions (e.g. 32-dim QK attention) still benefit
// from tiled MM when both output axes are large enough. Keep the
// existing MV choice for small outputs and other precision modes.
props_dev->has_simdgroup_mm &&
(ne00 >= 64 || (ne00 >= 32 && ne01 >= 64 && ne11 >= 32 &&
op->src[0]->type == GGML_TYPE_F32 &&
op->src[1]->type == GGML_TYPE_F32 &&
ggml_get_op_params_i32(op, 0) == GGML_PREC_F32)) &&
ne11 > ne11_mm_min) {
//GGML_LOG_INFO("matrix: ne00 = %6d, ne01 = %6d, ne02 = %6d, ne11 = %6d, ne12 = %6d\n", ne00, ne01, ne02, ne11, ne12);

// some Metal matrix data types require aligned pointers
Expand Down
11 changes: 9 additions & 2 deletions external/ggml/src/ggml-metal/ggml-metal.metal
Original file line number Diff line number Diff line change
Expand Up @@ -7913,6 +7913,11 @@ kernel void kernel_cpy_t_t(
const int i01 = ntg[1] == 1 ? tgpig[0]%args.ne01 : tgpig[0]*ntg[1] + tiitg/ntg[0];
const int iw0 = ntg[1] == 1 ? tgpig[0]/args.ne01 : 0;

// Padded rows must not cross channel strides or the destination boundary.
if (i01 >= args.ne01) {
return;
}

const int64_t n = i03*args.ne02*args.ne01*args.ne00 + i02*args.ne01*args.ne00 + i01*args.ne00;

const int64_t i3 = n/(args.ne2*args.ne1*args.ne0);
Expand Down Expand Up @@ -10262,7 +10267,7 @@ kernel void kernel_mul_mm(
ushort sgitg[[simdgroup_index_in_threadgroup]]) {

threadgroup S0 * sa = (threadgroup S0 *)(shmem);
threadgroup S1 * sb = (threadgroup S1 *)(shmem + 4096);
threadgroup S1 * sb = (threadgroup S1 *)(shmem + 64*32*sizeof(S0));

constexpr int NR0 = 64;
constexpr int NR1 = 32;
Expand Down Expand Up @@ -10909,7 +10914,9 @@ template [[host_name("kernel_set_rows_iq4_nl_i32")]] kernel set_rows_q32_t kerne

typedef decltype(kernel_mul_mm<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, float4x4, 1, dequantize_f32, float, float4x4, float, float2x4>) mul_mm_t;

template [[host_name("kernel_mul_mm_f32_f32")]] kernel mul_mm_t kernel_mul_mm<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, float4x4, 1, dequantize_f32, float, float4x4, float, float2x4>;
template [[host_name("kernel_mul_mm_f32_f32")]] kernel mul_mm_t kernel_mul_mm<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, float4x4, 1, dequantize_f32, float, float4x4, float, float2x4>;
// GGML_PREC_F32 must preserve F32 operands instead of staging them as half.
template [[host_name("kernel_mul_mm_f32_f32_prec_f32")]] kernel mul_mm_t kernel_mul_mm<float, float4x4, simdgroup_float8x8, float, float2x4, simdgroup_float8x8, float4x4, 1, dequantize_f32, float, float4x4, float, float2x4>;
template [[host_name("kernel_mul_mm_f16_f32")]] kernel mul_mm_t kernel_mul_mm<half, half4x4, simdgroup_half8x8, half, half2x4, simdgroup_half8x8, half4x4, 1, dequantize_f16, half, half4x4, float, float2x4>;
#if defined(GGML_METAL_HAS_BF16)
template [[host_name("kernel_mul_mm_bf16_f32")]] kernel mul_mm_t kernel_mul_mm<bfloat, bfloat4x4, simdgroup_bfloat8x8, bfloat, bfloat2x4, simdgroup_bfloat8x8, bfloat4x4, 1, dequantize_bf16, bfloat, bfloat4x4, float, float2x4>;
Expand Down
14 changes: 14 additions & 0 deletions external/ggml/tests/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -354,3 +354,17 @@ if (NOT GGML_BACKEND_DL)
add_test(NAME ${TEST_TARGET} COMMAND $<TARGET_FILE:${TEST_TARGET}>)
set_property(TEST ${TEST_TARGET} PROPERTY ENVIRONMENT "LLVM_PROFILE_FILE=${TEST_TARGET}.profraw")
endif()

if (GGML_METAL)
add_executable(test-metal-copy-bounds test-metal-copy-bounds.cpp)
target_link_libraries(test-metal-copy-bounds PRIVATE ggml)
add_test(NAME test-metal-copy-bounds COMMAND test-metal-copy-bounds)
set_tests_properties(test-metal-copy-bounds PROPERTIES SKIP_RETURN_CODE 77)
endif()

if (GGML_METAL)
add_executable(test-metal-f32-matmul test-metal-f32-matmul.cpp)
target_link_libraries(test-metal-f32-matmul PRIVATE ggml)
add_test(NAME test-metal-f32-matmul COMMAND test-metal-f32-matmul)
set_tests_properties(test-metal-f32-matmul PROPERTIES SKIP_RETURN_CODE 77)
endif()
48 changes: 48 additions & 0 deletions external/ggml/tests/test-metal-copy-bounds.cpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,48 @@
#include "ggml.h"
#include "ggml-backend.h"
#include "ggml-metal.h"
#include <vector>
#include <cmath>
#include <cstdio>
#include <chrono>
#include <algorithm>
#include <stdexcept>
struct Case {
ggml_context *ctx=ggml_init({32*1024*1024,nullptr,true});
ggml_backend_buffer_t buf=nullptr;
ggml_cgraph *g=nullptr;
~Case(){if(buf)ggml_backend_buffer_free(buf);ggml_free(ctx);}
void alloc(ggml_backend_t b,ggml_tensor*out){g=ggml_new_graph_custom(ctx,2048,false);ggml_build_forward_expand(g,out);buf=ggml_backend_alloc_ctx_tensors(ctx,b);if(!buf)throw std::runtime_error("alloc");}
void run(ggml_backend_t b){if(ggml_backend_graph_compute(b,g)!=GGML_STATUS_SUCCESS)throw std::runtime_error("compute");ggml_backend_synchronize(b);}
};
void copy(ggml_backend_t backend,int K,int R){
Case c;const int C=5,B=2,guard=8192;int count=K*R*C*B;
// Permute [C,K,R,B] into [K,R,C,B], a real non-contiguous ggml view.
auto *storage=ggml_new_tensor_1d(c.ctx,GGML_TYPE_F32,count+guard);
auto *source=ggml_view_4d(c.ctx,storage,C,K,R,B,C*4,C*K*4,C*K*R*4,0);
source=ggml_permute(c.ctx,source,2,0,1,3);
auto *dest=ggml_new_tensor_1d(c.ctx,GGML_TYPE_F32,count+guard);
auto *view=ggml_view_4d(c.ctx,dest,K,R,C,B,K*4,K*R*4,K*R*C*4,0);
auto *out=ggml_cpy(c.ctx,source,view);c.alloc(backend,out);
std::vector<float>in(count+guard),init(count+guard,-1234567.f),got(count+guard);
for(size_t i=0;i<in.size();i++)in[i]=float(i+1);
ggml_backend_tensor_set(storage,in.data(),0,in.size()*4);
int worst=0,tail=0;
for(int t=0;t<5;t++){
ggml_backend_tensor_set(dest,init.data(),0,init.size()*4);c.run(backend);ggml_backend_tensor_get(dest,got.data(),0,got.size()*4);
int bad=0,over=0;
for(int b=0;b<B;b++)for(int ch=0;ch<C;ch++)for(int r=0;r<R;r++)for(int k=0;k<K;k++)
bad+=got[((b*C+ch)*R+r)*K+k]!=in[((b*R+r)*K+k)*C+ch];
for(int i=count;i<count+guard;i++)over+=got[i]!=init[i];
worst=std::max(worst,bad);tail=std::max(tail,over);
}
printf("COPY K=%d rows=%d channels=%d batch=%d wrong=%d guard_overwrites=%d\n",K,R,C,B,worst,tail);fflush(stdout);
if(worst || tail)throw std::runtime_error("copy output or guard mismatch");
}
int main() {
auto b=ggml_backend_metal_init();if(!b)return 77;
try {
for(int k:{7,32,48,64,129,256})for(int r:{1,3,7,16})copy(b,k,r);
} catch (const std::exception & e) { std::fprintf(stderr,"%s\n",e.what());ggml_backend_free(b);return 1; }
ggml_backend_free(b);
}
52 changes: 52 additions & 0 deletions external/ggml/tests/test-metal-f32-matmul.cpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,52 @@
#include "ggml.h"
#include "ggml-backend.h"
#include "ggml-metal.h"
#include <vector>
#include <cmath>
#include <cstdio>
#include <chrono>
#include <algorithm>
#include <stdexcept>
struct Case {
ggml_context *ctx=ggml_init({32*1024*1024,nullptr,true});
ggml_backend_buffer_t buf=nullptr;
ggml_cgraph *g=nullptr;
~Case(){if(buf)ggml_backend_buffer_free(buf);ggml_free(ctx);}
void alloc(ggml_backend_t b,ggml_tensor*out){g=ggml_new_graph_custom(ctx,2048,false);ggml_build_forward_expand(g,out);buf=ggml_backend_alloc_ctx_tensors(ctx,b);if(!buf)throw std::runtime_error("alloc");}
void run(ggml_backend_t b){if(ggml_backend_graph_compute(b,g)!=GGML_STATUS_SUCCESS)throw std::runtime_error("compute");ggml_backend_synchronize(b);}
};
void mm(ggml_backend_t backend,int K,int M,int N,int broadcast,int mode,bool f32,int reps=5,ggml_type dtype=GGML_TYPE_F32){
Case c; int B=2,BA=broadcast?1:B;
auto*a=ggml_new_tensor_3d(c.ctx,dtype,K,M,BA);auto*b=ggml_new_tensor_3d(c.ctx,GGML_TYPE_F32,K,N,B);
auto*out=ggml_mul_mat(c.ctx,a,b);if(f32)ggml_mul_mat_set_prec(out,GGML_PREC_F32);
c.alloc(backend,out);
std::vector<float> av(K*M*BA),bv(K*N*B),got(M*N*B);
for(size_t i=0;i<av.size();i++)av[i]=mode==1?70000.f+float(i%31):mode==2?1.f+float(i%23)*0.00001f:std::sin(float(i)*0.013f);
for(size_t i=0;i<bv.size();i++)bv[i]=mode==2?(i%K%2?-1.f:1.f):std::cos(float(i)*0.017f);
if(dtype==GGML_TYPE_F16){std::vector<ggml_fp16_t>h(av.size());ggml_fp32_to_fp16_row(av.data(),h.data(),h.size());ggml_backend_tensor_set(a,h.data(),0,h.size()*2);ggml_fp16_to_fp32_row(h.data(),av.data(),h.size());}
else ggml_backend_tensor_set(a,av.data(),0,av.size()*4);
ggml_backend_tensor_set(b,bv.data(),0,bv.size()*4);c.run(backend);
auto t0=std::chrono::steady_clock::now();for(int r=0;r<reps;r++)c.run(backend);
double ms=std::chrono::duration<double,std::milli>(std::chrono::steady_clock::now()-t0).count()/reps;
ggml_backend_tensor_get(out,got.data(),0,got.size()*4);
double maxerr=0,se=0,ss=0,maxscaled=0;int nonfinite=0;
for(int z=0;z<B;z++)for(int n=0;n<N;n++)for(int m=0;m<M;m++){
double ref=0,sumabs=0;for(int k=0;k<K;k++){double v=double(av[((broadcast?0:z)*M+m)*K+k])*bv[(z*N+n)*K+k];ref+=v;sumabs+=std::abs(v);}
float v=got[(z*N+n)*M+m];if(!std::isfinite(v)){nonfinite++;continue;}
double e=std::abs(v-ref);maxerr=std::max(maxerr,e);maxscaled=std::max(maxscaled,e/(1+sumabs));se+=e*e;ss+=ref*ref;
}
printf("MM K=%d M=%d N=%d bc=%d mode=%d prec=%s dtype=%s nonfinite=%d maxerr=%.9g relrmse=%.9g scaled=%.9g ms=%.4f\n",K,M,N,broadcast,mode,f32?"f32":"default",ggml_type_name(dtype),nonfinite,maxerr,std::sqrt(se/(ss+1e-30)),maxscaled,ms);fflush(stdout);
if(nonfinite || (f32 && dtype==GGML_TYPE_F32 && maxscaled>2e-6))
throw std::runtime_error("explicit F32 matmul precision regression");
}
int main() {
auto b=ggml_backend_metal_init();if(!b)return 77;
try {
for(int k:{31,32,33,48,63,64,65,128})for(int bc:{0,1})mm(b,k,67,35,bc,0,true);
for(int mode:{1,2})for(int k:{32,48,64,128})mm(b,k,67,35,0,mode,true);
for(int m:{63,64,67})for(int n:{8,9,31,32,35})mm(b,32,m,n,0,0,true);
for(int k:{32,48})mm(b,k,512,512,0,0,true,20);
for(bool prec:{false,true}){mm(b,128,67,35,0,0,prec);mm(b,128,67,35,0,0,prec,5,GGML_TYPE_F16);}
} catch (const std::exception & e) {std::fprintf(stderr,"%s\n",e.what());ggml_backend_free(b);return 1;}
ggml_backend_free(b);
}
Loading