This commit is contained in:
Georgi Gerganov 2024-01-24 10:16:05 +02:00
parent 035c4f01e6
commit 5cbdba693d
No known key found for this signature in database
GPG key ID: BF970631944C16B7
2 changed files with 21 additions and 22 deletions

View file

@ -2253,16 +2253,16 @@ static bool ggml_metal_graph_compute(
[encoder setBytes:&ne3 length:sizeof( int64_t) atIndex:26];
[encoder setBytes:&scale length:sizeof( float) atIndex:27];
const int64_t nwarps = 4;
const int64_t nhptg = 2; // heads per threadgroup !! sync with kernel template arguments !!
const int64_t nqptg = 4; // queries per threadgroup !! sync with kernel template arguments !!
const int64_t nsg = 4; // simdgroups per threadgroup (a.k.a. warps)
const int64_t nhptg = 2; // heads per threadgroup !! sync with kernel template arguments !!
const int64_t nqptg = 4; // queries per threadgroup !! sync with kernel template arguments !!
const size_t smem = nqptg*(nhptg*ne00 + nwarps*(nhptg*ne00 + 256))*(sizeof(float)/2);
const size_t smem = nqptg*(nhptg*ne00 + nsg*(nhptg*ne00 + 256))*(sizeof(float)/2);
GGML_ASSERT(smem <= ctx->device.maxThreadgroupMemoryLength);
[encoder setThreadgroupMemoryLength:smem atIndex:0];
[encoder dispatchThreadgroups:MTLSizeMake((ne01 + nqptg - 1)/nqptg, (ne02 + nhptg - 1)/(nhptg), ne03) threadsPerThreadgroup:MTLSizeMake(32, nwarps, 1)];
[encoder dispatchThreadgroups:MTLSizeMake((ne01 + nqptg - 1)/nqptg, (ne02 + nhptg - 1)/(nhptg), ne03) threadsPerThreadgroup:MTLSizeMake(32, nsg, 1)];
} break;
case GGML_OP_DUP:
case GGML_OP_CPY: