Skip to content
Open
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
50 changes: 48 additions & 2 deletions h3_gpu.m
Original file line number Diff line number Diff line change
Expand Up @@ -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];
Expand Down Expand Up @@ -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<MPSGraphTensor *> *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];
Expand Down
13 changes: 12 additions & 1 deletion main.c
Original file line number Diff line number Diff line change
Expand Up @@ -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'},
Expand All @@ -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,
Expand Down Expand Up @@ -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;
Expand Down