From 39e95698543b9d78ad0152a61c2b36b43cc4e25c Mon Sep 17 00:00:00 2001 From: ZYP Date: Sat, 12 Sep 2026 23:55:18 +0800 Subject: [PATCH] Split long non-causal SDPA queries into blocks MPSGraph scaledDotProductAttention can silently return garbage on M2/M3 Ultra above ~15k query rows (#16). Chunk non-causal queries (default block 2048, only when sequence > 12288). Tune with --sdpa-query-block / env H3_SDPA_MAX_QUERY_ROWS and H3_SDPA_SPLIT_MIN_ROWS; set MAX to 0 to disable. --- h3_gpu.m | 50 ++++++++++++++++++++++++++++++++++++++++++++++++-- main.c | 13 ++++++++++++- 2 files changed, 60 insertions(+), 3 deletions(-) diff --git a/h3_gpu.m b/h3_gpu.m index c61d04c6..af64e1cb 100644 --- a/h3_gpu.m +++ b/h3_gpu.m @@ -1445,17 +1445,43 @@ int h3_gpu_qkv_rope_f32(h3_gpu *opaque, h3_gpu_tensor *query, }); } + +/* MPSGraph native SDPA silently corrupts output on M2/M3 Ultra above ~15k + * query rows (antirez/h3.c#16 / Apple FB24605554). Split non-causal queries + * into blocks. Defaults: split when sequence > 12288, block size 2048. + * H3_SDPA_MAX_QUERY_ROWS=0 disables splitting. */ +static unsigned h3_gpu_env_u32(const char *name, unsigned fallback) { + const char *text = getenv(name); + if (!text || !*text) return fallback; + char *end = NULL; + unsigned long value = strtoul(text, &end, 10); + if (end == text || *end) return fallback; + if (value > UINT32_MAX) return fallback; + return (unsigned)value; +} + +static unsigned h3_gpu_sdpa_query_block(void) { + /* 0 disables query splitting. */ + return h3_gpu_env_u32("H3_SDPA_MAX_QUERY_ROWS", 2048u); +} + +static unsigned h3_gpu_sdpa_split_min_rows(void) { + return h3_gpu_env_u32("H3_SDPA_SPLIT_MIN_ROWS", 12288u); +} + static H3SDPA *h3_gpu_sdpa_graph(H3GPU *gpu, uint32_t batch, uint32_t sequence, uint32_t heads, uint32_t head_dim, float scale, MPSDataType dataType, int causal, int headMajor, int outputHeadMajor) { @autoreleasepool { + unsigned query_block = h3_gpu_sdpa_query_block(); + unsigned split_min = h3_gpu_sdpa_split_min_rows(); NSString *cacheKey = [NSString stringWithFormat: - @"%u:%u:%u:%u:%u:%.9g:%d:%d:%d", + @"%u:%u:%u:%u:%u:%.9g:%d:%d:%d:%u:%u", (unsigned)dataType, batch, sequence, heads, head_dim, scale, causal, headMajor, - outputHeadMajor]; + outputHeadMajor, query_block, split_min]; H3SDPA *cached = gpu.sdpaCache[cacheKey]; if (cached) return cached; MPSGraph *graph = [[MPSGraph alloc] init]; @@ -1500,6 +1526,26 @@ int h3_gpu_qkv_rope_f32(h3_gpu *opaque, h3_gpu_tensor *query, attention = [graph scaledDotProductAttentionWithQueryTensor:qt keyTensor:kt valueTensor:vt maskTensor:mask scale:scale name:nil]; + } else if (query_block && sequence > split_min && + sequence > query_block) { + /* Non-causal rows are independent; chunk queries so each + * MPSGraph SDPA op stays under the Ultra silent-corruption + * threshold (see antirez/h3.c#16). */ + NSMutableArray *parts = + [NSMutableArray array]; + for (uint32_t start = 0; start < sequence; start += query_block) { + uint32_t rows = query_block; + if (start + rows > sequence) rows = sequence - start; + MPSGraphTensor *q_rows = [graph sliceTensor:qt + dimension:2 + start:start + length:rows + name:nil]; + [parts addObject:[graph + scaledDotProductAttentionWithQueryTensor:q_rows + keyTensor:kt valueTensor:vt scale:scale name:nil]]; + } + attention = [graph concatTensors:parts dimension:2 name:nil]; } else { attention = [graph scaledDotProductAttentionWithQueryTensor:qt keyTensor:kt valueTensor:vt scale:scale name:nil]; diff --git a/main.c b/main.c index 7f11e470..7070eabb 100644 --- a/main.c +++ b/main.c @@ -251,7 +251,8 @@ int main(int argc, char **argv) { OPT_FIRST, OPT_LAST, OPT_REF_IMAGE, OPT_REF_IMAGE_SIZE, OPT_REF_VIDEO, OPT_REF_SILENT_VIDEO, OPT_REF_VIDEO_AUDIO, OPT_REF_AUDIO, OPT_FRAMES_DIR, OPT_SHOW, OPT_ZOOM, - OPT_PROFILE, OPT_INFO }; + OPT_PROFILE, OPT_INFO, + OPT_SDPA_QUERY_BLOCK }; static const struct option options[] = { {"model-dir", required_argument, NULL, 'd'}, {"prompt", required_argument, NULL, 'p'}, @@ -268,6 +269,7 @@ int main(int argc, char **argv) { {"core-reuse", required_argument, NULL, OPT_CORE_REUSE}, {"token-reduction", no_argument, NULL, OPT_TOKEN_REDUCTION}, {"ssd-streaming", no_argument, NULL, OPT_SSD_STREAMING}, + {"sdpa-query-block", required_argument, NULL, OPT_SDPA_QUERY_BLOCK}, {"use-int8-row-fc2", no_argument, NULL, OPT_USE_INT8_ROW_FC2}, {"use-reference-rope", no_argument, NULL, OPT_USE_REFERENCE_ROPE}, {"use-slower-bf16-mlp", no_argument, NULL, @@ -355,6 +357,15 @@ int main(int argc, char **argv) { break; case OPT_TOKEN_REDUCTION: params.token_reduction = 1; break; case OPT_SSD_STREAMING: params.ssd_streaming = 1; break; + case OPT_SDPA_QUERY_BLOCK: { + /* Sets H3_SDPA_MAX_QUERY_ROWS for the Ultra SDPA workaround (#16). */ + int rows = parse_int(optarg, "sdpa query block"); + if (rows < 0) return 1; + char buf[32]; + snprintf(buf, sizeof(buf), "%d", rows); + setenv("H3_SDPA_MAX_QUERY_ROWS", buf, 1); + break; + } case OPT_USE_INT8_ROW_FC2: params.use_int8_row_fc2 = 1; break;