diff --git a/Makefile b/Makefile index 9711dc1a4..36b5a2738 100644 --- a/Makefile +++ b/Makefile @@ -217,6 +217,12 @@ ds4_rocm.o: ds4_rocm.cu ds4_gpu.h ds4_iq2_tables_cuda.inc $(ROCM_SRCS) tests/cuda_long_context_smoke: tests/cuda_long_context_smoke.o ds4_cuda.o $(NVCC) $(NVCCFLAGS) -o $@ $^ $(CUDA_LDLIBS) +tests/two_sessions_smoke.o: tests/two_sessions_smoke.c ds4.h + $(CC) $(CFLAGS) -I. -c -o $@ tests/two_sessions_smoke.c + +tests/two_sessions_smoke: tests/two_sessions_smoke.o $(CORE_OBJS) + $(DS4_LINK) -o $@ $^ $(DS4_LINK_LIBS) + ds4_test: ds4_test.o ds4_help.o ds4_kvstore.o rax.o $(CORE_OBJS) ifeq ($(UNAME_S),Darwin) $(CC) $(CFLAGS) -o $@ ds4_test.o ds4_help.o ds4_kvstore.o rax.o $(CORE_OBJS) $(METAL_LDLIBS) diff --git a/ds4.c b/ds4.c index 640511eb0..26ac783ac 100644 --- a/ds4.c +++ b/ds4.c @@ -10309,10 +10309,22 @@ static void print_vec_stats(const char *name, const float *x, uint64_t n) { enum { DS4_STREAMING_PREFILL_CACHE_SEED_MAX_TOKENS = 64 }; +/* Shared per-engine work buffers. Everything in here is written and + * consumed inside a single ds4_session_sync()/ds4_session_eval() call: + * no content survives across calls, so sessions that share one engine can + * also share one scratch as long as they are driven from one thread at a + * time (the same serialization the server's single graph worker already + * provides). Persistent KV caches, compressor/indexer frontiers and MTP + * verification state stay in ds4_gpu_graph: they belong to one session's + * timeline. Sharing the scratch is what makes a second live session + * affordable: the prefill batch buffers dominate a session's non-KV + * footprint (~1.8 MiB per prefill-chunk token). */ typedef struct { - /* One-token decode tensors. These stay allocated for the life of a - * session; a generated token enters as an embedding in cur_hc and leaves as - * logits after all 43 layers update their raw/compressed/indexer caches. */ + uint32_t cap_prefill; /* prefill_cap the buffers were sized for */ + uint32_t cap_comp; /* comp_cap ceiling across sessions */ + uint32_t cap_attn_comp_stage; + bool mtp_ready; /* MTP work tensors are allocated */ + int refcount; /* live sessions attached to this scratch */ ds4_gpu_tensor *cur_hc; ds4_gpu_tensor *flat_hc; ds4_gpu_tensor *hc_mix; @@ -10327,51 +10339,7 @@ typedef struct { ds4_gpu_tensor *q; ds4_gpu_tensor *kv_raw; ds4_gpu_tensor *kv; - - /* Persistent KV state. Raw KV is a sliding-window ring per layer. Ratio-4 - * layers also keep an indexer-compressed cache; ratio-128 layers keep only - * the attention-compressed cache. The small state tensors are compressor - * frontiers for the next compressed row, so they must be snapshotted with - * the row counters whenever a checkpoint is saved or partially rewound. */ - ds4_gpu_tensor *layer_raw_cache[DS4_MAX_LAYER]; - ds4_gpu_tensor *layer_attn_comp_cache[DS4_MAX_LAYER]; - ds4_gpu_tensor *layer_attn_state_kv[DS4_MAX_LAYER]; - ds4_gpu_tensor *layer_attn_state_score[DS4_MAX_LAYER]; - ds4_gpu_tensor *layer_index_comp_cache[DS4_MAX_LAYER]; - ds4_gpu_tensor *layer_index_state_kv[DS4_MAX_LAYER]; - ds4_gpu_tensor *layer_index_state_score[DS4_MAX_LAYER]; - - /* Speculative decoding scratch. MTP is allowed to mutate graph state only - * if the target verifier can either commit it or restore the saved - * frontiers. The prefix1 buffers are the cheap partial-accept state for the - * common N=2 case. */ - ds4_gpu_tensor *spec_attn_state_kv[DS4_MAX_LAYER]; - ds4_gpu_tensor *spec_attn_state_score[DS4_MAX_LAYER]; - ds4_gpu_tensor *spec_index_state_kv[DS4_MAX_LAYER]; - ds4_gpu_tensor *spec_index_state_score[DS4_MAX_LAYER]; - ds4_gpu_tensor *spec_prefix1_attn_state_kv[DS4_MAX_LAYER]; - ds4_gpu_tensor *spec_prefix1_attn_state_score[DS4_MAX_LAYER]; - ds4_gpu_tensor *spec_prefix1_index_state_kv[DS4_MAX_LAYER]; - ds4_gpu_tensor *spec_prefix1_index_state_score[DS4_MAX_LAYER]; ds4_gpu_tensor *spec_logits; - uint32_t layer_n_comp[DS4_MAX_LAYER]; - uint32_t layer_n_index_comp[DS4_MAX_LAYER]; - uint32_t spec_prefix1_n_comp[DS4_MAX_LAYER]; - uint32_t spec_prefix1_n_index_comp[DS4_MAX_LAYER]; - bool spec_capture_prefix1; - uint32_t raw_cap; - /* Maximum compressed-row capacity across layers. Shared work buffers use - * this worst-case size because ratio-4 indexer layers can still reach it. */ - uint32_t comp_cap; - /* Persistent compressed caches are per layer, so size them from the actual - * layer compression ratio instead of pessimistically using the ratio-4 cap - * for every ratio-128 layer. */ - uint32_t layer_comp_cap[DS4_MAX_LAYER]; - uint32_t attn_comp_stage_cap; - - /* Per-layer work tensors. They are reused in place by every layer instead - * of allocating a generic graph arena. This is why the code is verbose but - * predictable: each pointer names an actual DS4 stage. */ ds4_gpu_tensor *comp_kv_cur; ds4_gpu_tensor *comp_sc_cur; ds4_gpu_tensor *attn_comp_stage; @@ -10406,10 +10374,6 @@ typedef struct { ds4_gpu_tensor *output_embd; ds4_gpu_tensor *output_norm; ds4_gpu_tensor *logits; - - /* Optional MTP model state. It has its own raw cache because the drafter - * runs on speculative future tokens; target KV state is updated only after - * verification accepts draft tokens. */ ds4_gpu_tensor *mtp_embed; ds4_gpu_tensor *mtp_enorm; ds4_gpu_tensor *mtp_eproj; @@ -10419,15 +10383,6 @@ typedef struct { ds4_gpu_tensor *mtp_input_hc; ds4_gpu_tensor *mtp_state_hc; ds4_gpu_tensor *mtp_next_hc; - ds4_gpu_tensor *mtp_raw_cache; - uint32_t mtp_n_raw; - uint32_t prefill_cap; - uint32_t raw_window; - - /* Batched prefill tensors. Prefill is layer-major: a chunk of prompt - * tokens moves through layer 0, then layer 1, and so on, updating the same - * persistent caches used by decode. Keeping this separate from decode - * avoids a slow loop of one-token graph steps for long prompts. */ ds4_gpu_tensor *prefill_tokens; ds4_gpu_tensor *batch_cur_hc; ds4_gpu_tensor *batch_next_hc; @@ -10464,6 +10419,79 @@ typedef struct { ds4_gpu_tensor *batch_router_weights; ds4_gpu_tensor *prefill_seed_router_selected; uint32_t prefill_seed_tokens; + ds4_gpu_tensor *batch_routed_gate; + ds4_gpu_tensor *batch_routed_up; + ds4_gpu_tensor *batch_routed_mid; + ds4_gpu_tensor *batch_routed_down; + ds4_gpu_tensor *batch_routed_out; + bool batch_routed_mid_is_f16; + bool materialize_ffn_out; + float *cpu_router_norm; +} ds4_gpu_scratch; + +typedef struct { + /* One-token decode tensors. These stay allocated for the life of a + * session; a generated token enters as an embedding in cur_hc and leaves as + * logits after all 43 layers update their raw/compressed/indexer caches. */ + + /* Persistent KV state. Raw KV is a sliding-window ring per layer. Ratio-4 + * layers also keep an indexer-compressed cache; ratio-128 layers keep only + * the attention-compressed cache. The small state tensors are compressor + * frontiers for the next compressed row, so they must be snapshotted with + * the row counters whenever a checkpoint is saved or partially rewound. */ + ds4_gpu_tensor *layer_raw_cache[DS4_MAX_LAYER]; + ds4_gpu_tensor *layer_attn_comp_cache[DS4_MAX_LAYER]; + ds4_gpu_tensor *layer_attn_state_kv[DS4_MAX_LAYER]; + ds4_gpu_tensor *layer_attn_state_score[DS4_MAX_LAYER]; + ds4_gpu_tensor *layer_index_comp_cache[DS4_MAX_LAYER]; + ds4_gpu_tensor *layer_index_state_kv[DS4_MAX_LAYER]; + ds4_gpu_tensor *layer_index_state_score[DS4_MAX_LAYER]; + + /* Speculative decoding scratch. MTP is allowed to mutate graph state only + * if the target verifier can either commit it or restore the saved + * frontiers. The prefix1 buffers are the cheap partial-accept state for the + * common N=2 case. */ + ds4_gpu_tensor *spec_attn_state_kv[DS4_MAX_LAYER]; + ds4_gpu_tensor *spec_attn_state_score[DS4_MAX_LAYER]; + ds4_gpu_tensor *spec_index_state_kv[DS4_MAX_LAYER]; + ds4_gpu_tensor *spec_index_state_score[DS4_MAX_LAYER]; + ds4_gpu_tensor *spec_prefix1_attn_state_kv[DS4_MAX_LAYER]; + ds4_gpu_tensor *spec_prefix1_attn_state_score[DS4_MAX_LAYER]; + ds4_gpu_tensor *spec_prefix1_index_state_kv[DS4_MAX_LAYER]; + ds4_gpu_tensor *spec_prefix1_index_state_score[DS4_MAX_LAYER]; + uint32_t layer_n_comp[DS4_MAX_LAYER]; + uint32_t layer_n_index_comp[DS4_MAX_LAYER]; + uint32_t spec_prefix1_n_comp[DS4_MAX_LAYER]; + uint32_t spec_prefix1_n_index_comp[DS4_MAX_LAYER]; + bool spec_capture_prefix1; + /* Engine-owned shared work buffers (see ds4_gpu_scratch). */ + ds4_gpu_scratch *scr; + uint32_t raw_cap; + /* Maximum compressed-row capacity across layers. Shared work buffers use + * this worst-case size because ratio-4 indexer layers can still reach it. */ + uint32_t comp_cap; + /* Persistent compressed caches are per layer, so size them from the actual + * layer compression ratio instead of pessimistically using the ratio-4 cap + * for every ratio-128 layer. */ + uint32_t layer_comp_cap[DS4_MAX_LAYER]; + uint32_t attn_comp_stage_cap; + + /* Per-layer work tensors. They are reused in place by every layer instead + * of allocating a generic graph arena. This is why the code is verbose but + * predictable: each pointer names an actual DS4 stage. */ + + /* Optional MTP model state. It has its own raw cache because the drafter + * runs on speculative future tokens; target KV state is updated only after + * verification accepts draft tokens. */ + ds4_gpu_tensor *mtp_raw_cache; + uint32_t mtp_n_raw; + uint32_t prefill_cap; + uint32_t raw_window; + + /* Batched prefill tensors. Prefill is layer-major: a chunk of prompt + * tokens moves through layer 0, then layer 1, and so on, updating the same + * persistent caches used by decode. Keeping this separate from decode + * avoids a slow loop of one-token graph steps for long prompts. */ uint64_t prefill_selected_profile_rows; uint64_t prefill_selected_profile_unique; uint64_t prefill_selected_profile_selected_bytes; @@ -10471,14 +10499,7 @@ typedef struct { uint32_t prefill_selected_profile_layers; uint32_t prefill_selected_profile_min_unique; uint32_t prefill_selected_profile_max_unique; - ds4_gpu_tensor *batch_routed_gate; - ds4_gpu_tensor *batch_routed_up; - ds4_gpu_tensor *batch_routed_mid; - ds4_gpu_tensor *batch_routed_down; - ds4_gpu_tensor *batch_routed_out; - bool batch_routed_mid_is_f16; ds4_gpu_tensor *batch_ffn_out; - bool materialize_ffn_out; ds4_gpu_tensor *directional_steering_dirs; float directional_steering_attn_scale; float directional_steering_ffn_scale; @@ -10491,7 +10512,6 @@ typedef struct { bool ssd_streaming_cold; bool streaming_static_decode_map_current; bool mtp_enabled; - float *cpu_router_norm; } ds4_gpu_graph; static bool graph_power_throttle_enabled(const ds4_gpu_graph *g) { @@ -10533,93 +10553,12 @@ static void graph_power_note_decode_token(ds4_gpu_graph *g, double elapsed_sec) /* Release every Metal tensor owned by the whole-model graph runtime. */ static void metal_graph_free(ds4_gpu_graph *g) { + /* The shared scratch is engine-owned: just detach from it. The last + * engine close frees it. */ + if (g->scr && g->scr->refcount > 0) g->scr->refcount--; ds4_gpu_tensor_free(g->directional_steering_dirs); ds4_gpu_tensor_free(g->batch_ffn_out); - ds4_gpu_tensor_free(g->batch_routed_out); - ds4_gpu_tensor_free(g->batch_routed_down); - ds4_gpu_tensor_free(g->batch_routed_mid); - ds4_gpu_tensor_free(g->batch_routed_up); - ds4_gpu_tensor_free(g->batch_routed_gate); - ds4_gpu_tensor_free(g->batch_router_weights); - ds4_gpu_tensor_free(g->prefill_seed_router_selected); - ds4_gpu_tensor_free(g->batch_router_selected); - ds4_gpu_tensor_free(g->batch_router_probs); - ds4_gpu_tensor_free(g->batch_router_logits); - ds4_gpu_tensor_free(g->batch_shared_out); - ds4_gpu_tensor_free(g->batch_shared_mid); - ds4_gpu_tensor_free(g->batch_shared_up); - ds4_gpu_tensor_free(g->batch_shared_gate); - ds4_gpu_tensor_free(g->batch_ffn_norm); - ds4_gpu_tensor_free(g->batch_ffn_cur); - ds4_gpu_tensor_free(g->batch_after_attn_hc); - ds4_gpu_tensor_free(g->batch_low_tmp); - ds4_gpu_tensor_free(g->batch_group_tmp); - ds4_gpu_tensor_free(g->batch_attn_out); - ds4_gpu_tensor_free(g->batch_attn_low); - ds4_gpu_tensor_free(g->batch_heads); - ds4_gpu_tensor_free(g->batch_indexer_weights); - ds4_gpu_tensor_free(g->batch_indexer_q); - ds4_gpu_tensor_free(g->batch_comp_sc); - ds4_gpu_tensor_free(g->batch_comp_kv); - ds4_gpu_tensor_free(g->batch_kv); - ds4_gpu_tensor_free(g->batch_kv_raw); - ds4_gpu_tensor_free(g->batch_q_half); - ds4_gpu_tensor_free(g->batch_q); - ds4_gpu_tensor_free(g->batch_qr_norm); - ds4_gpu_tensor_free(g->batch_qr); - ds4_gpu_tensor_free(g->batch_attn_norm); - ds4_gpu_tensor_free(g->batch_attn_cur); - ds4_gpu_tensor_free(g->batch_hc_split); - ds4_gpu_tensor_free(g->batch_hc_mix); - ds4_gpu_tensor_free(g->batch_flat_hc); - ds4_gpu_tensor_free(g->batch_next_hc); - ds4_gpu_tensor_free(g->batch_cur_hc); - ds4_gpu_tensor_free(g->prefill_tokens); - ds4_gpu_tensor_free(g->logits); ds4_gpu_tensor_free(g->mtp_raw_cache); - ds4_gpu_tensor_free(g->mtp_next_hc); - ds4_gpu_tensor_free(g->mtp_state_hc); - ds4_gpu_tensor_free(g->mtp_input_hc); - ds4_gpu_tensor_free(g->mtp_hproj_hc); - ds4_gpu_tensor_free(g->mtp_hnorm_hc); - ds4_gpu_tensor_free(g->mtp_eproj_hc); - ds4_gpu_tensor_free(g->mtp_eproj); - ds4_gpu_tensor_free(g->mtp_enorm); - ds4_gpu_tensor_free(g->mtp_embed); - ds4_gpu_tensor_free(g->spec_logits); - ds4_gpu_tensor_free(g->output_norm); - ds4_gpu_tensor_free(g->output_embd); - ds4_gpu_tensor_free(g->output_weights); - ds4_gpu_tensor_free(g->output_pre); - ds4_gpu_tensor_free(g->after_ffn_hc); - ds4_gpu_tensor_free(g->ffn_out); - ds4_gpu_tensor_free(g->routed_out); - ds4_gpu_tensor_free(g->routed_down); - ds4_gpu_tensor_free(g->routed_mid); - ds4_gpu_tensor_free(g->routed_up); - ds4_gpu_tensor_free(g->routed_gate); - ds4_gpu_tensor_free(g->router_weights); - ds4_gpu_tensor_free(g->router_selected); - ds4_gpu_tensor_free(g->router_probs); - ds4_gpu_tensor_free(g->router_logits); - ds4_gpu_tensor_free(g->shared_out); - ds4_gpu_tensor_free(g->shared_mid); - ds4_gpu_tensor_free(g->shared_up); - ds4_gpu_tensor_free(g->shared_gate); - ds4_gpu_tensor_free(g->ffn_norm); - ds4_gpu_tensor_free(g->ffn_cur); - ds4_gpu_tensor_free(g->after_attn_hc); - ds4_gpu_tensor_free(g->attn_out); - ds4_gpu_tensor_free(g->attn_low); - ds4_gpu_tensor_free(g->heads); - ds4_gpu_tensor_free(g->comp_sc_cur); - ds4_gpu_tensor_free(g->comp_kv_cur); - ds4_gpu_tensor_free(g->attn_comp_stage); - ds4_gpu_tensor_free(g->comp_mask); - ds4_gpu_tensor_free(g->comp_selected); - ds4_gpu_tensor_free(g->indexer_scores); - ds4_gpu_tensor_free(g->indexer_weights); - ds4_gpu_tensor_free(g->indexer_q); for (uint32_t il = 0; il < DS4_N_LAYER; il++) { ds4_gpu_tensor_free(g->layer_raw_cache[il]); } @@ -10651,21 +10590,6 @@ static void metal_graph_free(ds4_gpu_graph *g) { ds4_gpu_tensor_free(g->spec_prefix1_index_state_kv[il]); ds4_gpu_tensor_free(g->spec_prefix1_index_state_score[il]); } - ds4_gpu_tensor_free(g->kv); - ds4_gpu_tensor_free(g->kv_raw); - ds4_gpu_tensor_free(g->q); - ds4_gpu_tensor_free(g->qr_norm); - ds4_gpu_tensor_free(g->qr); - ds4_gpu_tensor_free(g->attn_norm); - ds4_gpu_tensor_free(g->attn_cur); - ds4_gpu_tensor_free(g->hc_comb); - ds4_gpu_tensor_free(g->hc_post); - ds4_gpu_tensor_free(g->hc_pre); - ds4_gpu_tensor_free(g->hc_split); - ds4_gpu_tensor_free(g->hc_mix); - ds4_gpu_tensor_free(g->flat_hc); - ds4_gpu_tensor_free(g->cur_hc); - free(g->cpu_router_norm); memset(g, 0, sizeof(*g)); } @@ -10930,15 +10854,15 @@ static void metal_graph_debug_dump_i32_tensor( static bool metal_graph_needs_ffn_out(const ds4_gpu_graph *g, uint32_t il, uint32_t pos) { return metal_graph_directional_steering_ffn_enabled(g) || - g->materialize_ffn_out || + g->scr->materialize_ffn_out || metal_graph_debug_wants("ffn_out", il, pos); } static bool metal_graph_ensure_ffn_out(ds4_gpu_graph *g) { - if (!g->ffn_out) { - g->ffn_out = ds4_gpu_tensor_alloc((uint64_t)DS4_N_EMBD * sizeof(float)); + if (!g->scr->ffn_out) { + g->scr->ffn_out = ds4_gpu_tensor_alloc((uint64_t)DS4_N_EMBD * sizeof(float)); } - return g->ffn_out != NULL; + return g->scr->ffn_out != NULL; } static bool metal_graph_ensure_batch_ffn_out(ds4_gpu_graph *g) { @@ -10952,10 +10876,333 @@ static bool metal_graph_ensure_batch_ffn_out(ds4_gpu_graph *g) { * Metal Release Graph Allocation. * ========================================================================= */ +static void metal_scratch_free(ds4_gpu_scratch *scr) { + if (!scr) return; + const int refs = scr->refcount; + ds4_gpu_tensor_free(scr->batch_routed_out); + ds4_gpu_tensor_free(scr->batch_routed_down); + ds4_gpu_tensor_free(scr->batch_routed_mid); + ds4_gpu_tensor_free(scr->batch_routed_up); + ds4_gpu_tensor_free(scr->batch_routed_gate); + ds4_gpu_tensor_free(scr->batch_router_weights); + ds4_gpu_tensor_free(scr->prefill_seed_router_selected); + ds4_gpu_tensor_free(scr->batch_router_selected); + ds4_gpu_tensor_free(scr->batch_router_probs); + ds4_gpu_tensor_free(scr->batch_router_logits); + ds4_gpu_tensor_free(scr->batch_shared_out); + ds4_gpu_tensor_free(scr->batch_shared_mid); + ds4_gpu_tensor_free(scr->batch_shared_up); + ds4_gpu_tensor_free(scr->batch_shared_gate); + ds4_gpu_tensor_free(scr->batch_ffn_norm); + ds4_gpu_tensor_free(scr->batch_ffn_cur); + ds4_gpu_tensor_free(scr->batch_after_attn_hc); + ds4_gpu_tensor_free(scr->batch_low_tmp); + ds4_gpu_tensor_free(scr->batch_group_tmp); + ds4_gpu_tensor_free(scr->batch_attn_out); + ds4_gpu_tensor_free(scr->batch_attn_low); + ds4_gpu_tensor_free(scr->batch_heads); + ds4_gpu_tensor_free(scr->batch_indexer_weights); + ds4_gpu_tensor_free(scr->batch_indexer_q); + ds4_gpu_tensor_free(scr->batch_comp_sc); + ds4_gpu_tensor_free(scr->batch_comp_kv); + ds4_gpu_tensor_free(scr->batch_kv); + ds4_gpu_tensor_free(scr->batch_kv_raw); + ds4_gpu_tensor_free(scr->batch_q_half); + ds4_gpu_tensor_free(scr->batch_q); + ds4_gpu_tensor_free(scr->batch_qr_norm); + ds4_gpu_tensor_free(scr->batch_qr); + ds4_gpu_tensor_free(scr->batch_attn_norm); + ds4_gpu_tensor_free(scr->batch_attn_cur); + ds4_gpu_tensor_free(scr->batch_hc_split); + ds4_gpu_tensor_free(scr->batch_hc_mix); + ds4_gpu_tensor_free(scr->batch_flat_hc); + ds4_gpu_tensor_free(scr->batch_next_hc); + ds4_gpu_tensor_free(scr->batch_cur_hc); + ds4_gpu_tensor_free(scr->prefill_tokens); + ds4_gpu_tensor_free(scr->logits); + ds4_gpu_tensor_free(scr->mtp_next_hc); + ds4_gpu_tensor_free(scr->mtp_state_hc); + ds4_gpu_tensor_free(scr->mtp_input_hc); + ds4_gpu_tensor_free(scr->mtp_hproj_hc); + ds4_gpu_tensor_free(scr->mtp_hnorm_hc); + ds4_gpu_tensor_free(scr->mtp_eproj_hc); + ds4_gpu_tensor_free(scr->mtp_eproj); + ds4_gpu_tensor_free(scr->mtp_enorm); + ds4_gpu_tensor_free(scr->mtp_embed); + ds4_gpu_tensor_free(scr->spec_logits); + ds4_gpu_tensor_free(scr->output_norm); + ds4_gpu_tensor_free(scr->output_embd); + ds4_gpu_tensor_free(scr->output_weights); + ds4_gpu_tensor_free(scr->output_pre); + ds4_gpu_tensor_free(scr->after_ffn_hc); + ds4_gpu_tensor_free(scr->ffn_out); + ds4_gpu_tensor_free(scr->routed_out); + ds4_gpu_tensor_free(scr->routed_down); + ds4_gpu_tensor_free(scr->routed_mid); + ds4_gpu_tensor_free(scr->routed_up); + ds4_gpu_tensor_free(scr->routed_gate); + ds4_gpu_tensor_free(scr->router_weights); + ds4_gpu_tensor_free(scr->router_selected); + ds4_gpu_tensor_free(scr->router_probs); + ds4_gpu_tensor_free(scr->router_logits); + ds4_gpu_tensor_free(scr->shared_out); + ds4_gpu_tensor_free(scr->shared_mid); + ds4_gpu_tensor_free(scr->shared_up); + ds4_gpu_tensor_free(scr->shared_gate); + ds4_gpu_tensor_free(scr->ffn_norm); + ds4_gpu_tensor_free(scr->ffn_cur); + ds4_gpu_tensor_free(scr->after_attn_hc); + ds4_gpu_tensor_free(scr->attn_out); + ds4_gpu_tensor_free(scr->attn_low); + ds4_gpu_tensor_free(scr->heads); + ds4_gpu_tensor_free(scr->comp_sc_cur); + ds4_gpu_tensor_free(scr->comp_kv_cur); + ds4_gpu_tensor_free(scr->attn_comp_stage); + ds4_gpu_tensor_free(scr->comp_mask); + ds4_gpu_tensor_free(scr->comp_selected); + ds4_gpu_tensor_free(scr->indexer_scores); + ds4_gpu_tensor_free(scr->indexer_weights); + ds4_gpu_tensor_free(scr->indexer_q); + ds4_gpu_tensor_free(scr->kv); + ds4_gpu_tensor_free(scr->kv_raw); + ds4_gpu_tensor_free(scr->q); + ds4_gpu_tensor_free(scr->qr_norm); + ds4_gpu_tensor_free(scr->qr); + ds4_gpu_tensor_free(scr->attn_norm); + ds4_gpu_tensor_free(scr->attn_cur); + ds4_gpu_tensor_free(scr->hc_comb); + ds4_gpu_tensor_free(scr->hc_post); + ds4_gpu_tensor_free(scr->hc_pre); + ds4_gpu_tensor_free(scr->hc_split); + ds4_gpu_tensor_free(scr->hc_mix); + ds4_gpu_tensor_free(scr->flat_hc); + ds4_gpu_tensor_free(scr->cur_hc); + free(scr->cpu_router_norm); + memset(scr, 0, sizeof(*scr)); + scr->refcount = refs; +} + +/* Allocate the shared work buffers for the requested capacities. Sizing + * mirrors what the graph allocator used to do when every session owned a + * private copy of these tensors. */ +static bool metal_scratch_alloc( + ds4_gpu_scratch *scr, + const ds4_weights *weights, + const ds4_layer_weights *layer, + uint32_t comp_cap, + uint32_t attn_comp_stage_cap, + uint32_t prefill_cap, + bool enable_mtp) { + const int refs = scr->refcount; + memset(scr, 0, sizeof(*scr)); + scr->refcount = refs; + scr->cap_prefill = prefill_cap; + scr->cap_comp = comp_cap; + scr->cap_attn_comp_stage = attn_comp_stage_cap; + scr->mtp_ready = enable_mtp; + + const uint64_t hc_dim = (uint64_t)DS4_N_HC * DS4_N_EMBD; + const uint64_t mix_hc = 2ull * DS4_N_HC + (uint64_t)DS4_N_HC * DS4_N_HC; + const uint64_t q_rank = layer->attn_q_a->dim[1]; + const uint64_t q_dim = (uint64_t)DS4_N_HEAD * DS4_N_HEAD_DIM; + const uint64_t low_dim = (uint64_t)DS4_N_OUT_GROUP * DS4_N_LORA_O; + const uint64_t group_dim = (uint64_t)DS4_N_HEAD_DIM * (DS4_N_HEAD / DS4_N_OUT_GROUP); + const uint64_t shared_dim = layer->ffn_gate_shexp->dim[1]; + const uint64_t routed_mid_dim = layer->ffn_gate_exps->dim[1]; + const uint64_t vocab_dim = weights->output ? weights->output->dim[1] : DS4_N_VOCAB; + const uint64_t comp_width_max = 2ull * (DS4_N_HEAD_DIM > DS4_N_INDEXER_HEAD_DIM + ? DS4_N_HEAD_DIM + : DS4_N_INDEXER_HEAD_DIM); + const uint64_t indexer_q_dim = (uint64_t)DS4_N_INDEXER_HEAD * DS4_N_INDEXER_HEAD_DIM; + const uint64_t pc = prefill_cap; + + scr->cur_hc = ds4_gpu_tensor_alloc(hc_dim * sizeof(float)); + scr->flat_hc = ds4_gpu_tensor_alloc(hc_dim * sizeof(float)); + scr->hc_mix = ds4_gpu_tensor_alloc(mix_hc * sizeof(float)); + scr->hc_split = ds4_gpu_tensor_alloc(mix_hc * sizeof(float)); + scr->hc_pre = ds4_gpu_tensor_view(scr->hc_split, 0, (uint64_t)DS4_N_HC * sizeof(float)); + scr->hc_post = ds4_gpu_tensor_view(scr->hc_split, + (uint64_t)DS4_N_HC * sizeof(float), + (uint64_t)DS4_N_HC * sizeof(float)); + scr->hc_comb = ds4_gpu_tensor_view(scr->hc_split, + 2ull * DS4_N_HC * sizeof(float), + (uint64_t)DS4_N_HC * DS4_N_HC * sizeof(float)); + scr->attn_cur = ds4_gpu_tensor_alloc((uint64_t)DS4_N_EMBD * sizeof(float)); + scr->attn_norm = ds4_gpu_tensor_alloc((uint64_t)DS4_N_EMBD * sizeof(float)); + scr->qr = ds4_gpu_tensor_alloc(q_rank * sizeof(float)); + scr->qr_norm = ds4_gpu_tensor_alloc(q_rank * sizeof(float)); + scr->q = ds4_gpu_tensor_alloc(q_dim * sizeof(float)); + scr->kv_raw = ds4_gpu_tensor_alloc((uint64_t)DS4_N_HEAD_DIM * sizeof(float)); + scr->kv = ds4_gpu_tensor_alloc((uint64_t)DS4_N_HEAD_DIM * sizeof(float)); + scr->comp_kv_cur = ds4_gpu_tensor_alloc(comp_width_max * sizeof(float)); + scr->comp_sc_cur = ds4_gpu_tensor_alloc(comp_width_max * sizeof(float)); + scr->indexer_q = ds4_gpu_tensor_alloc(indexer_q_dim * sizeof(float)); + scr->indexer_weights = ds4_gpu_tensor_alloc((uint64_t)DS4_N_INDEXER_HEAD * sizeof(float)); + scr->indexer_scores = ds4_gpu_tensor_alloc((uint64_t)comp_cap * pc * sizeof(float)); + scr->comp_mask = ds4_gpu_tensor_alloc((uint64_t)comp_cap * pc * sizeof(float)); + scr->comp_selected = ds4_gpu_tensor_alloc((uint64_t)(DS4_N_INDEXER_TOP_K ? DS4_N_INDEXER_TOP_K : 1u) * + pc * sizeof(uint32_t)); + scr->heads = ds4_gpu_tensor_alloc(q_dim * sizeof(float)); + scr->attn_low = ds4_gpu_tensor_alloc(low_dim * sizeof(float)); + scr->attn_out = ds4_gpu_tensor_alloc((uint64_t)DS4_N_EMBD * sizeof(float)); + scr->after_attn_hc = ds4_gpu_tensor_alloc(hc_dim * sizeof(float)); + scr->ffn_cur = ds4_gpu_tensor_alloc((uint64_t)DS4_N_EMBD * sizeof(float)); + scr->ffn_norm = ds4_gpu_tensor_alloc((uint64_t)DS4_N_EMBD * sizeof(float)); + scr->cpu_router_norm = xmalloc((size_t)DS4_N_EMBD * sizeof(scr->cpu_router_norm[0])); + scr->shared_gate = ds4_gpu_tensor_alloc(shared_dim * sizeof(float)); + scr->shared_up = ds4_gpu_tensor_alloc(shared_dim * sizeof(float)); + scr->shared_mid = ds4_gpu_tensor_alloc(shared_dim * sizeof(float)); + scr->shared_out = ds4_gpu_tensor_alloc((uint64_t)DS4_N_EMBD * sizeof(float)); + scr->router_logits = ds4_gpu_tensor_alloc(DS4_N_EXPERT * sizeof(float)); + scr->router_probs = ds4_gpu_tensor_alloc(DS4_N_EXPERT * sizeof(float)); + scr->router_selected = ds4_gpu_tensor_alloc(DS4_N_EXPERT_USED * sizeof(int)); + scr->router_weights = ds4_gpu_tensor_alloc(DS4_N_EXPERT_USED * sizeof(float)); + scr->routed_gate = ds4_gpu_tensor_alloc((uint64_t)DS4_N_EXPERT_USED * routed_mid_dim * sizeof(float)); + scr->routed_up = ds4_gpu_tensor_alloc((uint64_t)DS4_N_EXPERT_USED * routed_mid_dim * sizeof(float)); + scr->routed_mid = ds4_gpu_tensor_alloc((uint64_t)DS4_N_EXPERT_USED * routed_mid_dim * sizeof(float)); + scr->routed_down = ds4_gpu_tensor_alloc((uint64_t)DS4_N_EXPERT_USED * DS4_N_EMBD * sizeof(float)); + scr->routed_out = ds4_gpu_tensor_alloc((uint64_t)DS4_N_EMBD * sizeof(float)); + scr->after_ffn_hc = ds4_gpu_tensor_alloc(hc_dim * sizeof(float)); + scr->output_pre = ds4_gpu_tensor_alloc((uint64_t)DS4_N_HC * sizeof(float)); + scr->output_weights = ds4_gpu_tensor_alloc((uint64_t)DS4_N_HC * sizeof(float)); + scr->output_embd = ds4_gpu_tensor_alloc((uint64_t)DS4_N_EMBD * sizeof(float)); + scr->output_norm = ds4_gpu_tensor_alloc((uint64_t)DS4_N_EMBD * sizeof(float)); + scr->logits = ds4_gpu_tensor_alloc(vocab_dim * sizeof(float)); + scr->prefill_tokens = ds4_gpu_tensor_alloc(pc * sizeof(int32_t)); + scr->batch_cur_hc = ds4_gpu_tensor_alloc(pc * hc_dim * sizeof(float)); + scr->batch_next_hc = ds4_gpu_tensor_alloc(pc * hc_dim * sizeof(float)); + scr->batch_flat_hc = ds4_gpu_tensor_alloc(pc * hc_dim * sizeof(float)); + scr->batch_hc_mix = ds4_gpu_tensor_alloc(pc * mix_hc * sizeof(float)); + scr->batch_hc_split = ds4_gpu_tensor_alloc(pc * mix_hc * sizeof(float)); + scr->batch_attn_cur = ds4_gpu_tensor_alloc(pc * DS4_N_EMBD * sizeof(float)); + scr->batch_attn_norm = ds4_gpu_tensor_alloc(pc * DS4_N_EMBD * sizeof(float)); + scr->batch_qr = ds4_gpu_tensor_alloc(pc * q_rank * sizeof(float)); + scr->batch_qr_norm = ds4_gpu_tensor_alloc(pc * q_rank * sizeof(float)); + scr->batch_q = ds4_gpu_tensor_alloc(pc * q_dim * sizeof(float)); + scr->batch_q_half = ds4_gpu_tensor_alloc(pc * q_dim * sizeof(uint16_t)); + scr->batch_kv_raw = ds4_gpu_tensor_alloc(pc * DS4_N_HEAD_DIM * sizeof(float)); + scr->batch_kv = ds4_gpu_tensor_alloc(pc * DS4_N_HEAD_DIM * sizeof(float)); + scr->batch_comp_kv = ds4_gpu_tensor_alloc(pc * comp_width_max * sizeof(float)); + scr->batch_comp_sc = ds4_gpu_tensor_alloc(pc * comp_width_max * sizeof(float)); + scr->batch_indexer_q = ds4_gpu_tensor_alloc(pc * indexer_q_dim * sizeof(float)); + scr->batch_indexer_weights = ds4_gpu_tensor_alloc(pc * DS4_N_INDEXER_HEAD * sizeof(float)); + scr->batch_heads = ds4_gpu_tensor_alloc(pc * q_dim * sizeof(float)); + scr->batch_attn_low = ds4_gpu_tensor_alloc(pc * low_dim * sizeof(float)); + scr->batch_attn_out = ds4_gpu_tensor_alloc(pc * DS4_N_EMBD * sizeof(float)); + scr->batch_group_tmp = ds4_gpu_tensor_alloc(pc * group_dim * sizeof(float)); + scr->batch_low_tmp = ds4_gpu_tensor_alloc(pc * DS4_N_LORA_O * sizeof(float)); + scr->batch_after_attn_hc = ds4_gpu_tensor_alloc(pc * hc_dim * sizeof(float)); + scr->batch_ffn_cur = ds4_gpu_tensor_alloc(pc * DS4_N_EMBD * sizeof(float)); + scr->batch_ffn_norm = ds4_gpu_tensor_alloc(pc * DS4_N_EMBD * sizeof(float)); + scr->batch_shared_gate = ds4_gpu_tensor_alloc(pc * shared_dim * sizeof(float)); + scr->batch_shared_up = ds4_gpu_tensor_alloc(pc * shared_dim * sizeof(float)); + scr->batch_shared_mid = ds4_gpu_tensor_alloc(pc * shared_dim * sizeof(float)); + scr->batch_shared_out = ds4_gpu_tensor_alloc(pc * DS4_N_EMBD * sizeof(float)); + scr->batch_router_logits = ds4_gpu_tensor_alloc(pc * DS4_N_EXPERT * sizeof(float)); + scr->batch_router_probs = ds4_gpu_tensor_alloc(pc * DS4_N_EXPERT * sizeof(float)); + scr->batch_router_selected = ds4_gpu_tensor_alloc(pc * DS4_N_EXPERT_USED * sizeof(int)); + scr->batch_router_weights = ds4_gpu_tensor_alloc(pc * DS4_N_EXPERT_USED * sizeof(float)); + scr->prefill_seed_router_selected = + ds4_gpu_tensor_alloc((uint64_t)DS4_N_LAYER * + DS4_STREAMING_PREFILL_CACHE_SEED_MAX_TOKENS * + DS4_N_EXPERT_USED * + sizeof(int32_t)); + scr->batch_routed_gate = ds4_gpu_tensor_alloc(pc * DS4_N_EXPERT_USED * routed_mid_dim * sizeof(float)); + scr->batch_routed_up = ds4_gpu_tensor_alloc(pc * DS4_N_EXPERT_USED * routed_mid_dim * sizeof(float)); + scr->batch_routed_mid = ds4_gpu_tensor_alloc(pc * DS4_N_EXPERT_USED * routed_mid_dim * sizeof(float)); + scr->batch_routed_down = ds4_gpu_tensor_alloc(pc * DS4_N_EXPERT_USED * DS4_N_EMBD * sizeof(float)); + scr->batch_routed_out = ds4_gpu_tensor_alloc(pc * DS4_N_EMBD * sizeof(float)); + if (DS4_GPU_ATTN_COMP_CACHE_F16) { + scr->attn_comp_stage = ds4_gpu_tensor_alloc((uint64_t)attn_comp_stage_cap * + DS4_N_HEAD_DIM * sizeof(float)); + } + if (enable_mtp) { + scr->mtp_embed = ds4_gpu_tensor_alloc((uint64_t)DS4_N_EMBD * sizeof(float)); + scr->mtp_enorm = ds4_gpu_tensor_alloc((uint64_t)DS4_N_EMBD * sizeof(float)); + scr->mtp_eproj = ds4_gpu_tensor_alloc((uint64_t)DS4_N_EMBD * sizeof(float)); + scr->mtp_eproj_hc = ds4_gpu_tensor_alloc(hc_dim * sizeof(float)); + scr->mtp_hnorm_hc = ds4_gpu_tensor_alloc(hc_dim * sizeof(float)); + scr->mtp_hproj_hc = ds4_gpu_tensor_alloc(hc_dim * sizeof(float)); + scr->mtp_input_hc = ds4_gpu_tensor_alloc(hc_dim * sizeof(float)); + scr->mtp_state_hc = ds4_gpu_tensor_alloc(hc_dim * sizeof(float)); + scr->mtp_next_hc = ds4_gpu_tensor_alloc(hc_dim * sizeof(float)); + scr->spec_logits = ds4_gpu_tensor_alloc((uint64_t)16 * DS4_N_VOCAB * sizeof(float)); + } + + const bool ok = scr->cpu_router_norm && + (!DS4_GPU_ATTN_COMP_CACHE_F16 || scr->attn_comp_stage) && + scr->cur_hc && scr->flat_hc && scr->hc_mix && + scr->hc_split && scr->hc_pre && scr->hc_post && + scr->hc_comb && scr->attn_cur && scr->attn_norm && + scr->qr && scr->qr_norm && scr->q && + scr->kv_raw && scr->kv && scr->comp_kv_cur && + scr->comp_sc_cur && scr->indexer_q && scr->indexer_weights && + scr->indexer_scores && scr->comp_mask && scr->comp_selected && + scr->heads && scr->attn_low && scr->attn_out && + scr->after_attn_hc && scr->ffn_cur && scr->ffn_norm && + scr->shared_gate && scr->shared_up && scr->shared_mid && + scr->shared_out && scr->router_logits && scr->router_probs && + scr->router_selected && scr->router_weights && scr->routed_gate && + scr->routed_up && scr->routed_mid && scr->routed_down && + scr->routed_out && scr->after_ffn_hc && scr->output_pre && + scr->output_weights && scr->output_embd && scr->output_norm && + scr->logits && scr->prefill_tokens && scr->batch_cur_hc && + scr->batch_next_hc && scr->batch_flat_hc && scr->batch_hc_mix && + scr->batch_hc_split && scr->batch_attn_cur && scr->batch_attn_norm && + scr->batch_qr && scr->batch_qr_norm && scr->batch_q && + scr->batch_q_half && scr->batch_kv_raw && scr->batch_kv && + scr->batch_comp_kv && scr->batch_comp_sc && scr->batch_indexer_q && + scr->batch_indexer_weights && scr->batch_heads && scr->batch_attn_low && + scr->batch_attn_out && scr->batch_group_tmp && scr->batch_low_tmp && + scr->batch_after_attn_hc && scr->batch_ffn_cur && scr->batch_ffn_norm && + scr->batch_shared_gate && scr->batch_shared_up && scr->batch_shared_mid && + scr->batch_shared_out && scr->batch_router_logits && scr->batch_router_probs && + scr->batch_router_selected && scr->batch_router_weights && scr->prefill_seed_router_selected && + scr->batch_routed_gate && scr->batch_routed_up && scr->batch_routed_mid && + scr->batch_routed_down && scr->batch_routed_out && + (!enable_mtp || + (scr->mtp_embed && scr->mtp_enorm && scr->mtp_eproj && + scr->mtp_eproj_hc && scr->mtp_hnorm_hc && scr->mtp_hproj_hc && + scr->mtp_input_hc && scr->mtp_state_hc && scr->mtp_next_hc && + scr->spec_logits)); + if (!ok) metal_scratch_free(scr); + return ok; +} + +/* Reuse the current scratch when it is already big enough, otherwise grow it + * to the union of the old and requested capacities. Sessions sharing an + * engine are driven from one thread at a time, so no live kernel can be using + * these buffers while a new session is being created. */ +static bool metal_scratch_ensure( + ds4_gpu_scratch *scr, + const ds4_weights *weights, + const ds4_layer_weights *layer, + uint32_t comp_cap, + uint32_t attn_comp_stage_cap, + uint32_t prefill_cap, + bool enable_mtp) { + if (scr->cur_hc && + scr->cap_prefill >= prefill_cap && + scr->cap_comp >= comp_cap && + scr->cap_attn_comp_stage >= attn_comp_stage_cap && + (!enable_mtp || scr->mtp_ready)) { + return true; + } + if (comp_cap < scr->cap_comp) comp_cap = scr->cap_comp; + if (attn_comp_stage_cap < scr->cap_attn_comp_stage) + attn_comp_stage_cap = scr->cap_attn_comp_stage; + if (prefill_cap < scr->cap_prefill) prefill_cap = scr->cap_prefill; + enable_mtp = enable_mtp || scr->mtp_ready; + metal_scratch_free(scr); + return metal_scratch_alloc(scr, weights, layer, comp_cap, + attn_comp_stage_cap, prefill_cap, enable_mtp); +} + /* Allocate the Metal graph state for a chosen raw-cache capacity. The model * weights are not copied here; tensors reference the mapped GGUF. */ static bool metal_graph_alloc_raw_cap( ds4_gpu_graph *g, + ds4_gpu_scratch *scr, const ds4_weights *weights, const ds4_layer_weights *layer, uint32_t raw_cap, @@ -10963,6 +11210,7 @@ static bool metal_graph_alloc_raw_cap( uint32_t prefill_cap, bool enable_mtp) { memset(g, 0, sizeof(*g)); + g->scr = scr; g->mtp_enabled = enable_mtp; if (raw_cap == 0) raw_cap = 1; if (ctx_size == 0) ctx_size = raw_cap; @@ -10998,20 +11246,15 @@ static bool metal_graph_alloc_raw_cap( } } - const uint64_t hc_dim = (uint64_t)DS4_N_HC * DS4_N_EMBD; - const uint64_t mix_hc = 2ull * DS4_N_HC + (uint64_t)DS4_N_HC * DS4_N_HC; - const uint64_t q_rank = layer->attn_q_a->dim[1]; - const uint64_t q_dim = (uint64_t)DS4_N_HEAD * DS4_N_HEAD_DIM; - const uint64_t low_dim = (uint64_t)DS4_N_OUT_GROUP * DS4_N_LORA_O; - const uint64_t group_dim = (uint64_t)DS4_N_HEAD_DIM * (DS4_N_HEAD / DS4_N_OUT_GROUP); - const uint64_t shared_dim = layer->ffn_gate_shexp->dim[1]; - const uint64_t routed_mid_dim = layer->ffn_gate_exps->dim[1]; - const uint64_t vocab_dim = weights->output ? weights->output->dim[1] : DS4_N_VOCAB; - const uint64_t comp_width_max = 2ull * (DS4_N_HEAD_DIM > DS4_N_INDEXER_HEAD_DIM - ? DS4_N_HEAD_DIM - : DS4_N_INDEXER_HEAD_DIM); - const uint64_t indexer_q_dim = (uint64_t)DS4_N_INDEXER_HEAD * DS4_N_INDEXER_HEAD_DIM; - const uint64_t pc = prefill_cap; + /* Work buffers are shared across the engine's sessions; grow them if + * this session needs more capacity than any session before it. */ + if (!metal_scratch_ensure(scr, weights, layer, g->comp_cap, + g->attn_comp_stage_cap, prefill_cap, + enable_mtp)) { + return false; + } + scr->refcount++; + uint64_t kv_cache_bytes = 0; const uint64_t context_bytes = metal_graph_context_bytes_for_kv_policy(ctx_size, raw_cap, prefill_cap, &kv_cache_bytes); @@ -11035,24 +11278,6 @@ static bool metal_graph_alloc_raw_cap( (double)context_bytes / 1073741824.0); } - g->cur_hc = ds4_gpu_tensor_alloc(hc_dim * sizeof(float)); - g->flat_hc = ds4_gpu_tensor_alloc(hc_dim * sizeof(float)); - g->hc_mix = ds4_gpu_tensor_alloc(mix_hc * sizeof(float)); - g->hc_split = ds4_gpu_tensor_alloc(mix_hc * sizeof(float)); - g->hc_pre = ds4_gpu_tensor_view(g->hc_split, 0, (uint64_t)DS4_N_HC * sizeof(float)); - g->hc_post = ds4_gpu_tensor_view(g->hc_split, - (uint64_t)DS4_N_HC * sizeof(float), - (uint64_t)DS4_N_HC * sizeof(float)); - g->hc_comb = ds4_gpu_tensor_view(g->hc_split, - 2ull * DS4_N_HC * sizeof(float), - (uint64_t)DS4_N_HC * DS4_N_HC * sizeof(float)); - g->attn_cur = ds4_gpu_tensor_alloc((uint64_t)DS4_N_EMBD * sizeof(float)); - g->attn_norm = ds4_gpu_tensor_alloc((uint64_t)DS4_N_EMBD * sizeof(float)); - g->qr = ds4_gpu_tensor_alloc(q_rank * sizeof(float)); - g->qr_norm = ds4_gpu_tensor_alloc(q_rank * sizeof(float)); - g->q = ds4_gpu_tensor_alloc(q_dim * sizeof(float)); - g->kv_raw = ds4_gpu_tensor_alloc((uint64_t)DS4_N_HEAD_DIM * sizeof(float)); - g->kv = ds4_gpu_tensor_alloc((uint64_t)DS4_N_HEAD_DIM * sizeof(float)); bool state_init_ok = true; for (uint32_t il = 0; il < DS4_N_LAYER; il++) { g->layer_raw_cache[il] = metal_graph_alloc_kv_cache_tensor( @@ -11109,44 +11334,6 @@ static bool metal_graph_alloc_raw_cap( } } } - g->comp_kv_cur = ds4_gpu_tensor_alloc(comp_width_max * sizeof(float)); - g->comp_sc_cur = ds4_gpu_tensor_alloc(comp_width_max * sizeof(float)); - if (DS4_GPU_ATTN_COMP_CACHE_F16) { - g->attn_comp_stage = ds4_gpu_tensor_alloc((uint64_t)g->attn_comp_stage_cap * - DS4_N_HEAD_DIM * sizeof(float)); - } - g->indexer_q = ds4_gpu_tensor_alloc(indexer_q_dim * sizeof(float)); - g->indexer_weights = ds4_gpu_tensor_alloc((uint64_t)DS4_N_INDEXER_HEAD * sizeof(float)); - g->indexer_scores = ds4_gpu_tensor_alloc((uint64_t)g->comp_cap * pc * sizeof(float)); - g->comp_mask = ds4_gpu_tensor_alloc((uint64_t)g->comp_cap * pc * sizeof(float)); - g->comp_selected = ds4_gpu_tensor_alloc((uint64_t)(DS4_N_INDEXER_TOP_K ? DS4_N_INDEXER_TOP_K : 1u) * - pc * sizeof(uint32_t)); - g->heads = ds4_gpu_tensor_alloc(q_dim * sizeof(float)); - g->attn_low = ds4_gpu_tensor_alloc(low_dim * sizeof(float)); - g->attn_out = ds4_gpu_tensor_alloc((uint64_t)DS4_N_EMBD * sizeof(float)); - g->after_attn_hc = ds4_gpu_tensor_alloc(hc_dim * sizeof(float)); - g->ffn_cur = ds4_gpu_tensor_alloc((uint64_t)DS4_N_EMBD * sizeof(float)); - g->ffn_norm = ds4_gpu_tensor_alloc((uint64_t)DS4_N_EMBD * sizeof(float)); - g->cpu_router_norm = xmalloc((size_t)DS4_N_EMBD * sizeof(g->cpu_router_norm[0])); - g->shared_gate = ds4_gpu_tensor_alloc(shared_dim * sizeof(float)); - g->shared_up = ds4_gpu_tensor_alloc(shared_dim * sizeof(float)); - g->shared_mid = ds4_gpu_tensor_alloc(shared_dim * sizeof(float)); - g->shared_out = ds4_gpu_tensor_alloc((uint64_t)DS4_N_EMBD * sizeof(float)); - g->router_logits = ds4_gpu_tensor_alloc(DS4_N_EXPERT * sizeof(float)); - g->router_probs = ds4_gpu_tensor_alloc(DS4_N_EXPERT * sizeof(float)); - g->router_selected = ds4_gpu_tensor_alloc(DS4_N_EXPERT_USED * sizeof(int)); - g->router_weights = ds4_gpu_tensor_alloc(DS4_N_EXPERT_USED * sizeof(float)); - g->routed_gate = ds4_gpu_tensor_alloc((uint64_t)DS4_N_EXPERT_USED * routed_mid_dim * sizeof(float)); - g->routed_up = ds4_gpu_tensor_alloc((uint64_t)DS4_N_EXPERT_USED * routed_mid_dim * sizeof(float)); - g->routed_mid = ds4_gpu_tensor_alloc((uint64_t)DS4_N_EXPERT_USED * routed_mid_dim * sizeof(float)); - g->routed_down = ds4_gpu_tensor_alloc((uint64_t)DS4_N_EXPERT_USED * DS4_N_EMBD * sizeof(float)); - g->routed_out = ds4_gpu_tensor_alloc((uint64_t)DS4_N_EMBD * sizeof(float)); - g->after_ffn_hc = ds4_gpu_tensor_alloc(hc_dim * sizeof(float)); - g->output_pre = ds4_gpu_tensor_alloc((uint64_t)DS4_N_HC * sizeof(float)); - g->output_weights = ds4_gpu_tensor_alloc((uint64_t)DS4_N_HC * sizeof(float)); - g->output_embd = ds4_gpu_tensor_alloc((uint64_t)DS4_N_EMBD * sizeof(float)); - g->output_norm = ds4_gpu_tensor_alloc((uint64_t)DS4_N_EMBD * sizeof(float)); - g->logits = ds4_gpu_tensor_alloc(vocab_dim * sizeof(float)); /* * MTP is deliberately outside the normal graph footprint. A session that * does not opt in with --mtp must allocate and execute exactly the same @@ -11154,66 +11341,12 @@ static bool metal_graph_alloc_raw_cap( * and no MTP scratch hidden behind otherwise unused tensors. */ if (enable_mtp) { - g->mtp_embed = ds4_gpu_tensor_alloc((uint64_t)DS4_N_EMBD * sizeof(float)); - g->mtp_enorm = ds4_gpu_tensor_alloc((uint64_t)DS4_N_EMBD * sizeof(float)); - g->mtp_eproj = ds4_gpu_tensor_alloc((uint64_t)DS4_N_EMBD * sizeof(float)); - g->mtp_eproj_hc = ds4_gpu_tensor_alloc(hc_dim * sizeof(float)); - g->mtp_hnorm_hc = ds4_gpu_tensor_alloc(hc_dim * sizeof(float)); - g->mtp_hproj_hc = ds4_gpu_tensor_alloc(hc_dim * sizeof(float)); - g->mtp_input_hc = ds4_gpu_tensor_alloc(hc_dim * sizeof(float)); - g->mtp_state_hc = ds4_gpu_tensor_alloc(hc_dim * sizeof(float)); - g->mtp_next_hc = ds4_gpu_tensor_alloc(hc_dim * sizeof(float)); g->mtp_raw_cache = metal_graph_alloc_kv_cache_tensor( managed_kv_cache, (uint64_t)raw_cap * DS4_N_HEAD_DIM * sizeof(float)); - g->spec_logits = ds4_gpu_tensor_alloc((uint64_t)16 * DS4_N_VOCAB * sizeof(float)); g->mtp_n_raw = 0; } - g->prefill_tokens = ds4_gpu_tensor_alloc(pc * sizeof(int32_t)); - g->batch_cur_hc = ds4_gpu_tensor_alloc(pc * hc_dim * sizeof(float)); - g->batch_next_hc = ds4_gpu_tensor_alloc(pc * hc_dim * sizeof(float)); - g->batch_flat_hc = ds4_gpu_tensor_alloc(pc * hc_dim * sizeof(float)); - g->batch_hc_mix = ds4_gpu_tensor_alloc(pc * mix_hc * sizeof(float)); - g->batch_hc_split = ds4_gpu_tensor_alloc(pc * mix_hc * sizeof(float)); - g->batch_attn_cur = ds4_gpu_tensor_alloc(pc * DS4_N_EMBD * sizeof(float)); - g->batch_attn_norm = ds4_gpu_tensor_alloc(pc * DS4_N_EMBD * sizeof(float)); - g->batch_qr = ds4_gpu_tensor_alloc(pc * q_rank * sizeof(float)); - g->batch_qr_norm = ds4_gpu_tensor_alloc(pc * q_rank * sizeof(float)); - g->batch_q = ds4_gpu_tensor_alloc(pc * q_dim * sizeof(float)); - g->batch_q_half = ds4_gpu_tensor_alloc(pc * q_dim * sizeof(uint16_t)); - g->batch_kv_raw = ds4_gpu_tensor_alloc(pc * DS4_N_HEAD_DIM * sizeof(float)); - g->batch_kv = ds4_gpu_tensor_alloc(pc * DS4_N_HEAD_DIM * sizeof(float)); - g->batch_comp_kv = ds4_gpu_tensor_alloc(pc * comp_width_max * sizeof(float)); - g->batch_comp_sc = ds4_gpu_tensor_alloc(pc * comp_width_max * sizeof(float)); - g->batch_indexer_q = ds4_gpu_tensor_alloc(pc * indexer_q_dim * sizeof(float)); - g->batch_indexer_weights = ds4_gpu_tensor_alloc(pc * DS4_N_INDEXER_HEAD * sizeof(float)); - g->batch_heads = ds4_gpu_tensor_alloc(pc * q_dim * sizeof(float)); - g->batch_attn_low = ds4_gpu_tensor_alloc(pc * low_dim * sizeof(float)); - g->batch_attn_out = ds4_gpu_tensor_alloc(pc * DS4_N_EMBD * sizeof(float)); - g->batch_group_tmp = ds4_gpu_tensor_alloc(pc * group_dim * sizeof(float)); - g->batch_low_tmp = ds4_gpu_tensor_alloc(pc * DS4_N_LORA_O * sizeof(float)); - g->batch_after_attn_hc = ds4_gpu_tensor_alloc(pc * hc_dim * sizeof(float)); - g->batch_ffn_cur = ds4_gpu_tensor_alloc(pc * DS4_N_EMBD * sizeof(float)); - g->batch_ffn_norm = ds4_gpu_tensor_alloc(pc * DS4_N_EMBD * sizeof(float)); - g->batch_shared_gate = ds4_gpu_tensor_alloc(pc * shared_dim * sizeof(float)); - g->batch_shared_up = ds4_gpu_tensor_alloc(pc * shared_dim * sizeof(float)); - g->batch_shared_mid = ds4_gpu_tensor_alloc(pc * shared_dim * sizeof(float)); - g->batch_shared_out = ds4_gpu_tensor_alloc(pc * DS4_N_EMBD * sizeof(float)); - g->batch_router_logits = ds4_gpu_tensor_alloc(pc * DS4_N_EXPERT * sizeof(float)); - g->batch_router_probs = ds4_gpu_tensor_alloc(pc * DS4_N_EXPERT * sizeof(float)); - g->batch_router_selected = ds4_gpu_tensor_alloc(pc * DS4_N_EXPERT_USED * sizeof(int)); - g->batch_router_weights = ds4_gpu_tensor_alloc(pc * DS4_N_EXPERT_USED * sizeof(float)); - g->prefill_seed_router_selected = - ds4_gpu_tensor_alloc((uint64_t)DS4_N_LAYER * - DS4_STREAMING_PREFILL_CACHE_SEED_MAX_TOKENS * - DS4_N_EXPERT_USED * - sizeof(int32_t)); - g->batch_routed_gate = ds4_gpu_tensor_alloc(pc * DS4_N_EXPERT_USED * routed_mid_dim * sizeof(float)); - g->batch_routed_up = ds4_gpu_tensor_alloc(pc * DS4_N_EXPERT_USED * routed_mid_dim * sizeof(float)); - g->batch_routed_mid = ds4_gpu_tensor_alloc(pc * DS4_N_EXPERT_USED * routed_mid_dim * sizeof(float)); - g->batch_routed_down = ds4_gpu_tensor_alloc(pc * DS4_N_EXPERT_USED * DS4_N_EMBD * sizeof(float)); - g->batch_routed_out = ds4_gpu_tensor_alloc(pc * DS4_N_EMBD * sizeof(float)); bool layer_cache_ok = true; for (uint32_t il = 0; layer_cache_ok && il < DS4_N_LAYER; il++) { @@ -11242,57 +11375,16 @@ static bool metal_graph_alloc_raw_cap( } const bool ok = state_init_ok && layer_cache_ok && - g->cur_hc && g->flat_hc && g->hc_mix && g->hc_split && - g->hc_pre && g->hc_post && g->hc_comb && - g->attn_cur && g->attn_norm && g->qr && g->qr_norm && - g->q && g->kv_raw && g->kv && - g->comp_kv_cur && g->comp_sc_cur && - (!DS4_GPU_ATTN_COMP_CACHE_F16 || g->attn_comp_stage) && - g->indexer_q && g->indexer_weights && g->indexer_scores && - g->comp_mask && g->comp_selected && - g->heads && g->attn_low && g->attn_out && - g->after_attn_hc && g->ffn_cur && g->ffn_norm && - g->shared_gate && g->shared_up && g->shared_mid && - g->shared_out && - g->router_logits && g->router_probs && g->router_selected && g->router_weights && - g->routed_gate && g->routed_up && g->routed_mid && - g->routed_down && g->routed_out && - g->after_ffn_hc && - g->output_pre && g->output_weights && g->output_embd && - g->output_norm && g->logits && - (!enable_mtp || - (g->mtp_embed && g->mtp_enorm && g->mtp_eproj && - g->mtp_eproj_hc && g->mtp_hnorm_hc && g->mtp_hproj_hc && - g->mtp_input_hc && g->mtp_state_hc && g->mtp_next_hc && - g->mtp_raw_cache && g->spec_logits)) && - g->prefill_tokens && - g->batch_cur_hc && g->batch_next_hc && g->batch_flat_hc && - g->batch_hc_mix && g->batch_hc_split && - g->batch_attn_cur && g->batch_attn_norm && - g->batch_qr && g->batch_qr_norm && g->batch_q && g->batch_q_half && - g->batch_kv_raw && g->batch_kv && - g->batch_comp_kv && g->batch_comp_sc && - g->batch_indexer_q && g->batch_indexer_weights && - g->batch_heads && g->batch_attn_low && g->batch_attn_out && - g->batch_group_tmp && g->batch_low_tmp && g->batch_after_attn_hc && - g->batch_ffn_cur && g->batch_ffn_norm && - g->batch_shared_gate && g->batch_shared_up && - g->batch_shared_mid && g->batch_shared_out && - g->batch_router_logits && g->batch_router_probs && - g->batch_router_selected && g->batch_router_weights && - g->prefill_seed_router_selected && - g->batch_routed_gate && g->batch_routed_up && - g->batch_routed_mid && g->batch_routed_down && - g->batch_routed_out; - if (!ok) metal_graph_free(g); + (!enable_mtp || g->mtp_raw_cache); return ok; } static bool metal_graph_alloc( ds4_gpu_graph *g, + ds4_gpu_scratch *scr, const ds4_weights *weights, const ds4_layer_weights *layer) { - return metal_graph_alloc_raw_cap(g, weights, layer, DS4_N_SWA, DS4_N_SWA, 1, false); + return metal_graph_alloc_raw_cap(g, scr, weights, layer, DS4_N_SWA, DS4_N_SWA, 1, false); } static bool metal_graph_install_model_spans( @@ -11899,7 +11991,7 @@ static bool rocm_graph_stream_seed_full_layer_selected( down_expert_bytes); if (ds4_gpu_stream_expert_cache_seed_from_layer_selected( &table, - g->batch_router_selected, + g->scr->batch_router_selected, n_tokens, rocm_graph_stream_prefill_full_layer_seed_tokens(), DS4_N_EXPERT_USED) == 0) { @@ -11947,7 +12039,7 @@ static bool metal_graph_stream_prefill_selected_profile_layer( uint32_t il, uint32_t n_tokens) { if (!metal_graph_stream_prefill_selected_profile_enabled(g)) return true; - if (!layer || !g->batch_router_selected || n_tokens == 0 || + if (!layer || !g->scr->batch_router_selected || n_tokens == 0 || DS4_N_EXPERT == 0 || DS4_N_EXPERT > DS4_MAX_EXPERT || DS4_N_EXPERT_USED == 0 || DS4_N_EXPERT_USED > DS4_MAX_EXPERT_USED) { return false; @@ -11956,7 +12048,7 @@ static bool metal_graph_stream_prefill_selected_profile_layer( const uint64_t n_ids = (uint64_t)n_tokens * DS4_N_EXPERT_USED; if (n_ids > SIZE_MAX / sizeof(int32_t)) return false; int32_t *selected = xmalloc((size_t)n_ids * sizeof(selected[0])); - const bool read_ok = ds4_gpu_tensor_read(g->batch_router_selected, + const bool read_ok = ds4_gpu_tensor_read(g->scr->batch_router_selected, 0, selected, n_ids * sizeof(selected[0])) != 0; @@ -12335,7 +12427,7 @@ static bool metal_graph_stream_prefill_selected_pagein_start( job->madvise_only = madvise_only; if (!metal_graph_stream_prefill_selected_pagein_enabled(g) && !madvise_only) return true; - if (!model || !layer || !g->batch_router_selected || n_tokens == 0) { + if (!model || !layer || !g->scr->batch_router_selected || n_tokens == 0) { return false; } @@ -12344,7 +12436,7 @@ static bool metal_graph_stream_prefill_selected_pagein_start( int32_t *selected = xmalloc((size_t)n_ids * sizeof(selected[0])); const double t_read0 = job->profile ? now_sec() : 0.0; - bool ok = ds4_gpu_tensor_read(g->batch_router_selected, + bool ok = ds4_gpu_tensor_read(g->scr->batch_router_selected, 0, selected, n_ids * sizeof(selected[0])) != 0; @@ -12978,7 +13070,7 @@ static bool metal_graph_stream_readahead_selected_experts_from_gpu( uint64_t gate_expert_bytes, uint64_t down_expert_bytes) { if (!metal_graph_stream_prefill_selected_readahead_enabled(g)) return true; - if (!model || !layer || !g || !g->batch_router_selected || n_tokens == 0) { + if (!model || !layer || !g || !g->scr->batch_router_selected || n_tokens == 0) { return false; } if (sizeof(int) != sizeof(int32_t) || DS4_N_EXPERT > DS4_MAX_EXPERT) { @@ -12992,7 +13084,7 @@ static bool metal_graph_stream_readahead_selected_experts_from_gpu( const bool profile = getenv("DS4_METAL_STREAMING_PREFILL_SELECTED_READAHEAD_PROFILE") != NULL; const double t0 = profile ? now_sec() : 0.0; - bool ok = ds4_gpu_tensor_read(g->batch_router_selected, + bool ok = ds4_gpu_tensor_read(g->scr->batch_router_selected, 0, selected, n_ids * sizeof(selected[0])) != 0; @@ -13494,7 +13586,7 @@ static bool metal_graph_store_attn_comp_stage( uint32_t rows) { if (!g || il >= DS4_N_LAYER) return false; if (rows == 0) return true; - if (!g->layer_attn_comp_cache[il] || !g->attn_comp_stage) return false; + if (!g->layer_attn_comp_cache[il] || !g->scr->attn_comp_stage) return false; if (rows > g->attn_comp_stage_cap || first_row > g->layer_comp_cap[il] || rows > g->layer_comp_cap[il] - first_row) { return false; @@ -13506,14 +13598,14 @@ static bool metal_graph_store_attn_comp_stage( if (DS4_GPU_ATTN_COMP_CACHE_F16) { return ds4_gpu_tensor_copy_f32_to_f16(g->layer_attn_comp_cache[il], dst_offset, - g->attn_comp_stage, + g->scr->attn_comp_stage, 0, count) != 0; } return ds4_gpu_tensor_copy(g->layer_attn_comp_cache[il], dst_offset, - g->attn_comp_stage, + g->scr->attn_comp_stage, 0, count * sizeof(float)) != 0; } @@ -13522,7 +13614,7 @@ static ds4_gpu_tensor *metal_graph_attn_comp_update_target( ds4_gpu_graph *g, uint32_t il) { return DS4_GPU_ATTN_COMP_CACHE_F16 - ? g->attn_comp_stage + ? g->scr->attn_comp_stage : g->layer_attn_comp_cache[il]; } @@ -13544,7 +13636,7 @@ static ds4_gpu_tensor *metal_graph_attn_comp_row_view( uint32_t il, uint32_t row) { if (DS4_GPU_ATTN_COMP_CACHE_F16) { - return ds4_gpu_tensor_view(g->attn_comp_stage, + return ds4_gpu_tensor_view(g->scr->attn_comp_stage, 0, (uint64_t)DS4_N_HEAD_DIM * sizeof(float)); } @@ -13558,7 +13650,7 @@ static ds4_gpu_tensor *metal_graph_attn_comp_prefill_target( uint32_t il, uint32_t first_row, uint32_t rows) { - if (DS4_GPU_ATTN_COMP_CACHE_F16) return g->attn_comp_stage; + if (DS4_GPU_ATTN_COMP_CACHE_F16) return g->scr->attn_comp_stage; const uint32_t view_rows = rows ? rows : 1u; return ds4_gpu_tensor_view(g->layer_attn_comp_cache[il], (uint64_t)first_row * DS4_N_HEAD_DIM * sizeof(float), @@ -14080,10 +14172,10 @@ static bool metal_graph_decode_cpu_router( const double t0 = profile ? now_sec() : 0.0; if (ds4_gpu_end_commands() == 0) return false; const double t_sync = profile ? now_sec() : 0.0; - if (ds4_gpu_tensor_read(g->ffn_norm, + if (ds4_gpu_tensor_read(g->scr->ffn_norm, 0, - g->cpu_router_norm, - (uint64_t)DS4_N_EMBD * sizeof(g->cpu_router_norm[0])) == 0) { + g->scr->cpu_router_norm, + (uint64_t)DS4_N_EMBD * sizeof(g->scr->cpu_router_norm[0])) == 0) { return false; } const double t_read = profile ? now_sec() : 0.0; @@ -14094,7 +14186,7 @@ static bool metal_graph_decode_cpu_router( int32_t selected_i32[DS4_MAX_EXPERT_USED]; float weights[DS4_MAX_EXPERT_USED]; - matvec_any(logits, model, layer->ffn_gate_inp, g->cpu_router_norm); + matvec_any(logits, model, layer->ffn_gate_inp, g->scr->cpu_router_norm); for (uint32_t i = 0; i < DS4_N_EXPERT; i++) { probs[i] = sqrtf(softplus_stable(logits[i])); } @@ -14109,19 +14201,19 @@ static bool metal_graph_decode_cpu_router( } const double t_cpu = profile ? now_sec() : 0.0; - if (ds4_gpu_tensor_write(g->router_logits, + if (ds4_gpu_tensor_write(g->scr->router_logits, 0, logits, (uint64_t)DS4_N_EXPERT * sizeof(logits[0])) == 0 || - ds4_gpu_tensor_write(g->router_probs, + ds4_gpu_tensor_write(g->scr->router_probs, 0, probs, (uint64_t)DS4_N_EXPERT * sizeof(probs[0])) == 0 || - ds4_gpu_tensor_write(g->router_selected, + ds4_gpu_tensor_write(g->scr->router_selected, 0, selected_i32, (uint64_t)DS4_N_EXPERT_USED * sizeof(selected_i32[0])) == 0 || - ds4_gpu_tensor_write(g->router_weights, + ds4_gpu_tensor_write(g->scr->router_weights, 0, weights, (uint64_t)DS4_N_EXPERT_USED * sizeof(weights[0])) == 0) { @@ -14161,7 +14253,7 @@ static bool metal_graph_decode_selected_readahead_override( uint32_t il, uint64_t gate_expert_bytes, uint64_t down_expert_bytes) { - if (!g || !model || !layer || !g->router_selected || + if (!g || !model || !layer || !g->scr->router_selected || DS4_N_EXPERT == 0 || DS4_N_EXPERT > DS4_MAX_EXPERT || DS4_N_EXPERT_USED == 0 || DS4_N_EXPERT_USED > DS4_MAX_EXPERT_USED) { return false; @@ -14174,7 +14266,7 @@ static bool metal_graph_decode_selected_readahead_override( const double t_sync = profile ? now_sec() : 0.0; int32_t selected_ids[DS4_MAX_EXPERT_USED] = {0}; - if (ds4_gpu_tensor_read(g->router_selected, + if (ds4_gpu_tensor_read(g->scr->router_selected, 0, selected_ids, (uint64_t)DS4_N_EXPERT_USED * sizeof(selected_ids[0])) == 0) { @@ -14271,7 +14363,7 @@ static bool metal_graph_decode_cuda_selected_load( #if !defined(DS4_ROCM_BUILD) && !defined(DS4_NO_GPU) && !defined(__APPLE__) if (!metal_graph_decode_cuda_selected_slots_expected(g, layer) || !model || - !g->router_selected || + !g->scr->router_selected || DS4_N_EXPERT == 0 || DS4_N_EXPERT > DS4_MAX_EXPERT || DS4_N_EXPERT_USED == 0 || @@ -14287,7 +14379,7 @@ static bool metal_graph_decode_cuda_selected_load( const double t_sync = profile ? now_sec() : 0.0; int32_t selected_ids[DS4_MAX_EXPERT_USED] = {0}; - bool ok = ds4_gpu_tensor_read(g->router_selected, + bool ok = ds4_gpu_tensor_read(g->scr->router_selected, 0, selected_ids, (uint64_t)DS4_N_EXPERT_USED * @@ -14344,7 +14436,7 @@ static bool metal_graph_cuda_stream_prefill_batch_selected_load( #if !defined(DS4_ROCM_BUILD) && !defined(DS4_NO_GPU) && !defined(__APPLE__) if (!metal_graph_decode_cuda_selected_slots_expected(g, layer) || !model || - !g->batch_router_selected || + !g->scr->batch_router_selected || n_tokens <= 1 || DS4_N_EXPERT == 0 || DS4_N_EXPERT_USED == 0 || @@ -14370,7 +14462,7 @@ static bool metal_graph_cuda_stream_prefill_batch_selected_load( const double t_sync = profile ? now_sec() : 0.0; int32_t *selected_ids = xmalloc((size_t)n_ids64 * sizeof(selected_ids[0])); - bool ok = ds4_gpu_tensor_read(g->batch_router_selected, + bool ok = ds4_gpu_tensor_read(g->scr->batch_router_selected, 0, selected_ids, n_ids64 * sizeof(selected_ids[0])) != 0; @@ -14447,7 +14539,7 @@ static void metal_graph_selected_async_load_run( metal_graph_selected_async_load *job) { job->ok = false; - if (!job->g || !job->model || !job->layer || !job->g->router_selected || + if (!job->g || !job->model || !job->layer || !job->g->scr->router_selected || DS4_N_EXPERT == 0 || DS4_N_EXPERT > DS4_MAX_EXPERT || DS4_N_EXPERT_USED == 0 || DS4_N_EXPERT_USED > DS4_MAX_EXPERT_USED) { return; @@ -14455,7 +14547,7 @@ static void metal_graph_selected_async_load_run( if (job->event_value != 0) { #ifdef DS4_ROCM_BUILD if (ds4_gpu_tensor_read_after_selected_event( - job->g->router_selected, + job->g->scr->router_selected, 0, job->selected_ids, (uint64_t)DS4_N_EXPERT_USED * @@ -14469,7 +14561,7 @@ static void metal_graph_selected_async_load_run( "selected-id async expert load") == 0) { return; } - if (ds4_gpu_tensor_read(job->g->router_selected, + if (ds4_gpu_tensor_read(job->g->scr->router_selected, 0, job->selected_ids, (uint64_t)DS4_N_EXPERT_USED * @@ -14635,7 +14727,7 @@ static void rocm_graph_batch_selected_async_load_run( rocm_graph_batch_selected_async_load *job) { job->ok = false; if (!job->g || !job->model || !job->layer || - !job->g->batch_router_selected || !job->selected_ids || + !job->g->scr->batch_router_selected || !job->selected_ids || job->n_tokens <= 1 || DS4_N_EXPERT == 0 || DS4_N_EXPERT > DS4_MAX_EXPERT || DS4_N_EXPERT_USED == 0 || DS4_N_EXPERT_USED > DS4_MAX_EXPERT_USED) { @@ -14648,7 +14740,7 @@ static void rocm_graph_batch_selected_async_load_run( const uint64_t n_ids = (uint64_t)job->n_tokens * DS4_N_EXPERT_USED; if (n_ids > SIZE_MAX / sizeof(job->selected_ids[0])) return; if (ds4_gpu_tensor_read_after_selected_event( - job->g->batch_router_selected, + job->g->scr->batch_router_selected, 0, job->selected_ids, n_ids * sizeof(job->selected_ids[0]), @@ -14801,7 +14893,7 @@ static bool metal_graph_profile_router_selection( uint32_t il, uint32_t pos) { if (!g_expert_profile.active) return true; - if (!g || !layer || !g->router_selected || !g->router_weights) return false; + if (!g || !layer || !g->scr->router_selected || !g->scr->router_weights) return false; if (ds4_gpu_end_commands() == 0) { fprintf(stderr, @@ -14812,11 +14904,11 @@ static bool metal_graph_profile_router_selection( int32_t selected[DS4_MAX_EXPERT_USED] = {0}; float weights[DS4_MAX_EXPERT_USED] = {0}; const bool read_ok = - ds4_gpu_tensor_read(g->router_selected, + ds4_gpu_tensor_read(g->scr->router_selected, 0, selected, (uint64_t)DS4_N_EXPERT_USED * sizeof(selected[0])) != 0 && - ds4_gpu_tensor_read(g->router_weights, + ds4_gpu_tensor_read(g->scr->router_weights, 0, weights, (uint64_t)DS4_N_EXPERT_USED * sizeof(weights[0])) != 0; @@ -14881,19 +14973,19 @@ static bool metal_graph_encode_decode_layer( ok = metal_graph_layer_stage_profile_boundary("decode", (name), il, pos, 1, &decode_stage_t0); \ } \ } while (0) - if (ok) ok = ds4_gpu_rms_norm_plain_tensor(g->flat_hc, g->cur_hc, (uint32_t)hc_dim, DS4_RMS_EPS) != 0; - if (ok) ok = metal_graph_matmul_plain_tensor(g->hc_mix, model, layer->hc_attn_fn, - hc_dim, mix_hc, g->flat_hc, 1); + if (ok) ok = ds4_gpu_rms_norm_plain_tensor(g->scr->flat_hc, g->scr->cur_hc, (uint32_t)hc_dim, DS4_RMS_EPS) != 0; + if (ok) ok = metal_graph_matmul_plain_tensor(g->scr->hc_mix, model, layer->hc_attn_fn, + hc_dim, mix_hc, g->scr->flat_hc, 1); const bool fuse_hc_norm = DS4_N_HC == 4 && !metal_graph_use_reference_hc_decode() && !metal_graph_use_reference_hc_norm_decode(); if (ok && fuse_hc_norm) { - ok = ds4_gpu_hc_split_weighted_sum_norm_tensor(g->attn_cur, - g->attn_norm, - g->hc_split, - g->hc_mix, - g->cur_hc, + ok = ds4_gpu_hc_split_weighted_sum_norm_tensor(g->scr->attn_cur, + g->scr->attn_norm, + g->scr->hc_split, + g->scr->hc_mix, + g->scr->cur_hc, model->map, model->size, layer->hc_attn_scale->abs_offset, @@ -14906,10 +14998,10 @@ static bool metal_graph_encode_decode_layer( DS4_RMS_EPS) != 0; if (ok) { ok = metal_graph_check_hc_norm_fusion("attn", - g->attn_cur, - g->attn_norm, - g->hc_mix, - g->cur_hc, + g->scr->attn_cur, + g->scr->attn_norm, + g->scr->hc_mix, + g->scr->cur_hc, model, layer->hc_attn_scale->abs_offset, layer->hc_attn_base->abs_offset, @@ -14918,36 +15010,36 @@ static bool metal_graph_encode_decode_layer( pos); } } else if (ok) { - ok = metal_graph_decode_hc_pre(g->attn_cur, - g->hc_split, - g->hc_mix, - g->cur_hc, + ok = metal_graph_decode_hc_pre(g->scr->attn_cur, + g->scr->hc_split, + g->scr->hc_mix, + g->scr->cur_hc, model, layer->hc_attn_scale->abs_offset, layer->hc_attn_base->abs_offset); } DS4_METAL_PROFILE_DECODE_STAGE("attn_hc_pre"); if (ok) { - metal_graph_debug_dump_tensor("hc_attn_pre_mixes", g->hc_mix, mix_hc, il, pos); - metal_graph_debug_dump_tensor("hc_attn_pre_weights", g->hc_pre, DS4_N_HC, il, pos); - metal_graph_debug_dump_tensor("hc_attn_pre_post_weights", g->hc_post, DS4_N_HC, il, pos); - metal_graph_debug_dump_tensor("hc_attn_pre_comb", g->hc_comb, (uint64_t)DS4_N_HC * DS4_N_HC, il, pos); + metal_graph_debug_dump_tensor("hc_attn_pre_mixes", g->scr->hc_mix, mix_hc, il, pos); + metal_graph_debug_dump_tensor("hc_attn_pre_weights", g->scr->hc_pre, DS4_N_HC, il, pos); + metal_graph_debug_dump_tensor("hc_attn_pre_post_weights", g->scr->hc_post, DS4_N_HC, il, pos); + metal_graph_debug_dump_tensor("hc_attn_pre_comb", g->scr->hc_comb, (uint64_t)DS4_N_HC * DS4_N_HC, il, pos); } if (ok) { - metal_graph_debug_dump_tensor("hc_attn_pre", g->attn_cur, DS4_N_EMBD, il, pos); + metal_graph_debug_dump_tensor("hc_attn_pre", g->scr->attn_cur, DS4_N_EMBD, il, pos); } - if (ok && !fuse_hc_norm) ok = ds4_gpu_rms_norm_weight_tensor(g->attn_norm, g->attn_cur, + if (ok && !fuse_hc_norm) ok = ds4_gpu_rms_norm_weight_tensor(g->scr->attn_norm, g->scr->attn_cur, model->map, model->size, layer->attn_norm->abs_offset, DS4_N_EMBD, DS4_RMS_EPS) != 0; DS4_METAL_PROFILE_DECODE_STAGE("attn_norm"); if (ok) { - metal_graph_debug_dump_tensor("attn_norm", g->attn_norm, DS4_N_EMBD, il, pos); + metal_graph_debug_dump_tensor("attn_norm", g->scr->attn_norm, DS4_N_EMBD, il, pos); } bool qkv_pair_projected = false; if (ok && qkv_rms_fused) { - qkv_pair_projected = ds4_gpu_matmul_q8_0_pair_tensor(g->qr, - g->kv_raw, + qkv_pair_projected = ds4_gpu_matmul_q8_0_pair_tensor(g->scr->qr, + g->scr->kv_raw, model->map, model->size, layer->attn_q_a->abs_offset, @@ -14955,64 +15047,64 @@ static bool metal_graph_encode_decode_layer( DS4_N_EMBD, q_rank, DS4_N_HEAD_DIM, - g->attn_norm, + g->scr->attn_norm, 1) != 0; } - if (ok && !qkv_pair_projected) ok = ds4_gpu_matmul_q8_0_tensor(g->qr, + if (ok && !qkv_pair_projected) ok = ds4_gpu_matmul_q8_0_tensor(g->scr->qr, model->map, model->size, layer->attn_q_a->abs_offset, DS4_N_EMBD, q_rank, - g->attn_norm, + g->scr->attn_norm, 1) != 0; if (ok) { - metal_graph_debug_dump_tensor("q_lora", g->qr, q_rank, il, pos); + metal_graph_debug_dump_tensor("q_lora", g->scr->qr, q_rank, il, pos); } if (qkv_rms_fused) { - if (ok && !qkv_pair_projected) ok = ds4_gpu_matmul_q8_0_tensor(g->kv_raw, model->map, model->size, + if (ok && !qkv_pair_projected) ok = ds4_gpu_matmul_q8_0_tensor(g->scr->kv_raw, model->map, model->size, layer->attn_kv->abs_offset, DS4_N_EMBD, DS4_N_HEAD_DIM, - g->attn_norm, 1) != 0; + g->scr->attn_norm, 1) != 0; if (ok) { - metal_graph_debug_dump_tensor("KVraw", g->kv_raw, DS4_N_HEAD_DIM, il, pos); + metal_graph_debug_dump_tensor("KVraw", g->scr->kv_raw, DS4_N_HEAD_DIM, il, pos); } - if (ok) ok = ds4_gpu_dsv4_qkv_rms_norm_rows_tensor(g->qr_norm, - g->qr, + if (ok) ok = ds4_gpu_dsv4_qkv_rms_norm_rows_tensor(g->scr->qr_norm, + g->scr->qr, model->map, model->size, layer->attn_q_a_norm->abs_offset, (uint32_t)q_rank, - g->kv, - g->kv_raw, + g->scr->kv, + g->scr->kv_raw, layer->attn_kv_a_norm->abs_offset, DS4_N_HEAD_DIM, 1, DS4_RMS_EPS) != 0; } else { - if (ok) ok = ds4_gpu_rms_norm_weight_tensor(g->qr_norm, g->qr, + if (ok) ok = ds4_gpu_rms_norm_weight_tensor(g->scr->qr_norm, g->scr->qr, model->map, model->size, layer->attn_q_a_norm->abs_offset, (uint32_t)q_rank, DS4_RMS_EPS) != 0; } if (ok) { - metal_graph_debug_dump_tensor("q_lora_norm", g->qr_norm, q_rank, il, pos); + metal_graph_debug_dump_tensor("q_lora_norm", g->scr->qr_norm, q_rank, il, pos); } if (qkv_rms_fused && ok) { - metal_graph_debug_dump_tensor("KVnorm", g->kv, DS4_N_HEAD_DIM, il, pos); + metal_graph_debug_dump_tensor("KVnorm", g->scr->kv, DS4_N_HEAD_DIM, il, pos); } - if (ok) ok = ds4_gpu_matmul_q8_0_tensor(g->q, model->map, model->size, + if (ok) ok = ds4_gpu_matmul_q8_0_tensor(g->scr->q, model->map, model->size, layer->attn_q_b->abs_offset, q_rank, q_dim, - g->qr_norm, 1) != 0; + g->scr->qr_norm, 1) != 0; if (ok) { - metal_graph_debug_dump_tensor("Qraw", g->q, q_dim, il, pos); + metal_graph_debug_dump_tensor("Qraw", g->scr->q, q_dim, il, pos); } const bool decode_q_norm_debug = metal_graph_debug_wants("Qnorm", il, pos); bool decode_q_norm_rope_fused = false; if (ok && !decode_q_norm_debug) { decode_q_norm_rope_fused = - ds4_gpu_head_rms_norm_rope_tail_tensor(g->q, + ds4_gpu_head_rms_norm_rope_tail_tensor(g->scr->q, 1, DS4_N_HEAD, DS4_N_HEAD_DIM, @@ -15029,11 +15121,11 @@ static bool metal_graph_encode_decode_layer( DS4_RMS_EPS) != 0; } if (!decode_q_norm_rope_fused) { - if (ok) ok = ds4_gpu_head_rms_norm_tensor(g->q, 1, DS4_N_HEAD, DS4_N_HEAD_DIM, DS4_RMS_EPS) != 0; + if (ok) ok = ds4_gpu_head_rms_norm_tensor(g->scr->q, 1, DS4_N_HEAD, DS4_N_HEAD_DIM, DS4_RMS_EPS) != 0; if (ok) { - metal_graph_debug_dump_tensor("Qnorm", g->q, q_dim, il, pos); + metal_graph_debug_dump_tensor("Qnorm", g->scr->q, q_dim, il, pos); } - if (ok) ok = ds4_gpu_rope_tail_tensor(g->q, 1, DS4_N_HEAD, DS4_N_HEAD_DIM, + if (ok) ok = ds4_gpu_rope_tail_tensor(g->scr->q, 1, DS4_N_HEAD, DS4_N_HEAD_DIM, DS4_N_ROT, pos, compressed ? (uint32_t)DS4_ROPE_ORIG_CTX : 0, false, freq_base, freq_scale, ext_factor, attn_factor, @@ -15041,39 +15133,39 @@ static bool metal_graph_encode_decode_layer( } DS4_METAL_PROFILE_DECODE_STAGE("q_path"); if (ok) { - metal_graph_debug_dump_tensor("Qcur", g->q, q_dim, il, pos); + metal_graph_debug_dump_tensor("Qcur", g->scr->q, q_dim, il, pos); } if (!qkv_rms_fused) { - if (ok) ok = ds4_gpu_matmul_q8_0_tensor(g->kv_raw, model->map, model->size, + if (ok) ok = ds4_gpu_matmul_q8_0_tensor(g->scr->kv_raw, model->map, model->size, layer->attn_kv->abs_offset, DS4_N_EMBD, DS4_N_HEAD_DIM, - g->attn_norm, 1) != 0; + g->scr->attn_norm, 1) != 0; if (ok) { - metal_graph_debug_dump_tensor("KVraw", g->kv_raw, DS4_N_HEAD_DIM, il, pos); + metal_graph_debug_dump_tensor("KVraw", g->scr->kv_raw, DS4_N_HEAD_DIM, il, pos); } - if (ok) ok = ds4_gpu_rms_norm_weight_tensor(g->kv, g->kv_raw, + if (ok) ok = ds4_gpu_rms_norm_weight_tensor(g->scr->kv, g->scr->kv_raw, model->map, model->size, layer->attn_kv_a_norm->abs_offset, DS4_N_HEAD_DIM, DS4_RMS_EPS) != 0; if (ok) { - metal_graph_debug_dump_tensor("KVnorm", g->kv, DS4_N_HEAD_DIM, il, pos); + metal_graph_debug_dump_tensor("KVnorm", g->scr->kv, DS4_N_HEAD_DIM, il, pos); } } - if (ok) ok = ds4_gpu_rope_tail_tensor(g->kv, 1, DS4_N_HEAD_KV, DS4_N_HEAD_DIM, + if (ok) ok = ds4_gpu_rope_tail_tensor(g->scr->kv, 1, DS4_N_HEAD_KV, DS4_N_HEAD_DIM, DS4_N_ROT, pos, compressed ? (uint32_t)DS4_ROPE_ORIG_CTX : 0, false, freq_base, freq_scale, ext_factor, attn_factor, DS4_ROPE_YARN_BETA_FAST, DS4_ROPE_YARN_BETA_SLOW) != 0; if (ok) { - metal_graph_debug_dump_tensor("KVrope", g->kv, DS4_N_HEAD_DIM, il, pos); + metal_graph_debug_dump_tensor("KVrope", g->scr->kv, DS4_N_HEAD_DIM, il, pos); } /* RoPE stays as the exact standalone kernel above. The decode fusion * starts after that, where FP8 KV quantization and raw-cache storage can * share one pass without changing the trigonometric path. */ - if (ok) ok = metal_graph_decode_kv_store(g->kv, raw_cache, raw_cap, raw_row); + if (ok) ok = metal_graph_decode_kv_store(g->scr->kv, raw_cache, raw_cap, raw_row); DS4_METAL_PROFILE_DECODE_STAGE("kv_path"); if (ok) { - metal_graph_debug_dump_tensor("KVcur", g->kv, DS4_N_HEAD_DIM, il, pos); + metal_graph_debug_dump_tensor("KVcur", g->scr->kv, DS4_N_HEAD_DIM, il, pos); } uint32_t n_comp = 0; @@ -15103,29 +15195,29 @@ static bool metal_graph_encode_decode_layer( ok = false; } if (ok && !metal_graph_use_reference_compressor_pair_proj()) { - ok = ds4_gpu_matmul_f16_pair_tensor(g->comp_kv_cur, - g->comp_sc_cur, + ok = ds4_gpu_matmul_f16_pair_tensor(g->scr->comp_kv_cur, + g->scr->comp_sc_cur, model->map, model->size, layer->attn_compressor_kv->abs_offset, layer->attn_compressor_gate->abs_offset, DS4_N_EMBD, comp_width, - g->attn_norm, + g->scr->attn_norm, 1) != 0; } else { - if (ok) ok = ds4_gpu_matmul_f16_tensor(g->comp_kv_cur, model->map, model->size, + if (ok) ok = ds4_gpu_matmul_f16_tensor(g->scr->comp_kv_cur, model->map, model->size, layer->attn_compressor_kv->abs_offset, DS4_N_EMBD, comp_width, - g->attn_norm, 1) != 0; - if (ok) ok = ds4_gpu_matmul_f16_tensor(g->comp_sc_cur, model->map, model->size, + g->scr->attn_norm, 1) != 0; + if (ok) ok = ds4_gpu_matmul_f16_tensor(g->scr->comp_sc_cur, model->map, model->size, layer->attn_compressor_gate->abs_offset, DS4_N_EMBD, comp_width, - g->attn_norm, 1) != 0; + g->scr->attn_norm, 1) != 0; } const uint32_t comp_row = g->layer_n_comp[il]; - if (ok) ok = ds4_gpu_compressor_update_tensor(g->comp_kv_cur, - g->comp_sc_cur, + if (ok) ok = ds4_gpu_compressor_update_tensor(g->scr->comp_kv_cur, + g->scr->comp_sc_cur, g->layer_attn_state_kv[il], g->layer_attn_state_score[il], metal_graph_attn_comp_update_target(g, il), @@ -15181,29 +15273,29 @@ static bool metal_graph_encode_decode_layer( ok = false; } if (ok && !metal_graph_use_reference_compressor_pair_proj()) { - ok = ds4_gpu_matmul_f16_pair_tensor(g->comp_kv_cur, - g->comp_sc_cur, + ok = ds4_gpu_matmul_f16_pair_tensor(g->scr->comp_kv_cur, + g->scr->comp_sc_cur, model->map, model->size, layer->indexer_compressor_kv->abs_offset, layer->indexer_compressor_gate->abs_offset, DS4_N_EMBD, index_width, - g->attn_norm, + g->scr->attn_norm, 1) != 0; } else { - if (ok) ok = ds4_gpu_matmul_f16_tensor(g->comp_kv_cur, model->map, model->size, + if (ok) ok = ds4_gpu_matmul_f16_tensor(g->scr->comp_kv_cur, model->map, model->size, layer->indexer_compressor_kv->abs_offset, DS4_N_EMBD, index_width, - g->attn_norm, 1) != 0; - if (ok) ok = ds4_gpu_matmul_f16_tensor(g->comp_sc_cur, model->map, model->size, + g->scr->attn_norm, 1) != 0; + if (ok) ok = ds4_gpu_matmul_f16_tensor(g->scr->comp_sc_cur, model->map, model->size, layer->indexer_compressor_gate->abs_offset, DS4_N_EMBD, index_width, - g->attn_norm, 1) != 0; + g->scr->attn_norm, 1) != 0; } const uint32_t index_row = g->layer_n_index_comp[il]; - if (ok) ok = ds4_gpu_compressor_update_tensor(g->comp_kv_cur, - g->comp_sc_cur, + if (ok) ok = ds4_gpu_compressor_update_tensor(g->scr->comp_kv_cur, + g->scr->comp_sc_cur, g->layer_index_state_kv[il], g->layer_index_state_score[il], g->layer_index_comp_cache[il], @@ -15261,14 +15353,14 @@ static bool metal_graph_encode_decode_layer( fprintf(stderr, "ds4: Metal graph indexer weight projection expects F16 weights\n"); ok = false; } - if (ok) ok = metal_graph_matmul_plain_tensor(g->indexer_q, + if (ok) ok = metal_graph_matmul_plain_tensor(g->scr->indexer_q, model, layer->indexer_attn_q_b, q_rank, indexer_q_dim, - g->qr_norm, + g->scr->qr_norm, 1); - if (ok) ok = ds4_gpu_rope_tail_tensor(g->indexer_q, 1, + if (ok) ok = ds4_gpu_rope_tail_tensor(g->scr->indexer_q, 1, DS4_N_INDEXER_HEAD, DS4_N_INDEXER_HEAD_DIM, DS4_N_ROT, @@ -15281,13 +15373,13 @@ static bool metal_graph_encode_decode_layer( attn_factor, DS4_ROPE_YARN_BETA_FAST, DS4_ROPE_YARN_BETA_SLOW) != 0; - if (ok) ok = ds4_gpu_dsv4_indexer_qat_tensor(g->indexer_q, + if (ok) ok = ds4_gpu_dsv4_indexer_qat_tensor(g->scr->indexer_q, DS4_N_INDEXER_HEAD, DS4_N_INDEXER_HEAD_DIM) != 0; - if (ok) ok = ds4_gpu_matmul_f16_tensor(g->indexer_weights, model->map, model->size, + if (ok) ok = ds4_gpu_matmul_f16_tensor(g->scr->indexer_weights, model->map, model->size, layer->indexer_proj->abs_offset, DS4_N_EMBD, DS4_N_INDEXER_HEAD, - g->attn_norm, 1) != 0; + g->scr->attn_norm, 1) != 0; const float index_scale = 1.0f / sqrtf((float)(DS4_N_INDEXER_HEAD_DIM * DS4_N_INDEXER_HEAD)); if (ok && decode_index_stage_profile) { ok = metal_graph_indexer_stage_profile_boundary(NULL, @@ -15297,9 +15389,9 @@ static bool metal_graph_encode_decode_layer( g->layer_n_index_comp[il], &decode_index_stage_t0); } - if (ok) ok = ds4_gpu_indexer_score_one_tensor(g->indexer_scores, - g->indexer_q, - g->indexer_weights, + if (ok) ok = ds4_gpu_indexer_score_one_tensor(g->scr->indexer_scores, + g->scr->indexer_q, + g->scr->indexer_weights, g->layer_index_comp_cache[il], g->layer_n_index_comp[il], DS4_N_INDEXER_HEAD, @@ -15313,8 +15405,8 @@ static bool metal_graph_encode_decode_layer( g->layer_n_index_comp[il], &decode_index_stage_t0); } - if (ok) ok = ds4_gpu_indexer_topk_tensor(g->comp_selected, - g->indexer_scores, + if (ok) ok = ds4_gpu_indexer_topk_tensor(g->scr->comp_selected, + g->scr->indexer_scores, g->layer_n_index_comp[il], 1, DS4_N_INDEXER_TOP_K) != 0; @@ -15335,7 +15427,7 @@ static bool metal_graph_encode_decode_layer( * plus the selected compressed rows, matching prefill and * avoiding the long-context decode failure. */ if (ok) { - comp_selected = g->comp_selected; + comp_selected = g->scr->comp_selected; /* * Contract: the indexer top-k is fixed by the model config * and must remain the full 512 rows. Do not reduce this for @@ -15377,11 +15469,11 @@ static bool metal_graph_encode_decode_layer( const uint32_t raw_start = metal_graph_raw_start_for_span(g, pos, n_raw); if (n_comp != 0 && comp_selected != NULL && n_selected != 0) { ok = ds4_gpu_attention_indexed_mixed_batch_heads_tensor( - g->heads, + g->scr->heads, model->map, model->size, layer->attn_sinks->abs_offset, - g->q, + g->scr->q, raw_cache, g->layer_attn_comp_cache[il], metal_graph_attn_comp_cache_is_f16(), @@ -15406,10 +15498,10 @@ static bool metal_graph_encode_decode_layer( &decode_index_stage_t0); } } else { - ok = ds4_gpu_attention_decode_heads_tensor(g->heads, + ok = ds4_gpu_attention_decode_heads_tensor(g->scr->heads, model->map, model->size, layer->attn_sinks->abs_offset, - g->q, raw_cache, n_raw, + g->scr->q, raw_cache, n_raw, raw_cap, raw_start, n_comp ? comp_cache : NULL, @@ -15422,9 +15514,9 @@ static bool metal_graph_encode_decode_layer( } DS4_METAL_PROFILE_DECODE_STAGE("attention"); if (ok) { - metal_graph_debug_dump_tensor("kqv_out", g->heads, q_dim, il, pos); + metal_graph_debug_dump_tensor("kqv_out", g->scr->heads, q_dim, il, pos); } - if (ok) ok = ds4_gpu_rope_tail_tensor(g->heads, + if (ok) ok = ds4_gpu_rope_tail_tensor(g->scr->heads, 1, DS4_N_HEAD, DS4_N_HEAD_DIM, DS4_N_ROT, pos, compressed ? (uint32_t)DS4_ROPE_ORIG_CTX : 0, @@ -15436,74 +15528,74 @@ static bool metal_graph_encode_decode_layer( DS4_ROPE_YARN_BETA_FAST, DS4_ROPE_YARN_BETA_SLOW) != 0; if (ok) { - metal_graph_debug_dump_tensor("kqv_back", g->heads, q_dim, il, pos); + metal_graph_debug_dump_tensor("kqv_back", g->scr->heads, q_dim, il, pos); } const bool fuse_attn_out_hc = !metal_graph_directional_steering_attn_enabled(g) && !metal_graph_use_reference_attn_out_hc(); if (ok && fuse_attn_out_hc) { - ok = ds4_gpu_attention_output_low_q8_tensor(g->attn_low, + ok = ds4_gpu_attention_output_low_q8_tensor(g->scr->attn_low, model->map, model->size, layer->attn_output_a->abs_offset, group_dim, rank, n_groups, - g->heads) != 0; + g->scr->heads) != 0; if (ok) { - ok = ds4_gpu_matmul_q8_0_hc_expand_tensor(g->after_attn_hc, - g->attn_out, + ok = ds4_gpu_matmul_q8_0_hc_expand_tensor(g->scr->after_attn_hc, + g->scr->attn_out, model->map, model->size, layer->attn_output_b->abs_offset, (uint64_t)n_groups * rank, DS4_N_EMBD, - g->attn_low, - g->cur_hc, - g->hc_split, + g->scr->attn_low, + g->scr->cur_hc, + g->scr->hc_split, DS4_N_EMBD, DS4_N_HC) != 0; } } else if (ok) { - ok = ds4_gpu_attention_output_q8_batch_tensor(g->attn_out, - g->attn_low, - g->batch_group_tmp, - g->batch_low_tmp, + ok = ds4_gpu_attention_output_q8_batch_tensor(g->scr->attn_out, + g->scr->attn_low, + g->scr->batch_group_tmp, + g->scr->batch_low_tmp, model->map, model->size, layer->attn_output_a->abs_offset, layer->attn_output_b->abs_offset, group_dim, rank, n_groups, DS4_N_EMBD, - g->heads, 1) != 0; + g->scr->heads, 1) != 0; } DS4_METAL_PROFILE_DECODE_STAGE("attn_output"); if (ok) { - metal_graph_debug_dump_tensor("attn_low", g->attn_low, (uint64_t)n_groups * rank, il, pos); + metal_graph_debug_dump_tensor("attn_low", g->scr->attn_low, (uint64_t)n_groups * rank, il, pos); } if (ok) { - metal_graph_debug_dump_tensor("attn_out", g->attn_out, DS4_N_EMBD, il, pos); + metal_graph_debug_dump_tensor("attn_out", g->scr->attn_out, DS4_N_EMBD, il, pos); } if (ok && metal_graph_directional_steering_attn_enabled(g)) { - ok = metal_graph_apply_directional_steering_attn(g, g->attn_out, il, 1); + ok = metal_graph_apply_directional_steering_attn(g, g->scr->attn_out, il, 1); } if (ok && !fuse_attn_out_hc) { - ok = ds4_gpu_hc_expand_tensor(g->after_attn_hc, g->attn_out, g->cur_hc, - g->hc_post, g->hc_comb, DS4_N_EMBD, DS4_N_HC) != 0; + ok = ds4_gpu_hc_expand_tensor(g->scr->after_attn_hc, g->scr->attn_out, g->scr->cur_hc, + g->scr->hc_post, g->scr->hc_comb, DS4_N_EMBD, DS4_N_HC) != 0; } DS4_METAL_PROFILE_DECODE_STAGE("attn_hc_post"); if (ok) { - metal_graph_debug_dump_tensor("hc_attn_post", g->after_attn_hc, hc_dim, il, pos); + metal_graph_debug_dump_tensor("hc_attn_post", g->scr->after_attn_hc, hc_dim, il, pos); } - if (ok) ok = ds4_gpu_rms_norm_plain_tensor(g->flat_hc, g->after_attn_hc, (uint32_t)hc_dim, DS4_RMS_EPS) != 0; - if (ok) ok = metal_graph_matmul_plain_tensor(g->hc_mix, model, layer->hc_ffn_fn, - hc_dim, mix_hc, g->flat_hc, 1); + if (ok) ok = ds4_gpu_rms_norm_plain_tensor(g->scr->flat_hc, g->scr->after_attn_hc, (uint32_t)hc_dim, DS4_RMS_EPS) != 0; + if (ok) ok = metal_graph_matmul_plain_tensor(g->scr->hc_mix, model, layer->hc_ffn_fn, + hc_dim, mix_hc, g->scr->flat_hc, 1); if (ok && fuse_hc_norm) { - ok = ds4_gpu_hc_split_weighted_sum_norm_tensor(g->ffn_cur, - g->ffn_norm, - g->hc_split, - g->hc_mix, - g->after_attn_hc, + ok = ds4_gpu_hc_split_weighted_sum_norm_tensor(g->scr->ffn_cur, + g->scr->ffn_norm, + g->scr->hc_split, + g->scr->hc_mix, + g->scr->after_attn_hc, model->map, model->size, layer->hc_ffn_scale->abs_offset, @@ -15516,10 +15608,10 @@ static bool metal_graph_encode_decode_layer( DS4_RMS_EPS) != 0; if (ok) { ok = metal_graph_check_hc_norm_fusion("ffn", - g->ffn_cur, - g->ffn_norm, - g->hc_mix, - g->after_attn_hc, + g->scr->ffn_cur, + g->scr->ffn_norm, + g->scr->hc_mix, + g->scr->after_attn_hc, model, layer->hc_ffn_scale->abs_offset, layer->hc_ffn_base->abs_offset, @@ -15528,31 +15620,31 @@ static bool metal_graph_encode_decode_layer( pos); } } else if (ok) { - ok = metal_graph_decode_hc_pre(g->ffn_cur, - g->hc_split, - g->hc_mix, - g->after_attn_hc, + ok = metal_graph_decode_hc_pre(g->scr->ffn_cur, + g->scr->hc_split, + g->scr->hc_mix, + g->scr->after_attn_hc, model, layer->hc_ffn_scale->abs_offset, layer->hc_ffn_base->abs_offset); } DS4_METAL_PROFILE_DECODE_STAGE("ffn_hc_pre"); if (ok) { - metal_graph_debug_dump_tensor("hc_ffn_pre_mixes", g->hc_mix, mix_hc, il, pos); - metal_graph_debug_dump_tensor("hc_ffn_pre_weights", g->hc_pre, DS4_N_HC, il, pos); - metal_graph_debug_dump_tensor("hc_ffn_pre_post_weights", g->hc_post, DS4_N_HC, il, pos); - metal_graph_debug_dump_tensor("hc_ffn_pre_comb", g->hc_comb, (uint64_t)DS4_N_HC * DS4_N_HC, il, pos); + metal_graph_debug_dump_tensor("hc_ffn_pre_mixes", g->scr->hc_mix, mix_hc, il, pos); + metal_graph_debug_dump_tensor("hc_ffn_pre_weights", g->scr->hc_pre, DS4_N_HC, il, pos); + metal_graph_debug_dump_tensor("hc_ffn_pre_post_weights", g->scr->hc_post, DS4_N_HC, il, pos); + metal_graph_debug_dump_tensor("hc_ffn_pre_comb", g->scr->hc_comb, (uint64_t)DS4_N_HC * DS4_N_HC, il, pos); } if (ok) { - metal_graph_debug_dump_tensor("hc_ffn_pre", g->ffn_cur, DS4_N_EMBD, il, pos); + metal_graph_debug_dump_tensor("hc_ffn_pre", g->scr->ffn_cur, DS4_N_EMBD, il, pos); } - if (ok && !fuse_hc_norm) ok = ds4_gpu_rms_norm_weight_tensor(g->ffn_norm, g->ffn_cur, + if (ok && !fuse_hc_norm) ok = ds4_gpu_rms_norm_weight_tensor(g->scr->ffn_norm, g->scr->ffn_cur, model->map, model->size, layer->ffn_norm->abs_offset, DS4_N_EMBD, DS4_RMS_EPS) != 0; DS4_METAL_PROFILE_DECODE_STAGE("ffn_norm"); if (ok) { - metal_graph_debug_dump_tensor("ffn_norm", g->ffn_norm, DS4_N_EMBD, il, pos); + metal_graph_debug_dump_tensor("ffn_norm", g->scr->ffn_norm, DS4_N_EMBD, il, pos); } const uint64_t gate_row_bytes = routed_expert_row_bytes(layer->ffn_gate_exps); const uint64_t gate_expert_bytes = expert_mid_dim * gate_row_bytes; @@ -15561,9 +15653,9 @@ static bool metal_graph_encode_decode_layer( if (ok && metal_graph_decode_cpu_router_applicable(g, layer)) { ok = metal_graph_decode_cpu_router(g, model, layer, il, (uint32_t)token); } else { - if (ok) ok = metal_graph_matmul_plain_tensor(g->router_logits, model, layer->ffn_gate_inp, - DS4_N_EMBD, DS4_N_EXPERT, g->ffn_norm, 1); - if (ok) ok = ds4_gpu_router_select_tensor(g->router_selected, g->router_weights, g->router_probs, + if (ok) ok = metal_graph_matmul_plain_tensor(g->scr->router_logits, model, layer->ffn_gate_inp, + DS4_N_EMBD, DS4_N_EXPERT, g->scr->ffn_norm, 1); + if (ok) ok = ds4_gpu_router_select_tensor(g->scr->router_selected, g->scr->router_weights, g->scr->router_probs, model->map, model->size, layer->ffn_exp_probs_b ? layer->ffn_exp_probs_b->abs_offset : 0, layer->ffn_gate_tid2eid ? layer->ffn_gate_tid2eid->abs_offset : 0, @@ -15576,7 +15668,7 @@ static bool metal_graph_encode_decode_layer( 0, layer->ffn_exp_probs_b != NULL, layer->ffn_gate_tid2eid != NULL, - g->router_logits) != 0; + g->scr->router_logits) != 0; if (ok) ok = metal_graph_decode_set_hash_selected_override(model, layer, il, @@ -15588,10 +15680,10 @@ static bool metal_graph_encode_decode_layer( DS4_METAL_PROFILE_DECODE_STAGE("router"); if (ok) ok = metal_graph_profile_router_selection(g, layer, il, pos); if (ok) { - metal_graph_debug_dump_tensor("ffn_moe_logits", g->router_logits, DS4_N_EXPERT, il, pos); - metal_graph_debug_dump_tensor("ffn_moe_probs", g->router_probs, DS4_N_EXPERT, il, pos); - metal_graph_debug_dump_i32_tensor("ffn_moe_topk", g->router_selected, DS4_N_EXPERT_USED, il, pos); - metal_graph_debug_dump_tensor("ffn_moe_weights_scaled", g->router_weights, DS4_N_EXPERT_USED, il, pos); + metal_graph_debug_dump_tensor("ffn_moe_logits", g->scr->router_logits, DS4_N_EXPERT, il, pos); + metal_graph_debug_dump_tensor("ffn_moe_probs", g->scr->router_probs, DS4_N_EXPERT, il, pos); + metal_graph_debug_dump_i32_tensor("ffn_moe_topk", g->scr->router_selected, DS4_N_EXPERT_USED, il, pos); + metal_graph_debug_dump_tensor("ffn_moe_weights_scaled", g->scr->router_weights, DS4_N_EXPERT_USED, il, pos); } const bool keep_ffn_out = metal_graph_needs_ffn_out(g, il, pos); #ifdef DS4_ROCM_BUILD @@ -15664,35 +15756,35 @@ static bool metal_graph_encode_decode_layer( down_expert_bytes); } if (ok && fuse_shared_gate_up) { - ok = ds4_gpu_shared_gate_up_swiglu_q8_0_tensor(g->shared_gate, - g->shared_up, - g->shared_mid, + ok = ds4_gpu_shared_gate_up_swiglu_q8_0_tensor(g->scr->shared_gate, + g->scr->shared_up, + g->scr->shared_mid, model->map, model->size, layer->ffn_gate_shexp->abs_offset, layer->ffn_up_shexp->abs_offset, DS4_N_EMBD, shared_dim, - g->ffn_norm, + g->scr->ffn_norm, DS4_SWIGLU_CLAMP_EXP) != 0; } else if (ok) { - if (ok) ok = ds4_gpu_matmul_q8_0_tensor(g->shared_gate, model->map, model->size, + if (ok) ok = ds4_gpu_matmul_q8_0_tensor(g->scr->shared_gate, model->map, model->size, layer->ffn_gate_shexp->abs_offset, DS4_N_EMBD, shared_dim, - g->ffn_norm, 1) != 0; - if (ok) ok = ds4_gpu_matmul_q8_0_tensor(g->shared_up, model->map, model->size, + g->scr->ffn_norm, 1) != 0; + if (ok) ok = ds4_gpu_matmul_q8_0_tensor(g->scr->shared_up, model->map, model->size, layer->ffn_up_shexp->abs_offset, DS4_N_EMBD, shared_dim, - g->ffn_norm, 1) != 0; - if (ok) ok = ds4_gpu_swiglu_tensor(g->shared_mid, g->shared_gate, g->shared_up, + g->scr->ffn_norm, 1) != 0; + if (ok) ok = ds4_gpu_swiglu_tensor(g->scr->shared_mid, g->scr->shared_gate, g->scr->shared_up, shared_dim, DS4_SWIGLU_CLAMP_EXP, 1.0f) != 0; } DS4_METAL_PROFILE_DECODE_STAGE("shared_gate_up"); - if (ok) ok = ds4_gpu_routed_moe_one_tensor(g->routed_out, - g->routed_gate, - g->routed_up, - g->routed_mid, - g->routed_down, + if (ok) ok = ds4_gpu_routed_moe_one_tensor(g->scr->routed_out, + g->scr->routed_gate, + g->scr->routed_up, + g->scr->routed_mid, + g->scr->routed_down, model->map, model->size, layer->ffn_gate_exps->abs_offset, layer->ffn_up_exps->abs_offset, @@ -15704,76 +15796,76 @@ static bool metal_graph_encode_decode_layer( (uint32_t)expert_in_dim, (uint32_t)down_in_dim, (uint32_t)routed_out_dim, - g->router_selected, g->router_weights, + g->scr->router_selected, g->scr->router_weights, DS4_N_EXPERT, - DS4_N_EXPERT_USED, DS4_SWIGLU_CLAMP_EXP, g->ffn_norm, + DS4_N_EXPERT_USED, DS4_SWIGLU_CLAMP_EXP, g->scr->ffn_norm, il) != 0; DS4_METAL_PROFILE_DECODE_STAGE("routed_moe"); if (ok) { - metal_graph_debug_dump_tensor("ffn_moe_gate_clamped", g->routed_gate, + metal_graph_debug_dump_tensor("ffn_moe_gate_clamped", g->scr->routed_gate, (uint64_t)DS4_N_EXPERT_USED * down_in_dim, il, pos); - metal_graph_debug_dump_tensor("ffn_moe_up_clamped", g->routed_up, + metal_graph_debug_dump_tensor("ffn_moe_up_clamped", g->scr->routed_up, (uint64_t)DS4_N_EXPERT_USED * down_in_dim, il, pos); - metal_graph_debug_dump_tensor("ffn_moe_weighted_swiglu", g->routed_mid, + metal_graph_debug_dump_tensor("ffn_moe_weighted_swiglu", g->scr->routed_mid, (uint64_t)DS4_N_EXPERT_USED * down_in_dim, il, pos); - metal_graph_debug_dump_tensor("ffn_moe_down", g->routed_down, + metal_graph_debug_dump_tensor("ffn_moe_down", g->scr->routed_down, (uint64_t)DS4_N_EXPERT_USED * DS4_N_EMBD, il, pos); - metal_graph_debug_dump_tensor("ffn_moe_out", g->routed_out, DS4_N_EMBD, il, pos); + metal_graph_debug_dump_tensor("ffn_moe_out", g->scr->routed_out, DS4_N_EMBD, il, pos); } if (ok && fuse_shared_down_hc) { - ok = ds4_gpu_shared_down_hc_expand_q8_0_tensor(g->after_ffn_hc, - g->shared_out, + ok = ds4_gpu_shared_down_hc_expand_q8_0_tensor(g->scr->after_ffn_hc, + g->scr->shared_out, model->map, model->size, layer->ffn_down_shexp->abs_offset, shared_dim, DS4_N_EMBD, - g->shared_mid, - g->routed_out, - g->after_attn_hc, - g->hc_split, + g->scr->shared_mid, + g->scr->routed_out, + g->scr->after_attn_hc, + g->scr->hc_split, DS4_N_EMBD, DS4_N_HC) != 0; } else if (ok) { - ok = ds4_gpu_matmul_q8_0_tensor(g->shared_out, model->map, model->size, + ok = ds4_gpu_matmul_q8_0_tensor(g->scr->shared_out, model->map, model->size, layer->ffn_down_shexp->abs_offset, shared_dim, DS4_N_EMBD, - g->shared_mid, 1) != 0; + g->scr->shared_mid, 1) != 0; } DS4_METAL_PROFILE_DECODE_STAGE("shared_down"); if (ok) { - metal_graph_debug_dump_tensor("ffn_shexp", g->shared_out, DS4_N_EMBD, il, pos); + metal_graph_debug_dump_tensor("ffn_shexp", g->scr->shared_out, DS4_N_EMBD, il, pos); } if (ok && keep_ffn_out) { ok = metal_graph_ensure_ffn_out(g) && - ds4_gpu_add_tensor(g->ffn_out, g->shared_out, g->routed_out, DS4_N_EMBD) != 0; + ds4_gpu_add_tensor(g->scr->ffn_out, g->scr->shared_out, g->scr->routed_out, DS4_N_EMBD) != 0; } if (ok && keep_ffn_out) { - metal_graph_debug_dump_tensor("ffn_out", g->ffn_out, DS4_N_EMBD, il, pos); + metal_graph_debug_dump_tensor("ffn_out", g->scr->ffn_out, DS4_N_EMBD, il, pos); } if (ok && metal_graph_directional_steering_ffn_enabled(g)) { - ok = metal_graph_apply_directional_steering_ffn(g, g->ffn_out, il, 1); + ok = metal_graph_apply_directional_steering_ffn(g, g->scr->ffn_out, il, 1); } if (ok && metal_graph_directional_steering_ffn_enabled(g)) { - ok = ds4_gpu_hc_expand_tensor(g->after_ffn_hc, - g->ffn_out, - g->after_attn_hc, - g->hc_post, - g->hc_comb, + ok = ds4_gpu_hc_expand_tensor(g->scr->after_ffn_hc, + g->scr->ffn_out, + g->scr->after_attn_hc, + g->scr->hc_post, + g->scr->hc_comb, DS4_N_EMBD, DS4_N_HC) != 0; } else if (ok && !fuse_shared_down_hc) { - ok = ds4_gpu_hc_expand_add_split_tensor(g->after_ffn_hc, - g->routed_out, - g->shared_out, - g->after_attn_hc, - g->hc_split, + ok = ds4_gpu_hc_expand_add_split_tensor(g->scr->after_ffn_hc, + g->scr->routed_out, + g->scr->shared_out, + g->scr->after_attn_hc, + g->scr->hc_split, DS4_N_EMBD, DS4_N_HC) != 0; } DS4_METAL_PROFILE_DECODE_STAGE("ffn_hc_post"); if (ok) { - metal_graph_debug_dump_tensor("hc_ffn_post", g->after_ffn_hc, hc_dim, il, pos); + metal_graph_debug_dump_tensor("hc_ffn_post", g->scr->after_ffn_hc, hc_dim, il, pos); } return ok; } @@ -15800,35 +15892,35 @@ static bool metal_graph_encode_decode_layer( ok = ds4_gpu_flush_commands() != 0; } if (ok && fuse_shared_gate_up) { - ok = ds4_gpu_shared_gate_up_swiglu_q8_0_tensor(g->shared_gate, - g->shared_up, - g->shared_mid, + ok = ds4_gpu_shared_gate_up_swiglu_q8_0_tensor(g->scr->shared_gate, + g->scr->shared_up, + g->scr->shared_mid, model->map, model->size, layer->ffn_gate_shexp->abs_offset, layer->ffn_up_shexp->abs_offset, DS4_N_EMBD, shared_dim, - g->ffn_norm, + g->scr->ffn_norm, DS4_SWIGLU_CLAMP_EXP) != 0; } else if (ok) { - if (ok) ok = ds4_gpu_matmul_q8_0_tensor(g->shared_gate, model->map, model->size, + if (ok) ok = ds4_gpu_matmul_q8_0_tensor(g->scr->shared_gate, model->map, model->size, layer->ffn_gate_shexp->abs_offset, DS4_N_EMBD, shared_dim, - g->ffn_norm, 1) != 0; - if (ok) ok = ds4_gpu_matmul_q8_0_tensor(g->shared_up, model->map, model->size, + g->scr->ffn_norm, 1) != 0; + if (ok) ok = ds4_gpu_matmul_q8_0_tensor(g->scr->shared_up, model->map, model->size, layer->ffn_up_shexp->abs_offset, DS4_N_EMBD, shared_dim, - g->ffn_norm, 1) != 0; - if (ok) ok = ds4_gpu_swiglu_tensor(g->shared_mid, g->shared_gate, g->shared_up, + g->scr->ffn_norm, 1) != 0; + if (ok) ok = ds4_gpu_swiglu_tensor(g->scr->shared_mid, g->scr->shared_gate, g->scr->shared_up, shared_dim, DS4_SWIGLU_CLAMP_EXP, 1.0f) != 0; } DS4_METAL_PROFILE_DECODE_STAGE("shared_gate_up"); if (ok && !fuse_shared_down_hc) { - ok = ds4_gpu_matmul_q8_0_tensor(g->shared_out, model->map, model->size, + ok = ds4_gpu_matmul_q8_0_tensor(g->scr->shared_out, model->map, model->size, layer->ffn_down_shexp->abs_offset, shared_dim, DS4_N_EMBD, - g->shared_mid, 1) != 0; + g->scr->shared_mid, 1) != 0; } DS4_METAL_PROFILE_DECODE_STAGE("shared_down"); if (async_load_started) { @@ -15842,7 +15934,7 @@ static bool metal_graph_encode_decode_layer( } if (ok && !async_load_started) { int32_t selected_ids[DS4_MAX_EXPERT_USED]; - ok = ds4_gpu_tensor_read(g->router_selected, + ok = ds4_gpu_tensor_read(g->scr->router_selected, 0, selected_ids, (uint64_t)DS4_N_EXPERT_USED * sizeof(selected_ids[0])) != 0 && @@ -15861,11 +15953,11 @@ static bool metal_graph_encode_decode_layer( DS4_N_EXPERT_USED) != 0; } } - if (ok) ok = ds4_gpu_routed_moe_one_tensor(g->routed_out, - g->routed_gate, - g->routed_up, - g->routed_mid, - g->routed_down, + if (ok) ok = ds4_gpu_routed_moe_one_tensor(g->scr->routed_out, + g->scr->routed_gate, + g->scr->routed_up, + g->scr->routed_mid, + g->scr->routed_down, model->map, model->size, layer->ffn_gate_exps->abs_offset, layer->ffn_up_exps->abs_offset, @@ -15877,79 +15969,79 @@ static bool metal_graph_encode_decode_layer( (uint32_t)expert_in_dim, (uint32_t)down_in_dim, (uint32_t)routed_out_dim, - g->router_selected, g->router_weights, + g->scr->router_selected, g->scr->router_weights, DS4_N_EXPERT, - DS4_N_EXPERT_USED, DS4_SWIGLU_CLAMP_EXP, g->ffn_norm, + DS4_N_EXPERT_USED, DS4_SWIGLU_CLAMP_EXP, g->scr->ffn_norm, il) != 0; DS4_METAL_PROFILE_DECODE_STAGE("routed_moe"); if (ok) { - metal_graph_debug_dump_tensor("ffn_moe_gate_clamped", g->routed_gate, + metal_graph_debug_dump_tensor("ffn_moe_gate_clamped", g->scr->routed_gate, (uint64_t)DS4_N_EXPERT_USED * down_in_dim, il, pos); - metal_graph_debug_dump_tensor("ffn_moe_up_clamped", g->routed_up, + metal_graph_debug_dump_tensor("ffn_moe_up_clamped", g->scr->routed_up, (uint64_t)DS4_N_EXPERT_USED * down_in_dim, il, pos); - metal_graph_debug_dump_tensor("ffn_moe_weighted_swiglu", g->routed_mid, + metal_graph_debug_dump_tensor("ffn_moe_weighted_swiglu", g->scr->routed_mid, (uint64_t)DS4_N_EXPERT_USED * down_in_dim, il, pos); - metal_graph_debug_dump_tensor("ffn_moe_down", g->routed_down, + metal_graph_debug_dump_tensor("ffn_moe_down", g->scr->routed_down, (uint64_t)DS4_N_EXPERT_USED * DS4_N_EMBD, il, pos); - metal_graph_debug_dump_tensor("ffn_moe_out", g->routed_out, DS4_N_EMBD, il, pos); + metal_graph_debug_dump_tensor("ffn_moe_out", g->scr->routed_out, DS4_N_EMBD, il, pos); } if (ok && fuse_shared_down_hc) { - ok = ds4_gpu_shared_down_hc_expand_q8_0_tensor(g->after_ffn_hc, - g->shared_out, + ok = ds4_gpu_shared_down_hc_expand_q8_0_tensor(g->scr->after_ffn_hc, + g->scr->shared_out, model->map, model->size, layer->ffn_down_shexp->abs_offset, shared_dim, DS4_N_EMBD, - g->shared_mid, - g->routed_out, - g->after_attn_hc, - g->hc_split, + g->scr->shared_mid, + g->scr->routed_out, + g->scr->after_attn_hc, + g->scr->hc_split, DS4_N_EMBD, DS4_N_HC) != 0; } DS4_METAL_PROFILE_DECODE_STAGE("shared_down"); if (ok) { - metal_graph_debug_dump_tensor("ffn_shexp", g->shared_out, DS4_N_EMBD, il, pos); + metal_graph_debug_dump_tensor("ffn_shexp", g->scr->shared_out, DS4_N_EMBD, il, pos); } if (ok && keep_ffn_out) { ok = metal_graph_ensure_ffn_out(g) && - ds4_gpu_add_tensor(g->ffn_out, g->shared_out, g->routed_out, DS4_N_EMBD) != 0; + ds4_gpu_add_tensor(g->scr->ffn_out, g->scr->shared_out, g->scr->routed_out, DS4_N_EMBD) != 0; } if (ok && keep_ffn_out) { - metal_graph_debug_dump_tensor("ffn_out", g->ffn_out, DS4_N_EMBD, il, pos); + metal_graph_debug_dump_tensor("ffn_out", g->scr->ffn_out, DS4_N_EMBD, il, pos); } if (ok && metal_graph_directional_steering_ffn_enabled(g)) { - ok = metal_graph_apply_directional_steering_ffn(g, g->ffn_out, il, 1); + ok = metal_graph_apply_directional_steering_ffn(g, g->scr->ffn_out, il, 1); } if (ok && metal_graph_directional_steering_ffn_enabled(g)) { - ok = ds4_gpu_hc_expand_tensor(g->after_ffn_hc, - g->ffn_out, - g->after_attn_hc, - g->hc_post, - g->hc_comb, + ok = ds4_gpu_hc_expand_tensor(g->scr->after_ffn_hc, + g->scr->ffn_out, + g->scr->after_attn_hc, + g->scr->hc_post, + g->scr->hc_comb, DS4_N_EMBD, DS4_N_HC) != 0; } else if (ok && !fuse_shared_down_hc) { - ok = ds4_gpu_hc_expand_add_split_tensor(g->after_ffn_hc, - g->routed_out, - g->shared_out, - g->after_attn_hc, - g->hc_split, + ok = ds4_gpu_hc_expand_add_split_tensor(g->scr->after_ffn_hc, + g->scr->routed_out, + g->scr->shared_out, + g->scr->after_attn_hc, + g->scr->hc_split, DS4_N_EMBD, DS4_N_HC) != 0; } DS4_METAL_PROFILE_DECODE_STAGE("ffn_hc_post"); if (ok) { - metal_graph_debug_dump_tensor("hc_ffn_post", g->after_ffn_hc, hc_dim, il, pos); + metal_graph_debug_dump_tensor("hc_ffn_post", g->scr->after_ffn_hc, hc_dim, il, pos); } return ok; } - if (ok) ok = ds4_gpu_routed_moe_one_tensor(g->routed_out, - g->routed_gate, - g->routed_up, - g->routed_mid, - g->routed_down, + if (ok) ok = ds4_gpu_routed_moe_one_tensor(g->scr->routed_out, + g->scr->routed_gate, + g->scr->routed_up, + g->scr->routed_mid, + g->scr->routed_down, model->map, model->size, layer->ffn_gate_exps->abs_offset, layer->ffn_up_exps->abs_offset, @@ -15961,108 +16053,108 @@ static bool metal_graph_encode_decode_layer( (uint32_t)expert_in_dim, (uint32_t)down_in_dim, (uint32_t)routed_out_dim, - g->router_selected, g->router_weights, + g->scr->router_selected, g->scr->router_weights, DS4_N_EXPERT, - DS4_N_EXPERT_USED, DS4_SWIGLU_CLAMP_EXP, g->ffn_norm, + DS4_N_EXPERT_USED, DS4_SWIGLU_CLAMP_EXP, g->scr->ffn_norm, il) != 0; DS4_METAL_PROFILE_DECODE_STAGE("routed_moe"); if (ok) { - metal_graph_debug_dump_tensor("ffn_moe_gate_clamped", g->routed_gate, + metal_graph_debug_dump_tensor("ffn_moe_gate_clamped", g->scr->routed_gate, (uint64_t)DS4_N_EXPERT_USED * down_in_dim, il, pos); - metal_graph_debug_dump_tensor("ffn_moe_up_clamped", g->routed_up, + metal_graph_debug_dump_tensor("ffn_moe_up_clamped", g->scr->routed_up, (uint64_t)DS4_N_EXPERT_USED * down_in_dim, il, pos); } if (ok) { - metal_graph_debug_dump_tensor("ffn_moe_weighted_swiglu", g->routed_mid, + metal_graph_debug_dump_tensor("ffn_moe_weighted_swiglu", g->scr->routed_mid, (uint64_t)DS4_N_EXPERT_USED * down_in_dim, il, pos); } if (ok) { - metal_graph_debug_dump_tensor("ffn_moe_down", g->routed_down, + metal_graph_debug_dump_tensor("ffn_moe_down", g->scr->routed_down, (uint64_t)DS4_N_EXPERT_USED * DS4_N_EMBD, il, pos); } if (ok) { - metal_graph_debug_dump_tensor("ffn_moe_out", g->routed_out, DS4_N_EMBD, il, pos); + metal_graph_debug_dump_tensor("ffn_moe_out", g->scr->routed_out, DS4_N_EMBD, il, pos); } if (ok && fuse_shared_gate_up) { - ok = ds4_gpu_shared_gate_up_swiglu_q8_0_tensor(g->shared_gate, - g->shared_up, - g->shared_mid, + ok = ds4_gpu_shared_gate_up_swiglu_q8_0_tensor(g->scr->shared_gate, + g->scr->shared_up, + g->scr->shared_mid, model->map, model->size, layer->ffn_gate_shexp->abs_offset, layer->ffn_up_shexp->abs_offset, DS4_N_EMBD, shared_dim, - g->ffn_norm, + g->scr->ffn_norm, DS4_SWIGLU_CLAMP_EXP) != 0; } else { - if (ok) ok = ds4_gpu_matmul_q8_0_tensor(g->shared_gate, model->map, model->size, + if (ok) ok = ds4_gpu_matmul_q8_0_tensor(g->scr->shared_gate, model->map, model->size, layer->ffn_gate_shexp->abs_offset, DS4_N_EMBD, shared_dim, - g->ffn_norm, 1) != 0; - if (ok) ok = ds4_gpu_matmul_q8_0_tensor(g->shared_up, model->map, model->size, + g->scr->ffn_norm, 1) != 0; + if (ok) ok = ds4_gpu_matmul_q8_0_tensor(g->scr->shared_up, model->map, model->size, layer->ffn_up_shexp->abs_offset, DS4_N_EMBD, shared_dim, - g->ffn_norm, 1) != 0; - if (ok) ok = ds4_gpu_swiglu_tensor(g->shared_mid, g->shared_gate, g->shared_up, + g->scr->ffn_norm, 1) != 0; + if (ok) ok = ds4_gpu_swiglu_tensor(g->scr->shared_mid, g->scr->shared_gate, g->scr->shared_up, shared_dim, DS4_SWIGLU_CLAMP_EXP, 1.0f) != 0; } DS4_METAL_PROFILE_DECODE_STAGE("shared_gate_up"); if (ok && fuse_shared_down_hc) { - ok = ds4_gpu_shared_down_hc_expand_q8_0_tensor(g->after_ffn_hc, - g->shared_out, + ok = ds4_gpu_shared_down_hc_expand_q8_0_tensor(g->scr->after_ffn_hc, + g->scr->shared_out, model->map, model->size, layer->ffn_down_shexp->abs_offset, shared_dim, DS4_N_EMBD, - g->shared_mid, - g->routed_out, - g->after_attn_hc, - g->hc_split, + g->scr->shared_mid, + g->scr->routed_out, + g->scr->after_attn_hc, + g->scr->hc_split, DS4_N_EMBD, DS4_N_HC) != 0; } else if (ok) { - ok = ds4_gpu_matmul_q8_0_tensor(g->shared_out, model->map, model->size, + ok = ds4_gpu_matmul_q8_0_tensor(g->scr->shared_out, model->map, model->size, layer->ffn_down_shexp->abs_offset, shared_dim, DS4_N_EMBD, - g->shared_mid, 1) != 0; + g->scr->shared_mid, 1) != 0; } DS4_METAL_PROFILE_DECODE_STAGE("shared_down"); if (ok) { - metal_graph_debug_dump_tensor("ffn_shexp", g->shared_out, DS4_N_EMBD, il, pos); + metal_graph_debug_dump_tensor("ffn_shexp", g->scr->shared_out, DS4_N_EMBD, il, pos); } if (ok && keep_ffn_out) { ok = metal_graph_ensure_ffn_out(g) && - ds4_gpu_add_tensor(g->ffn_out, g->shared_out, g->routed_out, DS4_N_EMBD) != 0; + ds4_gpu_add_tensor(g->scr->ffn_out, g->scr->shared_out, g->scr->routed_out, DS4_N_EMBD) != 0; } if (ok && keep_ffn_out) { - metal_graph_debug_dump_tensor("ffn_out", g->ffn_out, DS4_N_EMBD, il, pos); + metal_graph_debug_dump_tensor("ffn_out", g->scr->ffn_out, DS4_N_EMBD, il, pos); } if (ok && metal_graph_directional_steering_ffn_enabled(g)) { - ok = metal_graph_apply_directional_steering_ffn(g, g->ffn_out, il, 1); + ok = metal_graph_apply_directional_steering_ffn(g, g->scr->ffn_out, il, 1); } if (ok && metal_graph_directional_steering_ffn_enabled(g)) { - ok = ds4_gpu_hc_expand_tensor(g->after_ffn_hc, - g->ffn_out, - g->after_attn_hc, - g->hc_post, - g->hc_comb, + ok = ds4_gpu_hc_expand_tensor(g->scr->after_ffn_hc, + g->scr->ffn_out, + g->scr->after_attn_hc, + g->scr->hc_post, + g->scr->hc_comb, DS4_N_EMBD, DS4_N_HC) != 0; } else if (ok && !fuse_shared_down_hc) { - ok = ds4_gpu_hc_expand_add_split_tensor(g->after_ffn_hc, - g->routed_out, - g->shared_out, - g->after_attn_hc, - g->hc_split, + ok = ds4_gpu_hc_expand_add_split_tensor(g->scr->after_ffn_hc, + g->scr->routed_out, + g->scr->shared_out, + g->scr->after_attn_hc, + g->scr->hc_split, DS4_N_EMBD, DS4_N_HC) != 0; } DS4_METAL_PROFILE_DECODE_STAGE("ffn_hc_post"); #undef DS4_METAL_PROFILE_DECODE_STAGE if (ok) { - metal_graph_debug_dump_tensor("hc_ffn_post", g->after_ffn_hc, hc_dim, il, pos); + metal_graph_debug_dump_tensor("hc_ffn_post", g->scr->after_ffn_hc, hc_dim, il, pos); } return ok; } @@ -16074,20 +16166,20 @@ static bool metal_graph_encode_output_head( const ds4_weights *weights, uint64_t vocab_dim) { const uint64_t hc_dim = (uint64_t)DS4_N_HC * DS4_N_EMBD; - bool ok = ds4_gpu_rms_norm_plain_tensor(g->flat_hc, g->cur_hc, (uint32_t)hc_dim, DS4_RMS_EPS) != 0; - if (ok) ok = ds4_gpu_matmul_f16_tensor(g->output_pre, + bool ok = ds4_gpu_rms_norm_plain_tensor(g->scr->flat_hc, g->scr->cur_hc, (uint32_t)hc_dim, DS4_RMS_EPS) != 0; + if (ok) ok = ds4_gpu_matmul_f16_tensor(g->scr->output_pre, model->map, model->size, weights->output_hc_fn->abs_offset, hc_dim, DS4_N_HC, - g->flat_hc, + g->scr->flat_hc, 1) != 0; if (ok) { - metal_graph_debug_dump_tensor("result_hc_pre", g->output_pre, DS4_N_HC, DS4_N_LAYER, 0); + metal_graph_debug_dump_tensor("result_hc_pre", g->scr->output_pre, DS4_N_HC, DS4_N_LAYER, 0); } - if (ok) ok = ds4_gpu_output_hc_weights_tensor(g->output_weights, - g->output_pre, + if (ok) ok = ds4_gpu_output_hc_weights_tensor(g->scr->output_weights, + g->scr->output_pre, model->map, model->size, weights->output_hc_scale->abs_offset, @@ -16095,36 +16187,36 @@ static bool metal_graph_encode_output_head( DS4_N_HC, DS4_HC_EPS) != 0; if (ok) { - metal_graph_debug_dump_tensor("result_hc_weights", g->output_weights, DS4_N_HC, DS4_N_LAYER, 0); + metal_graph_debug_dump_tensor("result_hc_weights", g->scr->output_weights, DS4_N_HC, DS4_N_LAYER, 0); } - if (ok) ok = ds4_gpu_hc_weighted_sum_tensor(g->output_embd, - g->cur_hc, - g->output_weights, + if (ok) ok = ds4_gpu_hc_weighted_sum_tensor(g->scr->output_embd, + g->scr->cur_hc, + g->scr->output_weights, DS4_N_EMBD, DS4_N_HC) != 0; if (ok) { - metal_graph_debug_dump_tensor("result_hc", g->output_embd, DS4_N_EMBD, DS4_N_LAYER, 0); + metal_graph_debug_dump_tensor("result_hc", g->scr->output_embd, DS4_N_EMBD, DS4_N_LAYER, 0); } - if (ok) ok = ds4_gpu_rms_norm_weight_tensor(g->output_norm, - g->output_embd, + if (ok) ok = ds4_gpu_rms_norm_weight_tensor(g->scr->output_norm, + g->scr->output_embd, model->map, model->size, weights->output_norm->abs_offset, DS4_N_EMBD, DS4_RMS_EPS) != 0; if (ok) { - metal_graph_debug_dump_tensor("result_norm", g->output_norm, DS4_N_EMBD, DS4_N_LAYER, 0); + metal_graph_debug_dump_tensor("result_norm", g->scr->output_norm, DS4_N_EMBD, DS4_N_LAYER, 0); } - if (ok) ok = ds4_gpu_matmul_q8_0_tensor(g->logits, + if (ok) ok = ds4_gpu_matmul_q8_0_tensor(g->scr->logits, model->map, model->size, weights->output->abs_offset, DS4_N_EMBD, vocab_dim, - g->output_norm, + g->scr->output_norm, 1) != 0; if (ok) { - metal_graph_debug_dump_tensor("result_output", g->logits, vocab_dim, DS4_N_LAYER, 0); + metal_graph_debug_dump_tensor("result_output", g->scr->logits, vocab_dim, DS4_N_LAYER, 0); } return ok; } @@ -16143,7 +16235,7 @@ static bool metal_graph_encode_output_head_batch( const ds4_weights *weights, uint32_t n_tokens, uint64_t vocab_dim) { - if (n_tokens == 0 || n_tokens > g->prefill_cap || !g->spec_logits) return false; + if (n_tokens == 0 || n_tokens > g->prefill_cap || !g->scr->spec_logits) return false; const uint64_t hc_dim = (uint64_t)DS4_N_HC * DS4_N_EMBD; ds4_gpu_tensor *output_pre = NULL; @@ -16153,25 +16245,25 @@ static bool metal_graph_encode_output_head_batch( ds4_gpu_tensor *logits = NULL; bool ok = true; - output_pre = ds4_gpu_tensor_view(g->batch_hc_mix, + output_pre = ds4_gpu_tensor_view(g->scr->batch_hc_mix, 0, (uint64_t)n_tokens * DS4_N_HC * sizeof(float)); - output_weights = ds4_gpu_tensor_view(g->batch_hc_split, + output_weights = ds4_gpu_tensor_view(g->scr->batch_hc_split, 0, (uint64_t)n_tokens * DS4_N_HC * sizeof(float)); - output_embd = ds4_gpu_tensor_view(g->batch_ffn_cur, + output_embd = ds4_gpu_tensor_view(g->scr->batch_ffn_cur, 0, (uint64_t)n_tokens * DS4_N_EMBD * sizeof(float)); - output_norm = ds4_gpu_tensor_view(g->batch_ffn_norm, + output_norm = ds4_gpu_tensor_view(g->scr->batch_ffn_norm, 0, (uint64_t)n_tokens * DS4_N_EMBD * sizeof(float)); - logits = ds4_gpu_tensor_view(g->spec_logits, + logits = ds4_gpu_tensor_view(g->scr->spec_logits, 0, (uint64_t)n_tokens * vocab_dim * sizeof(float)); ok = output_pre && output_weights && output_embd && output_norm && logits; - if (ok) ok = ds4_gpu_rms_norm_plain_rows_tensor(g->batch_flat_hc, - g->batch_cur_hc, + if (ok) ok = ds4_gpu_rms_norm_plain_rows_tensor(g->scr->batch_flat_hc, + g->scr->batch_cur_hc, (uint32_t)hc_dim, n_tokens, DS4_RMS_EPS) != 0; @@ -16181,7 +16273,7 @@ static bool metal_graph_encode_output_head_batch( weights->output_hc_fn->abs_offset, hc_dim, DS4_N_HC, - g->batch_flat_hc, + g->scr->batch_flat_hc, n_tokens) != 0; if (ok) ok = ds4_gpu_output_hc_weights_tensor(output_weights, output_pre, @@ -16192,7 +16284,7 @@ static bool metal_graph_encode_output_head_batch( DS4_N_HC, DS4_HC_EPS) != 0; if (ok) ok = ds4_gpu_hc_weighted_sum_tensor(output_embd, - g->batch_cur_hc, + g->scr->batch_cur_hc, output_weights, DS4_N_EMBD, DS4_N_HC) != 0; @@ -16278,36 +16370,36 @@ static bool metal_graph_encode_output_head_mtp( const ds4_mtp_weights *mtp, uint64_t vocab_dim) { const uint64_t hc_dim = (uint64_t)DS4_N_HC * DS4_N_EMBD; - bool ok = ds4_gpu_rms_norm_plain_tensor(g->flat_hc, g->cur_hc, (uint32_t)hc_dim, DS4_RMS_EPS) != 0; - if (ok) ok = metal_graph_matmul_plain_tensor(g->output_pre, mtp_model, mtp->hc_head_fn, - hc_dim, DS4_N_HC, g->flat_hc, 1); - if (ok) ok = ds4_gpu_output_hc_weights_tensor(g->output_weights, - g->output_pre, + bool ok = ds4_gpu_rms_norm_plain_tensor(g->scr->flat_hc, g->scr->cur_hc, (uint32_t)hc_dim, DS4_RMS_EPS) != 0; + if (ok) ok = metal_graph_matmul_plain_tensor(g->scr->output_pre, mtp_model, mtp->hc_head_fn, + hc_dim, DS4_N_HC, g->scr->flat_hc, 1); + if (ok) ok = ds4_gpu_output_hc_weights_tensor(g->scr->output_weights, + g->scr->output_pre, mtp_model->map, mtp_model->size, mtp->hc_head_scale->abs_offset, mtp->hc_head_base->abs_offset, DS4_N_HC, DS4_HC_EPS) != 0; - if (ok) ok = ds4_gpu_hc_weighted_sum_tensor(g->output_embd, - g->cur_hc, - g->output_weights, + if (ok) ok = ds4_gpu_hc_weighted_sum_tensor(g->scr->output_embd, + g->scr->cur_hc, + g->scr->output_weights, DS4_N_EMBD, DS4_N_HC) != 0; - if (ok) ok = ds4_gpu_rms_norm_weight_tensor(g->output_norm, - g->output_embd, + if (ok) ok = ds4_gpu_rms_norm_weight_tensor(g->scr->output_norm, + g->scr->output_embd, mtp_model->map, mtp_model->size, mtp->norm->abs_offset, DS4_N_EMBD, DS4_RMS_EPS) != 0; - if (ok) ok = ds4_gpu_matmul_q8_0_tensor(g->logits, + if (ok) ok = ds4_gpu_matmul_q8_0_tensor(g->scr->logits, base_model->map, base_model->size, base_weights->output->abs_offset, DS4_N_EMBD, vocab_dim, - g->output_norm, + g->scr->output_norm, 1) != 0; return ok; } @@ -16436,24 +16528,24 @@ static void metal_graph_trace_layer_stages( int gpu_selected[DS4_MAX_EXPERT_USED]; float gpu_expert_weight[DS4_MAX_EXPERT_USED]; - bool ok = ds4_gpu_tensor_read(g->attn_cur, 0, gpu_attn_cur, (uint64_t)DS4_N_EMBD * sizeof(float)) != 0 && - ds4_gpu_tensor_read(g->attn_norm, 0, gpu_attn_norm, (uint64_t)DS4_N_EMBD * sizeof(float)) != 0 && - ds4_gpu_tensor_read(g->q, 0, gpu_q, q_dim * sizeof(float)) != 0 && - ds4_gpu_tensor_read(g->kv, 0, gpu_kv, (uint64_t)DS4_N_HEAD_DIM * sizeof(float)) != 0 && - ds4_gpu_tensor_read(g->attn_out, 0, gpu_attn_out, (uint64_t)DS4_N_EMBD * sizeof(float)) != 0 && - ds4_gpu_tensor_read(g->after_attn_hc, 0, gpu_after_attn_hc, hc_dim * sizeof(float)) != 0 && - ds4_gpu_tensor_read(g->ffn_cur, 0, gpu_ffn_cur, (uint64_t)DS4_N_EMBD * sizeof(float)) != 0 && - ds4_gpu_tensor_read(g->ffn_norm, 0, gpu_ffn_norm, (uint64_t)DS4_N_EMBD * sizeof(float)) != 0 && - ds4_gpu_tensor_read(g->shared_gate, 0, gpu_shared_gate, shared_dim * sizeof(float)) != 0 && - ds4_gpu_tensor_read(g->shared_up, 0, gpu_shared_up, shared_dim * sizeof(float)) != 0 && - ds4_gpu_tensor_read(g->shared_mid, 0, gpu_shared_mid, shared_dim * sizeof(float)) != 0 && - ds4_gpu_tensor_read(g->shared_out, 0, gpu_shared, (uint64_t)DS4_N_EMBD * sizeof(float)) != 0 && - ds4_gpu_tensor_read(g->router_selected, 0, gpu_selected, sizeof(gpu_selected)) != 0 && - ds4_gpu_tensor_read(g->router_weights, 0, gpu_expert_weight, sizeof(gpu_expert_weight)) != 0 && - ds4_gpu_tensor_read(g->routed_mid, 0, gpu_routed_mid_all, (uint64_t)DS4_N_EXPERT_USED * down_in_dim * sizeof(float)) != 0 && - ds4_gpu_tensor_read(g->routed_out, 0, gpu_routed, (uint64_t)DS4_N_EMBD * sizeof(float)) != 0 && - ds4_gpu_tensor_read(g->ffn_out, 0, gpu_ffn_out, (uint64_t)DS4_N_EMBD * sizeof(float)) != 0 && - ds4_gpu_tensor_read(g->cur_hc, 0, gpu_after_ffn_hc, hc_dim * sizeof(float)) != 0; + bool ok = ds4_gpu_tensor_read(g->scr->attn_cur, 0, gpu_attn_cur, (uint64_t)DS4_N_EMBD * sizeof(float)) != 0 && + ds4_gpu_tensor_read(g->scr->attn_norm, 0, gpu_attn_norm, (uint64_t)DS4_N_EMBD * sizeof(float)) != 0 && + ds4_gpu_tensor_read(g->scr->q, 0, gpu_q, q_dim * sizeof(float)) != 0 && + ds4_gpu_tensor_read(g->scr->kv, 0, gpu_kv, (uint64_t)DS4_N_HEAD_DIM * sizeof(float)) != 0 && + ds4_gpu_tensor_read(g->scr->attn_out, 0, gpu_attn_out, (uint64_t)DS4_N_EMBD * sizeof(float)) != 0 && + ds4_gpu_tensor_read(g->scr->after_attn_hc, 0, gpu_after_attn_hc, hc_dim * sizeof(float)) != 0 && + ds4_gpu_tensor_read(g->scr->ffn_cur, 0, gpu_ffn_cur, (uint64_t)DS4_N_EMBD * sizeof(float)) != 0 && + ds4_gpu_tensor_read(g->scr->ffn_norm, 0, gpu_ffn_norm, (uint64_t)DS4_N_EMBD * sizeof(float)) != 0 && + ds4_gpu_tensor_read(g->scr->shared_gate, 0, gpu_shared_gate, shared_dim * sizeof(float)) != 0 && + ds4_gpu_tensor_read(g->scr->shared_up, 0, gpu_shared_up, shared_dim * sizeof(float)) != 0 && + ds4_gpu_tensor_read(g->scr->shared_mid, 0, gpu_shared_mid, shared_dim * sizeof(float)) != 0 && + ds4_gpu_tensor_read(g->scr->shared_out, 0, gpu_shared, (uint64_t)DS4_N_EMBD * sizeof(float)) != 0 && + ds4_gpu_tensor_read(g->scr->router_selected, 0, gpu_selected, sizeof(gpu_selected)) != 0 && + ds4_gpu_tensor_read(g->scr->router_weights, 0, gpu_expert_weight, sizeof(gpu_expert_weight)) != 0 && + ds4_gpu_tensor_read(g->scr->routed_mid, 0, gpu_routed_mid_all, (uint64_t)DS4_N_EXPERT_USED * down_in_dim * sizeof(float)) != 0 && + ds4_gpu_tensor_read(g->scr->routed_out, 0, gpu_routed, (uint64_t)DS4_N_EMBD * sizeof(float)) != 0 && + ds4_gpu_tensor_read(g->scr->ffn_out, 0, gpu_ffn_out, (uint64_t)DS4_N_EMBD * sizeof(float)) != 0 && + ds4_gpu_tensor_read(g->scr->cur_hc, 0, gpu_after_ffn_hc, hc_dim * sizeof(float)) != 0; if (ok) { fprintf(stderr, @@ -16650,11 +16742,13 @@ static int metal_graph_decode_test( output_logits_one(cpu_logits, model, weights, cpu_after_ffn_hc); ds4_gpu_graph g; - bool ok = metal_graph_alloc(&g, weights, layer); + ds4_gpu_scratch scr; + memset(&scr, 0, sizeof(scr)); + bool ok = metal_graph_alloc(&g, &scr, weights, layer); g.quality = quality; - g.materialize_ffn_out = true; + if (ok) g.scr->materialize_ffn_out = true; if (ok) ok = ds4_gpu_begin_commands() != 0; - if (ok) ok = ds4_gpu_embed_token_hc_tensor(g.cur_hc, + if (ok) ok = ds4_gpu_embed_token_hc_tensor(g.scr->cur_hc, model->map, model->size, weights->token_embd->abs_offset, @@ -16673,31 +16767,31 @@ static int metal_graph_decode_test( 1, token); if (ok) { - ds4_gpu_tensor *embedded_hc = g.cur_hc; - g.cur_hc = g.after_ffn_hc; - g.after_ffn_hc = embedded_hc; + ds4_gpu_tensor *embedded_hc = g.scr->cur_hc; + g.scr->cur_hc = g.scr->after_ffn_hc; + g.scr->after_ffn_hc = embedded_hc; } if (ok) ok = metal_graph_encode_output_head(&g, model, weights, vocab_dim); if (ok) ok = ds4_gpu_end_commands() != 0; if (ok) { - ok = ds4_gpu_tensor_read(g.after_ffn_hc, 0, gpu_hc, hc_dim * sizeof(float)) != 0 && - ds4_gpu_tensor_read(g.attn_cur, 0, gpu_attn_cur, (uint64_t)DS4_N_EMBD * sizeof(float)) != 0 && - ds4_gpu_tensor_read(g.attn_norm, 0, gpu_attn_norm, (uint64_t)DS4_N_EMBD * sizeof(float)) != 0 && - ds4_gpu_tensor_read(g.q, 0, gpu_q, q_dim * sizeof(float)) != 0 && - ds4_gpu_tensor_read(g.kv, 0, gpu_kv, (uint64_t)DS4_N_HEAD_DIM * sizeof(float)) != 0 && + ok = ds4_gpu_tensor_read(g.scr->after_ffn_hc, 0, gpu_hc, hc_dim * sizeof(float)) != 0 && + ds4_gpu_tensor_read(g.scr->attn_cur, 0, gpu_attn_cur, (uint64_t)DS4_N_EMBD * sizeof(float)) != 0 && + ds4_gpu_tensor_read(g.scr->attn_norm, 0, gpu_attn_norm, (uint64_t)DS4_N_EMBD * sizeof(float)) != 0 && + ds4_gpu_tensor_read(g.scr->q, 0, gpu_q, q_dim * sizeof(float)) != 0 && + ds4_gpu_tensor_read(g.scr->kv, 0, gpu_kv, (uint64_t)DS4_N_HEAD_DIM * sizeof(float)) != 0 && ds4_gpu_tensor_read(g.layer_raw_cache[0], 0, gpu_raw, (uint64_t)DS4_N_HEAD_DIM * sizeof(float)) != 0 && - ds4_gpu_tensor_read(g.attn_out, 0, gpu_attn_out, (uint64_t)DS4_N_EMBD * sizeof(float)) != 0 && - ds4_gpu_tensor_read(g.after_attn_hc, 0, gpu_after_attn_hc, hc_dim * sizeof(float)) != 0 && - ds4_gpu_tensor_read(g.ffn_cur, 0, gpu_ffn_cur, (uint64_t)DS4_N_EMBD * sizeof(float)) != 0 && - ds4_gpu_tensor_read(g.ffn_norm, 0, gpu_ffn_norm, (uint64_t)DS4_N_EMBD * sizeof(float)) != 0 && - ds4_gpu_tensor_read(g.shared_out, 0, gpu_shared, (uint64_t)DS4_N_EMBD * sizeof(float)) != 0 && - ds4_gpu_tensor_read(g.router_selected, 0, gpu_selected, sizeof(gpu_selected)) != 0 && - ds4_gpu_tensor_read(g.router_weights, 0, gpu_expert_weight, sizeof(gpu_expert_weight)) != 0 && - ds4_gpu_tensor_read(g.routed_out, 0, gpu_routed, (uint64_t)DS4_N_EMBD * sizeof(float)) != 0 && - ds4_gpu_tensor_read(g.ffn_out, 0, gpu_ffn_out, (uint64_t)DS4_N_EMBD * sizeof(float)) != 0 && - ds4_gpu_tensor_read(g.cur_hc, 0, gpu_after_ffn_hc, hc_dim * sizeof(float)) != 0 && - ds4_gpu_tensor_read(g.logits, 0, gpu_logits, vocab_dim * sizeof(float)) != 0; + ds4_gpu_tensor_read(g.scr->attn_out, 0, gpu_attn_out, (uint64_t)DS4_N_EMBD * sizeof(float)) != 0 && + ds4_gpu_tensor_read(g.scr->after_attn_hc, 0, gpu_after_attn_hc, hc_dim * sizeof(float)) != 0 && + ds4_gpu_tensor_read(g.scr->ffn_cur, 0, gpu_ffn_cur, (uint64_t)DS4_N_EMBD * sizeof(float)) != 0 && + ds4_gpu_tensor_read(g.scr->ffn_norm, 0, gpu_ffn_norm, (uint64_t)DS4_N_EMBD * sizeof(float)) != 0 && + ds4_gpu_tensor_read(g.scr->shared_out, 0, gpu_shared, (uint64_t)DS4_N_EMBD * sizeof(float)) != 0 && + ds4_gpu_tensor_read(g.scr->router_selected, 0, gpu_selected, sizeof(gpu_selected)) != 0 && + ds4_gpu_tensor_read(g.scr->router_weights, 0, gpu_expert_weight, sizeof(gpu_expert_weight)) != 0 && + ds4_gpu_tensor_read(g.scr->routed_out, 0, gpu_routed, (uint64_t)DS4_N_EMBD * sizeof(float)) != 0 && + ds4_gpu_tensor_read(g.scr->ffn_out, 0, gpu_ffn_out, (uint64_t)DS4_N_EMBD * sizeof(float)) != 0 && + ds4_gpu_tensor_read(g.scr->cur_hc, 0, gpu_after_ffn_hc, hc_dim * sizeof(float)) != 0 && + ds4_gpu_tensor_read(g.scr->logits, 0, gpu_logits, vocab_dim * sizeof(float)) != 0; } if (ok) { @@ -16736,6 +16830,8 @@ static int metal_graph_decode_test( } metal_graph_free(&g); + + metal_scratch_free(&scr); free(routed_midq); free(routed_xq); free(routed_mid_all); @@ -16800,11 +16896,13 @@ static int metal_graph_first_token_full_test( output_logits_one(cpu_logits, model, weights, cpu_hc); ds4_gpu_graph g; - bool ok = metal_graph_alloc(&g, weights, &weights->layer[0]); + ds4_gpu_scratch scr; + memset(&scr, 0, sizeof(scr)); + bool ok = metal_graph_alloc(&g, &scr, weights, &weights->layer[0]); g.quality = quality; const bool trace_layers = getenv("DS4_METAL_GRAPH_TRACE_LAYERS") != NULL; if (trace_layers && ok) { - g.materialize_ffn_out = true; + g.scr->materialize_ffn_out = true; const bool teacher_force = getenv("DS4_METAL_GRAPH_TEACHER_FORCE") != NULL; const char *stage_layer_env = getenv("DS4_METAL_GRAPH_TRACE_STAGE_LAYER"); const long stage_layer = stage_layer_env ? strtol(stage_layer_env, NULL, 10) : -1; @@ -16815,7 +16913,7 @@ static int metal_graph_first_token_full_test( embed_token_f16(model, weights, token, plain); hc_from_plain_embedding(cpu_cur, plain, DS4_N_EMBD, DS4_N_HC); ok = ds4_gpu_begin_commands() != 0; - if (ok) ok = ds4_gpu_embed_token_hc_tensor(g.cur_hc, + if (ok) ok = ds4_gpu_embed_token_hc_tensor(g.scr->cur_hc, model->map, model->size, weights->token_embd->abs_offset, @@ -16827,18 +16925,18 @@ static int metal_graph_first_token_full_test( for (uint32_t il = 0; ok && il < DS4_N_LAYER; il++) { if (teacher_force) { - ok = ds4_gpu_tensor_write(g.cur_hc, 0, cpu_cur, hc_dim * sizeof(float)) != 0; + ok = ds4_gpu_tensor_write(g.scr->cur_hc, 0, cpu_cur, hc_dim * sizeof(float)) != 0; } ok = ds4_gpu_begin_commands() != 0; if (ok) ok = metal_graph_encode_decode_layer(&g, model, &weights->layer[il], il, 0, g.layer_raw_cache[il], g.raw_cap, 0, 1, token); - ds4_gpu_tensor *tmp = g.cur_hc; - g.cur_hc = g.after_ffn_hc; - g.after_ffn_hc = tmp; + ds4_gpu_tensor *tmp = g.scr->cur_hc; + g.scr->cur_hc = g.scr->after_ffn_hc; + g.scr->after_ffn_hc = tmp; if (ok) ok = ds4_gpu_end_commands() != 0; layer_forward_self_one(cpu_next, model, &weights->layer[il], cpu_cur, il, 0, token); - if (ok) ok = ds4_gpu_tensor_read(g.cur_hc, 0, gpu_hc, hc_dim * sizeof(float)) != 0; + if (ok) ok = ds4_gpu_tensor_read(g.scr->cur_hc, 0, gpu_hc, hc_dim * sizeof(float)) != 0; if (ok) { fprintf(stderr, "ds4: Metal full graph layer %u%s hc_max=%g hc_rms=%g\n", @@ -16864,7 +16962,7 @@ static int metal_graph_first_token_full_test( free(plain); } else { if (ok) ok = ds4_gpu_begin_commands() != 0; - if (ok) ok = ds4_gpu_embed_token_hc_tensor(g.cur_hc, + if (ok) ok = ds4_gpu_embed_token_hc_tensor(g.scr->cur_hc, model->map, model->size, weights->token_embd->abs_offset, @@ -16877,9 +16975,9 @@ static int metal_graph_first_token_full_test( ok = metal_graph_encode_decode_layer(&g, model, &weights->layer[il], il, 0, g.layer_raw_cache[il], g.raw_cap, 0, 1, token); - ds4_gpu_tensor *tmp = g.cur_hc; - g.cur_hc = g.after_ffn_hc; - g.after_ffn_hc = tmp; + ds4_gpu_tensor *tmp = g.scr->cur_hc; + g.scr->cur_hc = g.scr->after_ffn_hc; + g.scr->after_ffn_hc = tmp; } if (ok) ok = metal_graph_encode_output_head(&g, model, weights, vocab_dim); @@ -16887,8 +16985,8 @@ static int metal_graph_first_token_full_test( } if (ok) { - ok = ds4_gpu_tensor_read(g.cur_hc, 0, gpu_hc, hc_dim * sizeof(float)) != 0 && - ds4_gpu_tensor_read(g.logits, 0, gpu_logits, vocab_dim * sizeof(float)) != 0; + ok = ds4_gpu_tensor_read(g.scr->cur_hc, 0, gpu_hc, hc_dim * sizeof(float)) != 0 && + ds4_gpu_tensor_read(g.scr->logits, 0, gpu_logits, vocab_dim * sizeof(float)) != 0; } if (ok) { @@ -16912,6 +17010,8 @@ static int metal_graph_first_token_full_test( } metal_graph_free(&g); + + metal_scratch_free(&scr); free(gpu_logits); free(cpu_logits); free(gpu_hc); @@ -16958,7 +17058,7 @@ static bool metal_graph_encode_token_raw_swa( const uint32_t raw_row = pos % g->raw_cap; const uint32_t n_raw = metal_graph_raw_span_for_batch(g, pos, 1); - bool ok = ds4_gpu_embed_token_hc_tensor(g->cur_hc, + bool ok = ds4_gpu_embed_token_hc_tensor(g->scr->cur_hc, model->map, model->size, weights->token_embd->abs_offset, @@ -16987,9 +17087,9 @@ static bool metal_graph_encode_token_raw_swa( raw_row, n_raw, token); - ds4_gpu_tensor *tmp = g->cur_hc; - g->cur_hc = g->after_ffn_hc; - g->after_ffn_hc = tmp; + ds4_gpu_tensor *tmp = g->scr->cur_hc; + g->scr->cur_hc = g->scr->after_ffn_hc; + g->scr->after_ffn_hc = tmp; if (ok && allow_split_flush && split_after_layers != 0 && il + 1u == split_after_layers) { ok = ds4_gpu_flush_commands() != 0; } @@ -17061,12 +17161,12 @@ static bool metal_graph_refresh_ratio4_compressor_state( * decisions in later chunks. */ ds4_gpu_tensor *tail_hc = ds4_gpu_tensor_view( - g->batch_attn_norm, + g->scr->batch_attn_norm, (uint64_t)(n_tokens - 4u) * DS4_N_EMBD * sizeof(float), 4ull * DS4_N_EMBD * sizeof(float)); bool ok = tail_hc != NULL; if (ok) { - ok = ds4_gpu_matmul_f16_tensor(g->batch_comp_kv, + ok = ds4_gpu_matmul_f16_tensor(g->scr->batch_comp_kv, model->map, model->size, kv_weight->abs_offset, @@ -17076,7 +17176,7 @@ static bool metal_graph_refresh_ratio4_compressor_state( 4) != 0; } if (ok) { - ok = ds4_gpu_matmul_f16_tensor(g->batch_comp_sc, + ok = ds4_gpu_matmul_f16_tensor(g->scr->batch_comp_sc, model->map, model->size, score_weight->abs_offset, @@ -17088,8 +17188,8 @@ static bool metal_graph_refresh_ratio4_compressor_state( if (ok) { ok = ds4_gpu_compressor_prefill_state_ratio4_tensor(state_kv, state_score, - g->batch_comp_kv, - g->batch_comp_sc, + g->scr->batch_comp_kv, + g->scr->batch_comp_sc, model->map, model->size, ape->abs_offset, @@ -17200,13 +17300,13 @@ static bool metal_graph_warmup_prefill_kernels( bool ok = ds4_gpu_begin_commands() != 0; if (ok) { - ok = ds4_gpu_matmul_f16_tensor(g->batch_hc_mix, + ok = ds4_gpu_matmul_f16_tensor(g->scr->batch_hc_mix, model->map, model->size, weights->layer[0].hc_attn_fn->abs_offset, hc_dim, mix_hc, - g->batch_flat_hc, + g->scr->batch_flat_hc, n_tokens) != 0; } if (ok) ok = ds4_gpu_end_commands() != 0; @@ -17356,19 +17456,19 @@ static bool metal_graph_encode_layer_attention_batch( uint32_t *index_counts = ratio == 4 ? xcalloc(n_tokens, sizeof(index_counts[0])) : NULL; const bool qkv_rms_fused = !metal_graph_use_reference_qkv_norm(); ds4_gpu_tensor *hc_mix_view = ds4_gpu_tensor_view( - g->batch_hc_mix, 0, (uint64_t)n_tokens * mix_hc * sizeof(float)); + g->scr->batch_hc_mix, 0, (uint64_t)n_tokens * mix_hc * sizeof(float)); ds4_gpu_tensor *hc_split_view = ds4_gpu_tensor_view( - g->batch_hc_split, 0, (uint64_t)n_tokens * mix_hc * sizeof(float)); + g->scr->batch_hc_split, 0, (uint64_t)n_tokens * mix_hc * sizeof(float)); ds4_gpu_tensor *attn_cur_view = ds4_gpu_tensor_view( - g->batch_attn_cur, 0, (uint64_t)n_tokens * DS4_N_EMBD * sizeof(float)); + g->scr->batch_attn_cur, 0, (uint64_t)n_tokens * DS4_N_EMBD * sizeof(float)); ds4_gpu_tensor *after_attn_hc_view = ds4_gpu_tensor_view( - g->batch_after_attn_hc, 0, (uint64_t)n_tokens * hc_dim * sizeof(float)); + g->scr->batch_after_attn_hc, 0, (uint64_t)n_tokens * hc_dim * sizeof(float)); bool ok = hc_mix_view && hc_split_view && attn_cur_view && after_attn_hc_view; const bool fuse_hc_norm = DS4_N_HC == 4 && !metal_graph_use_reference_hc_decode() && metal_graph_enable_batch_hc_norm_fusion(); - if (ok) ok = ds4_gpu_rms_norm_plain_rows_tensor(g->batch_flat_hc, - g->batch_cur_hc, + if (ok) ok = ds4_gpu_rms_norm_plain_rows_tensor(g->scr->batch_flat_hc, + g->scr->batch_cur_hc, (uint32_t)hc_dim, n_tokens, DS4_RMS_EPS) != 0; @@ -17378,7 +17478,7 @@ static bool metal_graph_encode_layer_attention_batch( layer->hc_attn_fn->abs_offset, hc_dim, mix_hc, - g->batch_flat_hc, + g->scr->batch_flat_hc, n_tokens) != 0; if (metal_graph_use_reference_hc_decode()) { if (ok) ok = ds4_gpu_hc_split_sinkhorn_tensor(hc_split_view, @@ -17391,16 +17491,16 @@ static bool metal_graph_encode_layer_attention_batch( DS4_N_HC_SINKHORN_ITER, DS4_HC_EPS) != 0; if (ok) ok = ds4_gpu_hc_weighted_sum_split_tensor(attn_cur_view, - g->batch_cur_hc, + g->scr->batch_cur_hc, hc_split_view, DS4_N_EMBD, DS4_N_HC) != 0; } else if (fuse_hc_norm) { if (ok) ok = ds4_gpu_hc_split_weighted_sum_norm_tensor(attn_cur_view, - g->batch_attn_norm, + g->scr->batch_attn_norm, hc_split_view, hc_mix_view, - g->batch_cur_hc, + g->scr->batch_cur_hc, model->map, model->size, layer->hc_attn_scale->abs_offset, @@ -17415,7 +17515,7 @@ static bool metal_graph_encode_layer_attention_batch( if (ok) ok = ds4_gpu_hc_split_weighted_sum_tensor(attn_cur_view, hc_split_view, hc_mix_view, - g->batch_cur_hc, + g->scr->batch_cur_hc, model->map, model->size, layer->hc_attn_scale->abs_offset, @@ -17426,13 +17526,13 @@ static bool metal_graph_encode_layer_attention_batch( DS4_HC_EPS) != 0; } if (ok) { - metal_graph_debug_dump_tensor("hc_attn_pre", g->batch_attn_cur, + metal_graph_debug_dump_tensor("hc_attn_pre", g->scr->batch_attn_cur, (uint64_t)n_tokens * DS4_N_EMBD, il, pos0); } DS4_METAL_PROFILE_ATTN_STAGE("hc_pre"); if (ok && !fuse_hc_norm) { - ok = ds4_gpu_rms_norm_weight_rows_tensor(g->batch_attn_norm, - g->batch_attn_cur, + ok = ds4_gpu_rms_norm_weight_rows_tensor(g->scr->batch_attn_norm, + g->scr->batch_attn_cur, model->map, model->size, layer->attn_norm->abs_offset, @@ -17441,7 +17541,7 @@ static bool metal_graph_encode_layer_attention_batch( DS4_RMS_EPS) != 0; } if (ok) { - metal_graph_debug_dump_tensor("attn_norm", g->batch_attn_norm, + metal_graph_debug_dump_tensor("attn_norm", g->scr->batch_attn_norm, (uint64_t)n_tokens * DS4_N_EMBD, il, pos0); } DS4_METAL_PROFILE_ATTN_STAGE("norm"); @@ -17449,15 +17549,15 @@ static bool metal_graph_encode_layer_attention_batch( if (ok) ok = metal_graph_matmul_q8_0_named_tensor("attn_q_a", il, pos0, - g->batch_qr, + g->scr->batch_qr, model, layer->attn_q_a, DS4_N_EMBD, q_rank, - g->batch_attn_norm, + g->scr->batch_attn_norm, n_tokens); if (ok) { - metal_graph_debug_dump_tensor("q_lora", g->batch_qr, + metal_graph_debug_dump_tensor("q_lora", g->scr->batch_qr, (uint64_t)n_tokens * q_rank, il, pos0); } DS4_METAL_PROFILE_Q_STAGE("q_a"); @@ -17465,32 +17565,32 @@ static bool metal_graph_encode_layer_attention_batch( if (ok) ok = metal_graph_matmul_q8_0_named_tensor("attn_kv", il, pos0, - g->batch_kv_raw, + g->scr->batch_kv_raw, model, layer->attn_kv, DS4_N_EMBD, DS4_N_HEAD_DIM, - g->batch_attn_norm, + g->scr->batch_attn_norm, n_tokens); if (ok) { - metal_graph_debug_dump_tensor("KVraw", g->batch_kv_raw, + metal_graph_debug_dump_tensor("KVraw", g->scr->batch_kv_raw, (uint64_t)n_tokens * DS4_N_HEAD_DIM, il, pos0); } - if (ok) ok = ds4_gpu_dsv4_qkv_rms_norm_rows_tensor(g->batch_qr_norm, - g->batch_qr, + if (ok) ok = ds4_gpu_dsv4_qkv_rms_norm_rows_tensor(g->scr->batch_qr_norm, + g->scr->batch_qr, model->map, model->size, layer->attn_q_a_norm->abs_offset, (uint32_t)q_rank, - g->batch_kv, - g->batch_kv_raw, + g->scr->batch_kv, + g->scr->batch_kv_raw, layer->attn_kv_a_norm->abs_offset, DS4_N_HEAD_DIM, n_tokens, DS4_RMS_EPS) != 0; } else { - if (ok) ok = ds4_gpu_rms_norm_weight_rows_tensor(g->batch_qr_norm, - g->batch_qr, + if (ok) ok = ds4_gpu_rms_norm_weight_rows_tensor(g->scr->batch_qr_norm, + g->scr->batch_qr, model->map, model->size, layer->attn_q_a_norm->abs_offset, @@ -17499,11 +17599,11 @@ static bool metal_graph_encode_layer_attention_batch( DS4_RMS_EPS) != 0; } if (ok) { - metal_graph_debug_dump_tensor("q_lora_norm", g->batch_qr_norm, + metal_graph_debug_dump_tensor("q_lora_norm", g->scr->batch_qr_norm, (uint64_t)n_tokens * q_rank, il, pos0); } if (qkv_rms_fused && ok) { - metal_graph_debug_dump_tensor("KVnorm", g->batch_kv, + metal_graph_debug_dump_tensor("KVnorm", g->scr->batch_kv, (uint64_t)n_tokens * DS4_N_HEAD_DIM, il, pos0); } DS4_METAL_PROFILE_Q_STAGE("q_a_norm"); @@ -17512,14 +17612,14 @@ static bool metal_graph_encode_layer_attention_batch( metal_graph_debug_wants("Qnorm", il, pos0); bool q_b_f16_out = false; if (ok && !q_path_debug) { - q_b_f16_out = ds4_gpu_attn_q_b_f16_head_rms_rope_tail_tensor(g->batch_q, - g->batch_q_half, + q_b_f16_out = ds4_gpu_attn_q_b_f16_head_rms_rope_tail_tensor(g->scr->batch_q, + g->scr->batch_q_half, model->map, model->size, layer->attn_q_b->abs_offset, q_rank, q_dim, - g->batch_qr_norm, + g->scr->batch_qr_norm, n_tokens, DS4_N_HEAD, DS4_N_HEAD_DIM, @@ -17539,7 +17639,7 @@ static bool metal_graph_encode_layer_attention_batch( DS4_METAL_PROFILE_Q_STAGE("q_b"); DS4_METAL_PROFILE_Q_STAGE("head_norm"); if (ok) { - metal_graph_debug_dump_tensor("Qcur", g->batch_q, + metal_graph_debug_dump_tensor("Qcur", g->scr->batch_q, (uint64_t)n_tokens * q_dim, il, pos0); } DS4_METAL_PROFILE_Q_STAGE("rope"); @@ -17547,29 +17647,29 @@ static bool metal_graph_encode_layer_attention_batch( if (ok) ok = metal_graph_matmul_q8_0_named_tensor("attn_q_b", il, pos0, - g->batch_q, + g->scr->batch_q, model, layer->attn_q_b, q_rank, q_dim, - g->batch_qr_norm, + g->scr->batch_qr_norm, n_tokens); if (ok) { - metal_graph_debug_dump_tensor("Qraw", g->batch_q, + metal_graph_debug_dump_tensor("Qraw", g->scr->batch_q, (uint64_t)n_tokens * q_dim, il, pos0); } DS4_METAL_PROFILE_Q_STAGE("q_b"); - if (ok) ok = ds4_gpu_head_rms_norm_tensor(g->batch_q, + if (ok) ok = ds4_gpu_head_rms_norm_tensor(g->scr->batch_q, n_tokens, DS4_N_HEAD, DS4_N_HEAD_DIM, DS4_RMS_EPS) != 0; if (ok) { - metal_graph_debug_dump_tensor("Qnorm", g->batch_q, + metal_graph_debug_dump_tensor("Qnorm", g->scr->batch_q, (uint64_t)n_tokens * q_dim, il, pos0); } DS4_METAL_PROFILE_Q_STAGE("head_norm"); - if (ok) ok = ds4_gpu_rope_tail_tensor(g->batch_q, + if (ok) ok = ds4_gpu_rope_tail_tensor(g->scr->batch_q, n_tokens, DS4_N_HEAD, DS4_N_HEAD_DIM, @@ -17584,7 +17684,7 @@ static bool metal_graph_encode_layer_attention_batch( DS4_ROPE_YARN_BETA_FAST, DS4_ROPE_YARN_BETA_SLOW) != 0; if (ok) { - metal_graph_debug_dump_tensor("Qcur", g->batch_q, + metal_graph_debug_dump_tensor("Qcur", g->scr->batch_q, (uint64_t)n_tokens * q_dim, il, pos0); } DS4_METAL_PROFILE_Q_STAGE("rope"); @@ -17594,19 +17694,19 @@ static bool metal_graph_encode_layer_attention_batch( if (ok) ok = metal_graph_matmul_q8_0_named_tensor("attn_kv", il, pos0, - g->batch_kv_raw, + g->scr->batch_kv_raw, model, layer->attn_kv, DS4_N_EMBD, DS4_N_HEAD_DIM, - g->batch_attn_norm, + g->scr->batch_attn_norm, n_tokens); if (ok) { - metal_graph_debug_dump_tensor("KVraw", g->batch_kv_raw, + metal_graph_debug_dump_tensor("KVraw", g->scr->batch_kv_raw, (uint64_t)n_tokens * DS4_N_HEAD_DIM, il, pos0); } - if (ok) ok = ds4_gpu_rms_norm_weight_rows_tensor(g->batch_kv, - g->batch_kv_raw, + if (ok) ok = ds4_gpu_rms_norm_weight_rows_tensor(g->scr->batch_kv, + g->scr->batch_kv_raw, model->map, model->size, layer->attn_kv_a_norm->abs_offset, @@ -17614,11 +17714,11 @@ static bool metal_graph_encode_layer_attention_batch( n_tokens, DS4_RMS_EPS) != 0; if (ok) { - metal_graph_debug_dump_tensor("KVnorm", g->batch_kv, + metal_graph_debug_dump_tensor("KVnorm", g->scr->batch_kv, (uint64_t)n_tokens * DS4_N_HEAD_DIM, il, pos0); } } - if (ok) ok = ds4_gpu_rope_tail_tensor(g->batch_kv, + if (ok) ok = ds4_gpu_rope_tail_tensor(g->scr->batch_kv, n_tokens, DS4_N_HEAD_KV, DS4_N_HEAD_DIM, @@ -17633,15 +17733,15 @@ static bool metal_graph_encode_layer_attention_batch( DS4_ROPE_YARN_BETA_FAST, DS4_ROPE_YARN_BETA_SLOW) != 0; if (ok) { - metal_graph_debug_dump_tensor("KVrope", g->batch_kv, + metal_graph_debug_dump_tensor("KVrope", g->scr->batch_kv, (uint64_t)n_tokens * DS4_N_HEAD_DIM, il, pos0); } - if (ok) ok = ds4_gpu_dsv4_fp8_kv_quantize_tensor(g->batch_kv, + if (ok) ok = ds4_gpu_dsv4_fp8_kv_quantize_tensor(g->scr->batch_kv, n_tokens, DS4_N_HEAD_DIM, DS4_N_ROT) != 0; if (ok) { - metal_graph_debug_dump_tensor("KVcur", g->batch_kv, + metal_graph_debug_dump_tensor("KVcur", g->scr->batch_kv, (uint64_t)n_tokens * DS4_N_HEAD_DIM, il, pos0); } DS4_METAL_PROFILE_ATTN_STAGE("kv_path"); @@ -17654,7 +17754,7 @@ static bool metal_graph_encode_layer_attention_batch( * attention mask still enforces the 128-token logical window. */ if (ok && zero_prefix) ok = ds4_gpu_store_raw_kv_batch_tensor(g->layer_raw_cache[il], - g->batch_kv, + g->scr->batch_kv, g->raw_cap, pos0, n_tokens, @@ -17663,12 +17763,12 @@ static bool metal_graph_encode_layer_attention_batch( bool batch_attention_done = false; if (ok && raw_batch_attention) { - ok = ds4_gpu_attention_prefill_raw_heads_tensor(g->batch_heads, + ok = ds4_gpu_attention_prefill_raw_heads_tensor(g->scr->batch_heads, model->map, model->size, layer->attn_sinks->abs_offset, - g->batch_q, - g->batch_kv, + g->scr->batch_q, + g->scr->batch_kv, n_tokens, g->raw_window, DS4_N_HEAD, @@ -17689,7 +17789,7 @@ static bool metal_graph_encode_layer_attention_batch( pos0 + n_tokens - 1u, n_raw); ok = ds4_gpu_store_raw_kv_batch_tensor(g->layer_raw_cache[il], - g->batch_kv, + g->scr->batch_kv, g->raw_cap, pos0, n_tokens, @@ -17702,11 +17802,11 @@ static bool metal_graph_encode_layer_attention_batch( pos0); } if (ok) { - ok = ds4_gpu_attention_decode_raw_batch_heads_tensor(g->batch_heads, + ok = ds4_gpu_attention_decode_raw_batch_heads_tensor(g->scr->batch_heads, model->map, model->size, layer->attn_sinks->abs_offset, - g->batch_q, + g->scr->batch_q, g->layer_raw_cache[il], n_tokens, pos0, @@ -17728,30 +17828,30 @@ static bool metal_graph_encode_layer_attention_batch( ok = false; } if (ok) { - ok = ds4_gpu_matmul_f16_tensor(g->batch_comp_kv, + ok = ds4_gpu_matmul_f16_tensor(g->scr->batch_comp_kv, model->map, model->size, layer->attn_compressor_kv->abs_offset, DS4_N_EMBD, comp_width, - g->batch_attn_norm, + g->scr->batch_attn_norm, n_tokens) != 0; - if (ok) ok = ds4_gpu_matmul_f16_tensor(g->batch_comp_sc, + if (ok) ok = ds4_gpu_matmul_f16_tensor(g->scr->batch_comp_sc, model->map, model->size, layer->attn_compressor_gate->abs_offset, DS4_N_EMBD, comp_width, - g->batch_attn_norm, + g->scr->batch_attn_norm, n_tokens) != 0; } if (ok) metal_graph_debug_dump_tensor("attn_comp_kv_raw", - g->batch_comp_kv, + g->scr->batch_comp_kv, (uint64_t)comp_width * n_tokens, il, pos0); if (ok) metal_graph_debug_dump_tensor("attn_comp_score_raw", - g->batch_comp_sc, + g->scr->batch_comp_sc, (uint64_t)comp_width * n_tokens, il, pos0); @@ -17773,8 +17873,8 @@ static bool metal_graph_encode_layer_attention_batch( ds4_gpu_compressor_prefill_tensor(attn_comp_target, g->layer_attn_state_kv[il], g->layer_attn_state_score[il], - g->batch_comp_kv, - g->batch_comp_sc, + g->scr->batch_comp_kv, + g->scr->batch_comp_sc, model->map, model->size, layer->attn_compressor_ape->abs_offset, @@ -17857,8 +17957,8 @@ static bool metal_graph_encode_layer_attention_batch( attn_comp_target, g->layer_attn_state_kv[il], g->layer_attn_state_score[il], - g->batch_comp_kv, - g->batch_comp_sc, + g->scr->batch_comp_kv, + g->scr->batch_comp_sc, model->map, model->size, layer->attn_compressor_ape->abs_offset, @@ -17883,8 +17983,8 @@ static bool metal_graph_encode_layer_attention_batch( attn_comp_target, g->layer_attn_state_kv[il], g->layer_attn_state_score[il], - g->batch_comp_kv, - g->batch_comp_sc, + g->scr->batch_comp_kv, + g->scr->batch_comp_sc, model->map, model->size, layer->attn_compressor_ape->abs_offset, @@ -17955,8 +18055,8 @@ static bool metal_graph_encode_layer_attention_batch( ok = false; break; } - ds4_gpu_tensor *kv_view = metal_graph_tensor_row_view(g->batch_comp_kv, t, comp_width); - ds4_gpu_tensor *sc_view = metal_graph_tensor_row_view(g->batch_comp_sc, t, comp_width); + ds4_gpu_tensor *kv_view = metal_graph_tensor_row_view(g->scr->batch_comp_kv, t, comp_width); + ds4_gpu_tensor *sc_view = metal_graph_tensor_row_view(g->scr->batch_comp_sc, t, comp_width); const uint32_t comp_row = g->layer_n_comp[il]; ok = kv_view && sc_view && ds4_gpu_compressor_update_tensor(kv_view, @@ -18020,41 +18120,41 @@ static bool metal_graph_encode_layer_attention_batch( ok = false; } if (ok) { - ok = ds4_gpu_matmul_f16_tensor(g->batch_comp_kv, + ok = ds4_gpu_matmul_f16_tensor(g->scr->batch_comp_kv, model->map, model->size, layer->indexer_compressor_kv->abs_offset, DS4_N_EMBD, index_width, - g->batch_attn_norm, + g->scr->batch_attn_norm, n_tokens) != 0; - if (ok) ok = ds4_gpu_matmul_f16_tensor(g->batch_comp_sc, + if (ok) ok = ds4_gpu_matmul_f16_tensor(g->scr->batch_comp_sc, model->map, model->size, layer->indexer_compressor_gate->abs_offset, DS4_N_EMBD, index_width, - g->batch_attn_norm, + g->scr->batch_attn_norm, n_tokens) != 0; } if (ok) metal_graph_debug_dump_tensor("indexer_comp_kv_raw", - g->batch_comp_kv, + g->scr->batch_comp_kv, (uint64_t)index_width * n_tokens, il, pos0); if (ok) metal_graph_debug_dump_tensor("indexer_comp_score_raw", - g->batch_comp_sc, + g->scr->batch_comp_sc, (uint64_t)index_width * n_tokens, il, pos0); - if (ok) ok = metal_graph_matmul_plain_tensor(g->batch_indexer_q, + if (ok) ok = metal_graph_matmul_plain_tensor(g->scr->batch_indexer_q, model, layer->indexer_attn_q_b, q_rank, (uint64_t)DS4_N_INDEXER_HEAD * DS4_N_INDEXER_HEAD_DIM, - g->batch_qr_norm, + g->scr->batch_qr_norm, n_tokens); - if (ok) ok = ds4_gpu_rope_tail_tensor(g->batch_indexer_q, + if (ok) ok = ds4_gpu_rope_tail_tensor(g->scr->batch_indexer_q, n_tokens, DS4_N_INDEXER_HEAD, DS4_N_INDEXER_HEAD_DIM, @@ -18068,16 +18168,16 @@ static bool metal_graph_encode_layer_attention_batch( attn_factor, DS4_ROPE_YARN_BETA_FAST, DS4_ROPE_YARN_BETA_SLOW) != 0; - if (ok) ok = ds4_gpu_dsv4_indexer_qat_tensor(g->batch_indexer_q, + if (ok) ok = ds4_gpu_dsv4_indexer_qat_tensor(g->scr->batch_indexer_q, n_tokens * DS4_N_INDEXER_HEAD, DS4_N_INDEXER_HEAD_DIM) != 0; - if (ok) ok = ds4_gpu_matmul_f16_tensor(g->batch_indexer_weights, + if (ok) ok = ds4_gpu_matmul_f16_tensor(g->scr->batch_indexer_weights, model->map, model->size, layer->indexer_proj->abs_offset, DS4_N_EMBD, DS4_N_INDEXER_HEAD, - g->batch_attn_norm, + g->scr->batch_attn_norm, n_tokens) != 0; if (zero_prefix) { if (ok && n_comp > g->layer_comp_cap[il]) { @@ -18088,8 +18188,8 @@ static bool metal_graph_encode_layer_attention_batch( ok = ds4_gpu_compressor_prefill_tensor(g->layer_index_comp_cache[il], g->layer_index_state_kv[il], g->layer_index_state_score[il], - g->batch_comp_kv, - g->batch_comp_sc, + g->scr->batch_comp_kv, + g->scr->batch_comp_sc, model->map, model->size, layer->indexer_compressor_ape->abs_offset, @@ -18174,8 +18274,8 @@ static bool metal_graph_encode_layer_attention_batch( index_view, g->layer_index_state_kv[il], g->layer_index_state_score[il], - g->batch_comp_kv, - g->batch_comp_sc, + g->scr->batch_comp_kv, + g->scr->batch_comp_sc, model->map, model->size, layer->indexer_compressor_ape->abs_offset, @@ -18247,8 +18347,8 @@ static bool metal_graph_encode_layer_attention_batch( ok = false; break; } - ds4_gpu_tensor *kv_view = metal_graph_tensor_row_view(g->batch_comp_kv, t, index_width); - ds4_gpu_tensor *sc_view = metal_graph_tensor_row_view(g->batch_comp_sc, t, index_width); + ds4_gpu_tensor *kv_view = metal_graph_tensor_row_view(g->scr->batch_comp_kv, t, index_width); + ds4_gpu_tensor *sc_view = metal_graph_tensor_row_view(g->scr->batch_comp_sc, t, index_width); const uint32_t index_row = g->layer_n_index_comp[il]; ok = kv_view && sc_view && ds4_gpu_compressor_update_tensor(kv_view, @@ -18312,7 +18412,7 @@ static bool metal_graph_encode_layer_attention_batch( double index_stage_t0 = 0.0; ok = ds4_gpu_store_raw_kv_batch_tensor(g->layer_raw_cache[il], - g->batch_kv, + g->scr->batch_kv, g->raw_cap, pos0, n_tokens, @@ -18327,9 +18427,9 @@ static bool metal_graph_encode_layer_attention_batch( n_comp, &index_stage_t0); } - ok = ds4_gpu_indexer_scores_decode_batch_tensor(g->indexer_scores, - g->batch_indexer_q, - g->batch_indexer_weights, + ok = ds4_gpu_indexer_scores_decode_batch_tensor(g->scr->indexer_scores, + g->scr->batch_indexer_q, + g->scr->batch_indexer_weights, g->layer_index_comp_cache[il], n_comp, n_tokens, @@ -18348,14 +18448,14 @@ static bool metal_graph_encode_layer_attention_batch( } if (ok) { metal_graph_debug_dump_tensor("indexer_scores", - g->indexer_scores, + g->scr->indexer_scores, (uint64_t)n_comp * n_tokens, il, pos0); } if (ok) { - ok = ds4_gpu_indexer_topk_tensor(g->comp_selected, - g->indexer_scores, + ok = ds4_gpu_indexer_topk_tensor(g->scr->comp_selected, + g->scr->indexer_scores, n_comp, n_tokens, DS4_N_INDEXER_TOP_K) != 0; @@ -18369,7 +18469,7 @@ static bool metal_graph_encode_layer_attention_batch( } if (ok) { metal_graph_debug_dump_i32_tensor("indexer_topk", - g->comp_selected, + g->scr->comp_selected, (uint64_t)n_tokens * DS4_N_INDEXER_TOP_K, il, pos0); @@ -18382,15 +18482,15 @@ static bool metal_graph_encode_layer_attention_batch( } if (ok) { if (use_indexed_comp) { - ok = ds4_gpu_attention_indexed_mixed_batch_heads_tensor(g->batch_heads, + ok = ds4_gpu_attention_indexed_mixed_batch_heads_tensor(g->scr->batch_heads, model->map, model->size, layer->attn_sinks->abs_offset, - g->batch_q, + g->scr->batch_q, g->layer_raw_cache[il], g->layer_attn_comp_cache[il], metal_graph_attn_comp_cache_is_f16(), - g->comp_selected, + g->scr->comp_selected, n_tokens, pos0, n_raw, @@ -18411,15 +18511,15 @@ static bool metal_graph_encode_layer_attention_batch( &index_stage_t0); } } else { - ok = ds4_gpu_attention_decode_mixed_batch_heads_tensor(g->batch_heads, + ok = ds4_gpu_attention_decode_mixed_batch_heads_tensor(g->scr->batch_heads, model->map, model->size, layer->attn_sinks->abs_offset, - g->batch_q, + g->scr->batch_q, g->layer_raw_cache[il], g->layer_attn_comp_cache[il], metal_graph_attn_comp_cache_is_f16(), - use_comp_mask ? g->comp_mask : NULL, + use_comp_mask ? g->scr->comp_mask : NULL, use_comp_mask, n_tokens, pos0, @@ -18448,9 +18548,9 @@ static bool metal_graph_encode_layer_attention_batch( n_comp, &index_stage_t0); } - ok = ds4_gpu_indexer_scores_prefill_tensor(g->indexer_scores, - g->batch_indexer_q, - g->batch_indexer_weights, + ok = ds4_gpu_indexer_scores_prefill_tensor(g->scr->indexer_scores, + g->scr->batch_indexer_q, + g->scr->batch_indexer_weights, g->layer_index_comp_cache[il], n_comp, n_tokens, @@ -18468,14 +18568,14 @@ static bool metal_graph_encode_layer_attention_batch( } if (ok) { metal_graph_debug_dump_tensor("indexer_scores", - g->indexer_scores, + g->scr->indexer_scores, (uint64_t)n_comp * n_tokens, il, pos0); } if (ok) { - ok = ds4_gpu_indexer_topk_tensor(g->comp_selected, - g->indexer_scores, + ok = ds4_gpu_indexer_topk_tensor(g->scr->comp_selected, + g->scr->indexer_scores, n_comp, n_tokens, DS4_N_INDEXER_TOP_K) != 0; @@ -18489,22 +18589,22 @@ static bool metal_graph_encode_layer_attention_batch( } if (ok) { metal_graph_debug_dump_i32_tensor("indexer_topk", - g->comp_selected, + g->scr->comp_selected, (uint64_t)n_tokens * DS4_N_INDEXER_TOP_K, il, pos0); } } if (ok) { - ok = ds4_gpu_attention_indexed_mixed_batch_heads_tensor(g->batch_heads, + ok = ds4_gpu_attention_indexed_mixed_batch_heads_tensor(g->scr->batch_heads, model->map, model->size, layer->attn_sinks->abs_offset, - g->batch_q, + g->scr->batch_q, g->layer_raw_cache[il], g->layer_attn_comp_cache[il], metal_graph_attn_comp_cache_is_f16(), - g->comp_selected, + g->scr->comp_selected, n_tokens, pos0, n_tokens, @@ -18528,12 +18628,12 @@ static bool metal_graph_encode_layer_attention_batch( if (ok) batch_attention_done = true; } if (ok && zero_prefix && !topk_prefill_needed && n_comp != 0) { - ok = ds4_gpu_attention_prefill_static_mixed_heads_tensor(g->batch_heads, + ok = ds4_gpu_attention_prefill_static_mixed_heads_tensor(g->scr->batch_heads, model->map, model->size, layer->attn_sinks->abs_offset, - g->batch_q, - g->batch_kv, + g->scr->batch_q, + g->scr->batch_kv, g->layer_attn_comp_cache[il], metal_graph_attn_comp_cache_is_f16(), n_tokens, @@ -18555,12 +18655,12 @@ static bool metal_graph_encode_layer_attention_batch( } if (raw_prefix_tokens != 0) { - ok = ds4_gpu_attention_prefill_raw_heads_tensor(g->batch_heads, + ok = ds4_gpu_attention_prefill_raw_heads_tensor(g->scr->batch_heads, model->map, model->size, layer->attn_sinks->abs_offset, - g->batch_q, - g->batch_kv, + g->scr->batch_q, + g->scr->batch_kv, raw_prefix_tokens, g->raw_window, DS4_N_HEAD, @@ -18579,11 +18679,11 @@ static bool metal_graph_encode_layer_attention_batch( if (ratio == 4 && cur_comp > DS4_N_INDEXER_TOP_K) { const float index_scale = 1.0f / sqrtf((float)(DS4_N_INDEXER_HEAD_DIM * DS4_N_INDEXER_HEAD)); ds4_gpu_tensor *indexer_q_view = metal_graph_tensor_row_view( - g->batch_indexer_q, t, (uint64_t)DS4_N_INDEXER_HEAD * DS4_N_INDEXER_HEAD_DIM); + g->scr->batch_indexer_q, t, (uint64_t)DS4_N_INDEXER_HEAD * DS4_N_INDEXER_HEAD_DIM); ds4_gpu_tensor *indexer_w_view = metal_graph_tensor_row_view( - g->batch_indexer_weights, t, DS4_N_INDEXER_HEAD); + g->scr->batch_indexer_weights, t, DS4_N_INDEXER_HEAD); ok = indexer_q_view && indexer_w_view && - ds4_gpu_indexer_score_one_tensor(g->indexer_scores, + ds4_gpu_indexer_score_one_tensor(g->scr->indexer_scores, indexer_q_view, indexer_w_view, g->layer_index_comp_cache[il], @@ -18591,29 +18691,29 @@ static bool metal_graph_encode_layer_attention_batch( DS4_N_INDEXER_HEAD, DS4_N_INDEXER_HEAD_DIM, index_scale) != 0 && - ds4_gpu_indexer_topk_tensor(g->comp_selected, - g->indexer_scores, + ds4_gpu_indexer_topk_tensor(g->scr->comp_selected, + g->scr->indexer_scores, cur_index, 1, DS4_N_INDEXER_TOP_K) != 0 && - ds4_gpu_dsv4_topk_mask_tensor(g->comp_mask, - g->comp_selected, + ds4_gpu_dsv4_topk_mask_tensor(g->scr->comp_mask, + g->scr->comp_selected, cur_index, 1, DS4_N_INDEXER_TOP_K) != 0; ds4_gpu_tensor_free(indexer_w_view); ds4_gpu_tensor_free(indexer_q_view); if (ok) { - comp_mask = g->comp_mask; + comp_mask = g->scr->comp_mask; n_selected = DS4_N_INDEXER_TOP_K < cur_index ? DS4_N_INDEXER_TOP_K : cur_index; } } - ds4_gpu_tensor *q_view = metal_graph_tensor_row_view(g->batch_q, t, q_dim); - ds4_gpu_tensor *kv_cache_view = metal_graph_tensor_row_view(g->batch_kv, t, DS4_N_HEAD_DIM); - ds4_gpu_tensor *heads_view = metal_graph_tensor_row_view(g->batch_heads, t, q_dim); + ds4_gpu_tensor *q_view = metal_graph_tensor_row_view(g->scr->batch_q, t, q_dim); + ds4_gpu_tensor *kv_cache_view = metal_graph_tensor_row_view(g->scr->batch_kv, t, DS4_N_HEAD_DIM); + ds4_gpu_tensor *heads_view = metal_graph_tensor_row_view(g->scr->batch_heads, t, q_dim); ok = ok && q_view && kv_cache_view && heads_view; if (ok && !zero_prefix) { ok = ds4_gpu_store_raw_kv_tensor(g->layer_raw_cache[il], @@ -18631,7 +18731,7 @@ static bool metal_graph_encode_layer_attention_batch( g->layer_raw_cache[il], g->layer_attn_comp_cache[il], metal_graph_attn_comp_cache_is_f16(), - g->comp_selected, + g->scr->comp_selected, 1, pos, n_raw, @@ -18670,10 +18770,10 @@ static bool metal_graph_encode_layer_attention_batch( DS4_METAL_PROFILE_ATTN_STAGE("attention"); if (ok) { - metal_graph_debug_dump_tensor("kqv_out", g->batch_heads, + metal_graph_debug_dump_tensor("kqv_out", g->scr->batch_heads, (uint64_t)n_tokens * q_dim, il, pos0); } - if (ok) ok = ds4_gpu_rope_tail_tensor(g->batch_heads, + if (ok) ok = ds4_gpu_rope_tail_tensor(g->scr->batch_heads, n_tokens, DS4_N_HEAD, DS4_N_HEAD_DIM, @@ -18688,7 +18788,7 @@ static bool metal_graph_encode_layer_attention_batch( DS4_ROPE_YARN_BETA_FAST, DS4_ROPE_YARN_BETA_SLOW) != 0; if (ok) { - metal_graph_debug_dump_tensor("kqv_back", g->batch_heads, + metal_graph_debug_dump_tensor("kqv_back", g->scr->batch_heads, (uint64_t)n_tokens * q_dim, il, pos0); } DS4_METAL_PROFILE_ATTN_STAGE("inv_rope"); @@ -18699,8 +18799,8 @@ static bool metal_graph_encode_layer_attention_batch( if (ok && !attn_out_debug && !metal_graph_directional_steering_attn_enabled(g)) { - attn_out_f16 = ds4_gpu_attention_output_q8_batch_f16_tensor(g->batch_q_half, - g->batch_attn_low, + attn_out_f16 = ds4_gpu_attention_output_q8_batch_f16_tensor(g->scr->batch_q_half, + g->scr->batch_attn_low, model->map, model->size, layer->attn_output_a->abs_offset, @@ -18709,15 +18809,15 @@ static bool metal_graph_encode_layer_attention_batch( rank, n_groups, DS4_N_EMBD, - g->batch_heads, + g->scr->batch_heads, n_tokens) != 0; } if (!attn_out_f16) { if (ok) { - ok = ds4_gpu_attention_output_q8_batch_tensor(g->batch_attn_out, - g->batch_attn_low, - g->batch_group_tmp, - g->batch_low_tmp, + ok = ds4_gpu_attention_output_q8_batch_tensor(g->scr->batch_attn_out, + g->scr->batch_attn_low, + g->scr->batch_group_tmp, + g->scr->batch_low_tmp, model->map, model->size, layer->attn_output_a->abs_offset, @@ -18726,41 +18826,41 @@ static bool metal_graph_encode_layer_attention_batch( rank, n_groups, DS4_N_EMBD, - g->batch_heads, + g->scr->batch_heads, n_tokens) != 0; } if (ok) { - metal_graph_debug_dump_tensor("attn_low", g->batch_attn_low, + metal_graph_debug_dump_tensor("attn_low", g->scr->batch_attn_low, (uint64_t)n_tokens * n_groups * rank, il, pos0); } if (ok) { - metal_graph_debug_dump_tensor("attn_out", g->batch_attn_out, + metal_graph_debug_dump_tensor("attn_out", g->scr->batch_attn_out, (uint64_t)n_tokens * DS4_N_EMBD, il, pos0); } } DS4_METAL_PROFILE_ATTN_STAGE("output_proj"); if (ok && !attn_out_f16 && metal_graph_directional_steering_attn_enabled(g)) { - ok = metal_graph_apply_directional_steering_attn(g, g->batch_attn_out, il, n_tokens); + ok = metal_graph_apply_directional_steering_attn(g, g->scr->batch_attn_out, il, n_tokens); } if (ok && attn_out_f16) { ok = ds4_gpu_hc_expand_split_half_tensor(after_attn_hc_view, - g->batch_q_half, - g->batch_cur_hc, + g->scr->batch_q_half, + g->scr->batch_cur_hc, hc_split_view, DS4_N_EMBD, DS4_N_HC) != 0; } else if (ok) { ok = ds4_gpu_hc_expand_split_tensor(after_attn_hc_view, - g->batch_attn_out, - g->batch_cur_hc, + g->scr->batch_attn_out, + g->scr->batch_cur_hc, hc_split_view, DS4_N_EMBD, DS4_N_HC) != 0; } if (ok) { - metal_graph_debug_dump_tensor("hc_attn_post", g->batch_after_attn_hc, + metal_graph_debug_dump_tensor("hc_attn_post", g->scr->batch_after_attn_hc, (uint64_t)n_tokens * hc_dim, il, pos0); } DS4_METAL_PROFILE_ATTN_STAGE("hc_post"); @@ -18806,19 +18906,19 @@ static bool metal_graph_encode_layer_ffn_batch( } while (0) ds4_gpu_tensor *hc_mix_view = ds4_gpu_tensor_view( - g->batch_hc_mix, 0, (uint64_t)n_tokens * mix_hc * sizeof(float)); + g->scr->batch_hc_mix, 0, (uint64_t)n_tokens * mix_hc * sizeof(float)); ds4_gpu_tensor *hc_split_view = ds4_gpu_tensor_view( - g->batch_hc_split, 0, (uint64_t)n_tokens * mix_hc * sizeof(float)); + g->scr->batch_hc_split, 0, (uint64_t)n_tokens * mix_hc * sizeof(float)); ds4_gpu_tensor *ffn_cur_view = ds4_gpu_tensor_view( - g->batch_ffn_cur, 0, (uint64_t)n_tokens * DS4_N_EMBD * sizeof(float)); + g->scr->batch_ffn_cur, 0, (uint64_t)n_tokens * DS4_N_EMBD * sizeof(float)); ds4_gpu_tensor *next_hc_view = ds4_gpu_tensor_view( - g->batch_next_hc, 0, (uint64_t)n_tokens * hc_dim * sizeof(float)); + g->scr->batch_next_hc, 0, (uint64_t)n_tokens * hc_dim * sizeof(float)); bool ok = hc_mix_view && hc_split_view && ffn_cur_view && next_hc_view; const bool fuse_hc_norm = DS4_N_HC == 4 && !metal_graph_use_reference_hc_decode() && metal_graph_enable_batch_hc_norm_fusion(); - if (ok) ok = ds4_gpu_rms_norm_plain_rows_tensor(g->batch_flat_hc, - g->batch_after_attn_hc, + if (ok) ok = ds4_gpu_rms_norm_plain_rows_tensor(g->scr->batch_flat_hc, + g->scr->batch_after_attn_hc, (uint32_t)hc_dim, n_tokens, DS4_RMS_EPS) != 0; @@ -18828,7 +18928,7 @@ static bool metal_graph_encode_layer_ffn_batch( layer->hc_ffn_fn->abs_offset, hc_dim, mix_hc, - g->batch_flat_hc, + g->scr->batch_flat_hc, n_tokens) != 0; if (metal_graph_use_reference_hc_decode()) { if (ok) ok = ds4_gpu_hc_split_sinkhorn_tensor(hc_split_view, @@ -18841,16 +18941,16 @@ static bool metal_graph_encode_layer_ffn_batch( DS4_N_HC_SINKHORN_ITER, DS4_HC_EPS) != 0; if (ok) ok = ds4_gpu_hc_weighted_sum_split_tensor(ffn_cur_view, - g->batch_after_attn_hc, + g->scr->batch_after_attn_hc, hc_split_view, DS4_N_EMBD, DS4_N_HC) != 0; } else if (fuse_hc_norm) { if (ok) ok = ds4_gpu_hc_split_weighted_sum_norm_tensor(ffn_cur_view, - g->batch_ffn_norm, + g->scr->batch_ffn_norm, hc_split_view, hc_mix_view, - g->batch_after_attn_hc, + g->scr->batch_after_attn_hc, model->map, model->size, layer->hc_ffn_scale->abs_offset, @@ -18865,7 +18965,7 @@ static bool metal_graph_encode_layer_ffn_batch( if (ok) ok = ds4_gpu_hc_split_weighted_sum_tensor(ffn_cur_view, hc_split_view, hc_mix_view, - g->batch_after_attn_hc, + g->scr->batch_after_attn_hc, model->map, model->size, layer->hc_ffn_scale->abs_offset, @@ -18876,13 +18976,13 @@ static bool metal_graph_encode_layer_ffn_batch( DS4_HC_EPS) != 0; } if (ok) { - metal_graph_debug_dump_tensor("hc_ffn_pre", g->batch_ffn_cur, + metal_graph_debug_dump_tensor("hc_ffn_pre", g->scr->batch_ffn_cur, (uint64_t)n_tokens * DS4_N_EMBD, il, pos0); } DS4_METAL_PROFILE_FFN_STAGE("hc_pre"); if (ok && !fuse_hc_norm) { - ok = ds4_gpu_rms_norm_weight_rows_tensor(g->batch_ffn_norm, - g->batch_ffn_cur, + ok = ds4_gpu_rms_norm_weight_rows_tensor(g->scr->batch_ffn_norm, + g->scr->batch_ffn_cur, model->map, model->size, layer->ffn_norm->abs_offset, @@ -18891,22 +18991,22 @@ static bool metal_graph_encode_layer_ffn_batch( DS4_RMS_EPS) != 0; } if (ok) { - metal_graph_debug_dump_tensor("ffn_norm", g->batch_ffn_norm, + metal_graph_debug_dump_tensor("ffn_norm", g->scr->batch_ffn_norm, (uint64_t)n_tokens * DS4_N_EMBD, il, pos0); } DS4_METAL_PROFILE_FFN_STAGE("norm"); - if (ok) ok = ds4_gpu_matmul_f16_tensor(g->batch_router_logits, + if (ok) ok = ds4_gpu_matmul_f16_tensor(g->scr->batch_router_logits, model->map, model->size, layer->ffn_gate_inp->abs_offset, DS4_N_EMBD, DS4_N_EXPERT, - g->batch_ffn_norm, + g->scr->batch_ffn_norm, n_tokens) != 0; - if (ok) ok = ds4_gpu_router_select_batch_tensor(g->batch_router_selected, - g->batch_router_weights, - g->batch_router_probs, + if (ok) ok = ds4_gpu_router_select_batch_tensor(g->scr->batch_router_selected, + g->scr->batch_router_weights, + g->scr->batch_router_probs, model->map, model->size, layer->ffn_exp_probs_b ? layer->ffn_exp_probs_b->abs_offset : 0, @@ -18916,20 +19016,20 @@ static bool metal_graph_encode_layer_ffn_batch( 0, layer->ffn_exp_probs_b != NULL, layer->ffn_gate_tid2eid != NULL, - g->batch_router_logits, - g->prefill_tokens, + g->scr->batch_router_logits, + g->scr->prefill_tokens, DS4_N_EXPERT, DS4_N_EXPERT_USED, DS4_EXPERT_WEIGHT_SCALE, n_tokens) != 0; if (ok) { - metal_graph_debug_dump_tensor("ffn_moe_logits", g->batch_router_logits, + metal_graph_debug_dump_tensor("ffn_moe_logits", g->scr->batch_router_logits, (uint64_t)n_tokens * DS4_N_EXPERT, il, pos0); - metal_graph_debug_dump_tensor("ffn_moe_probs", g->batch_router_probs, + metal_graph_debug_dump_tensor("ffn_moe_probs", g->scr->batch_router_probs, (uint64_t)n_tokens * DS4_N_EXPERT, il, pos0); - metal_graph_debug_dump_i32_tensor("ffn_moe_topk", g->batch_router_selected, + metal_graph_debug_dump_i32_tensor("ffn_moe_topk", g->scr->batch_router_selected, (uint64_t)n_tokens * DS4_N_EXPERT_USED, il, pos0); - metal_graph_debug_dump_tensor("ffn_moe_weights_scaled", g->batch_router_weights, + metal_graph_debug_dump_tensor("ffn_moe_weights_scaled", g->scr->batch_router_weights, (uint64_t)n_tokens * DS4_N_EXPERT_USED, il, pos0); } DS4_METAL_PROFILE_FFN_STAGE("router"); @@ -19008,13 +19108,13 @@ static bool metal_graph_encode_layer_ffn_batch( #define DS4_METAL_TRY_SHARED_DOWN_F16() do { \ if (ok && !keep_ffn_out && !metal_graph_debug_wants("ffn_shexp", il, pos0)) { \ - shared_down_f16 = ds4_gpu_matmul_q8_0_f16_out_tensor(g->batch_q_half, \ + shared_down_f16 = ds4_gpu_matmul_q8_0_f16_out_tensor(g->scr->batch_q_half, \ model->map, \ model->size, \ layer->ffn_down_shexp->abs_offset, \ shared_dim, \ DS4_N_EMBD, \ - g->batch_shared_mid, \ + g->scr->batch_shared_mid, \ n_tokens) != 0; \ } \ } while (0) @@ -19023,27 +19123,27 @@ static bool metal_graph_encode_layer_ffn_batch( if (ok) ok = metal_graph_matmul_q8_0_named_tensor("shared_gate", \ il, \ pos0, \ - g->batch_shared_gate, \ + g->scr->batch_shared_gate, \ model, \ layer->ffn_gate_shexp, \ DS4_N_EMBD, \ shared_dim, \ - g->batch_ffn_norm, \ + g->scr->batch_ffn_norm, \ n_tokens); \ if (ok) ok = metal_graph_matmul_q8_0_named_tensor("shared_up", \ il, \ pos0, \ - g->batch_shared_up, \ + g->scr->batch_shared_up, \ model, \ layer->ffn_up_shexp, \ DS4_N_EMBD, \ shared_dim, \ - g->batch_ffn_norm, \ + g->scr->batch_ffn_norm, \ n_tokens); \ DS4_METAL_PROFILE_FFN_STAGE("shared_gate_up"); \ - if (ok) ok = ds4_gpu_swiglu_tensor(g->batch_shared_mid, \ - g->batch_shared_gate, \ - g->batch_shared_up, \ + if (ok) ok = ds4_gpu_swiglu_tensor(g->scr->batch_shared_mid, \ + g->scr->batch_shared_gate, \ + g->scr->batch_shared_up, \ (uint32_t)((uint64_t)n_tokens * shared_dim), \ DS4_SWIGLU_CLAMP_EXP, \ 1.0f) != 0; \ @@ -19051,16 +19151,16 @@ static bool metal_graph_encode_layer_ffn_batch( if (ok && !shared_down_f16) ok = metal_graph_matmul_q8_0_named_tensor("shared_down", \ il, \ pos0, \ - g->batch_shared_out, \ + g->scr->batch_shared_out, \ model, \ layer->ffn_down_shexp, \ shared_dim, \ DS4_N_EMBD, \ - g->batch_shared_mid, \ + g->scr->batch_shared_mid, \ n_tokens); \ DS4_METAL_PROFILE_FFN_STAGE("shared_down"); \ if (ok && !shared_down_f16) { \ - metal_graph_debug_dump_tensor("ffn_shexp", g->batch_shared_out, \ + metal_graph_debug_dump_tensor("ffn_shexp", g->scr->batch_shared_out, \ (uint64_t)n_tokens * DS4_N_EMBD, il, pos0); \ } \ } while (0) @@ -19144,11 +19244,11 @@ static bool metal_graph_encode_layer_ffn_batch( #endif if (ok) { - ok = ds4_gpu_routed_moe_batch_tensor(g->batch_routed_out, - g->batch_routed_gate, - g->batch_routed_up, - g->batch_routed_mid, - g->batch_routed_down, + ok = ds4_gpu_routed_moe_batch_tensor(g->scr->batch_routed_out, + g->scr->batch_routed_gate, + g->scr->batch_routed_up, + g->scr->batch_routed_mid, + g->scr->batch_routed_down, model->map, model->size, layer->ffn_gate_exps->abs_offset, @@ -19163,38 +19263,38 @@ static bool metal_graph_encode_layer_ffn_batch( (uint32_t)expert_in_dim, (uint32_t)down_in_dim, (uint32_t)routed_out_dim, - g->batch_router_selected, - g->batch_router_weights, + g->scr->batch_router_selected, + g->scr->batch_router_weights, DS4_N_EXPERT, DS4_N_EXPERT_USED, DS4_SWIGLU_CLAMP_EXP, - g->batch_ffn_norm, + g->scr->batch_ffn_norm, il, n_tokens, - &g->batch_routed_mid_is_f16) != 0; + &g->scr->batch_routed_mid_is_f16) != 0; } if (ok) { - metal_graph_debug_dump_tensor("ffn_moe_gate_clamped", g->batch_routed_gate, + metal_graph_debug_dump_tensor("ffn_moe_gate_clamped", g->scr->batch_routed_gate, (uint64_t)n_tokens * DS4_N_EXPERT_USED * down_in_dim, il, pos0); - metal_graph_debug_dump_tensor("ffn_moe_up_clamped", g->batch_routed_up, + metal_graph_debug_dump_tensor("ffn_moe_up_clamped", g->scr->batch_routed_up, (uint64_t)n_tokens * DS4_N_EXPERT_USED * down_in_dim, il, pos0); } if (ok) { const uint64_t routed_mid_elems = (uint64_t)n_tokens * DS4_N_EXPERT_USED * down_in_dim; - if (g->batch_routed_mid_is_f16) { - metal_graph_debug_dump_f16_tensor("ffn_moe_weighted_swiglu", g->batch_routed_mid, + if (g->scr->batch_routed_mid_is_f16) { + metal_graph_debug_dump_f16_tensor("ffn_moe_weighted_swiglu", g->scr->batch_routed_mid, routed_mid_elems, il, pos0); } else { - metal_graph_debug_dump_tensor("ffn_moe_weighted_swiglu", g->batch_routed_mid, + metal_graph_debug_dump_tensor("ffn_moe_weighted_swiglu", g->scr->batch_routed_mid, routed_mid_elems, il, pos0); } } if (ok) { - metal_graph_debug_dump_tensor("ffn_moe_down", g->batch_routed_down, + metal_graph_debug_dump_tensor("ffn_moe_down", g->scr->batch_routed_down, (uint64_t)n_tokens * DS4_N_EXPERT_USED * DS4_N_EMBD, il, pos0); } if (ok) { - metal_graph_debug_dump_tensor("ffn_moe_out", g->batch_routed_out, + metal_graph_debug_dump_tensor("ffn_moe_out", g->scr->batch_routed_out, (uint64_t)n_tokens * DS4_N_EMBD, il, pos0); } DS4_METAL_PROFILE_FFN_STAGE("routed_moe"); @@ -19207,8 +19307,8 @@ static bool metal_graph_encode_layer_ffn_batch( if (ok && keep_ffn_out) { ok = metal_graph_ensure_batch_ffn_out(g) && ds4_gpu_add_tensor(g->batch_ffn_out, - g->batch_shared_out, - g->batch_routed_out, + g->scr->batch_shared_out, + g->scr->batch_routed_out, (uint32_t)((uint64_t)n_tokens * DS4_N_EMBD)) != 0; } if (ok && keep_ffn_out) { @@ -19221,31 +19321,31 @@ static bool metal_graph_encode_layer_ffn_batch( if (ok && metal_graph_directional_steering_ffn_enabled(g)) { ok = ds4_gpu_hc_expand_split_tensor(next_hc_view, g->batch_ffn_out, - g->batch_after_attn_hc, + g->scr->batch_after_attn_hc, hc_split_view, DS4_N_EMBD, DS4_N_HC) != 0; } else if (ok && shared_down_f16) { ok = ds4_gpu_hc_expand_add_split_half_add_tensor(next_hc_view, - g->batch_routed_out, - g->batch_q_half, - g->batch_after_attn_hc, + g->scr->batch_routed_out, + g->scr->batch_q_half, + g->scr->batch_after_attn_hc, hc_split_view, DS4_N_EMBD, DS4_N_HC) != 0; } else if (ok) { ok = ds4_gpu_hc_expand_add_split_tensor(next_hc_view, - g->batch_routed_out, - g->batch_shared_out, - g->batch_after_attn_hc, + g->scr->batch_routed_out, + g->scr->batch_shared_out, + g->scr->batch_after_attn_hc, hc_split_view, DS4_N_EMBD, DS4_N_HC) != 0; } if (ok) { - metal_graph_debug_dump_tensor("hc_ffn_post", g->batch_next_hc, + metal_graph_debug_dump_tensor("hc_ffn_post", g->scr->batch_next_hc, (uint64_t)n_tokens * hc_dim, il, pos0); } DS4_METAL_PROFILE_FFN_STAGE("hc_post"); @@ -19276,9 +19376,9 @@ static bool metal_graph_encode_layer_batch( } } if (ok) { - ds4_gpu_tensor *tmp = g->batch_cur_hc; - g->batch_cur_hc = g->batch_next_hc; - g->batch_next_hc = tmp; + ds4_gpu_tensor *tmp = g->scr->batch_cur_hc; + g->scr->batch_cur_hc = g->scr->batch_next_hc; + g->scr->batch_next_hc = tmp; } return ok; } @@ -19321,7 +19421,7 @@ static bool metal_graph_eval_token_raw_swa_streaming( } if (ok) ok = ds4_gpu_begin_commands() != 0; if (ok) { - ok = ds4_gpu_embed_token_hc_tensor(g->cur_hc, + ok = ds4_gpu_embed_token_hc_tensor(g->scr->cur_hc, model->map, model->size, weights->token_embd->abs_offset, @@ -19343,9 +19443,9 @@ static bool metal_graph_eval_token_raw_swa_streaming( n_raw, token); if (ok) { - ds4_gpu_tensor *tmp = g->cur_hc; - g->cur_hc = g->after_ffn_hc; - g->after_ffn_hc = tmp; + ds4_gpu_tensor *tmp = g->scr->cur_hc; + g->scr->cur_hc = g->scr->after_ffn_hc; + g->scr->after_ffn_hc = tmp; } } if (ok && logits) { @@ -19355,7 +19455,7 @@ static bool metal_graph_eval_token_raw_swa_streaming( if (ok) ok = ds4_gpu_end_commands() != 0; const double t_done = (profile || throttle) ? now_sec() : 0.0; if (ok && logits) { - ok = ds4_gpu_tensor_read(g->logits, 0, logits, (uint64_t)DS4_N_VOCAB * sizeof(float)) != 0; + ok = ds4_gpu_tensor_read(g->scr->logits, 0, logits, (uint64_t)DS4_N_VOCAB * sizeof(float)) != 0; } const double t_read = (profile || throttle) ? now_sec() : 0.0; if (profile) { @@ -19409,9 +19509,9 @@ static bool metal_graph_eval_token_raw_swa_streaming( encoded_layer = true; } if (encoded_layer) { - ds4_gpu_tensor *tmp = g->cur_hc; - g->cur_hc = g->after_ffn_hc; - g->after_ffn_hc = tmp; + ds4_gpu_tensor *tmp = g->scr->cur_hc; + g->scr->cur_hc = g->scr->after_ffn_hc; + g->scr->after_ffn_hc = tmp; } const double tl_encoded = profile ? now_sec() : 0.0; if (ok) ok = ds4_gpu_end_commands() != 0; @@ -19430,7 +19530,7 @@ static bool metal_graph_eval_token_raw_swa_streaming( if (ok && logits) ok = ds4_gpu_end_commands() != 0; const double t_done = (profile || throttle) ? now_sec() : 0.0; if (ok && logits) { - ok = ds4_gpu_tensor_read(g->logits, 0, logits, (uint64_t)DS4_N_VOCAB * sizeof(float)) != 0; + ok = ds4_gpu_tensor_read(g->scr->logits, 0, logits, (uint64_t)DS4_N_VOCAB * sizeof(float)) != 0; } const double t_read = (profile || throttle) ? now_sec() : 0.0; @@ -19480,7 +19580,7 @@ static bool metal_graph_eval_token_raw_swa( const double t_done = (profile || throttle) ? now_sec() : 0.0; if (ok && logits) { - ok = ds4_gpu_tensor_read(g->logits, 0, logits, (uint64_t)DS4_N_VOCAB * sizeof(float)) != 0; + ok = ds4_gpu_tensor_read(g->scr->logits, 0, logits, (uint64_t)DS4_N_VOCAB * sizeof(float)) != 0; } const double t_read = (profile || throttle) ? now_sec() : 0.0; if (profile) { @@ -19661,8 +19761,8 @@ static bool metal_graph_capture_prefill_seed_router_selected( uint32_t k = metal_graph_streaming_prefill_cache_seed_k(g); if (k == 0) return true; if (k > n_tokens) k = n_tokens; - g->prefill_seed_tokens = k; - if (!g->prefill_seed_router_selected || !g->batch_router_selected || + g->scr->prefill_seed_tokens = k; + if (!g->scr->prefill_seed_router_selected || !g->scr->batch_router_selected || il >= DS4_N_LAYER || n_tokens == 0 || sizeof(int) != sizeof(int32_t)) { return false; } @@ -19673,9 +19773,9 @@ static bool metal_graph_capture_prefill_seed_router_selected( const uint64_t dst_off = (uint64_t)il * DS4_STREAMING_PREFILL_CACHE_SEED_MAX_TOKENS * DS4_N_EXPERT_USED * sizeof(int32_t); - return ds4_gpu_tensor_copy(g->prefill_seed_router_selected, + return ds4_gpu_tensor_copy(g->scr->prefill_seed_router_selected, dst_off, - g->batch_router_selected, + g->scr->batch_router_selected, src_off, bytes) != 0; } @@ -19684,9 +19784,9 @@ static bool metal_graph_seed_streaming_expert_cache_from_prefill( ds4_gpu_graph *g, const ds4_model *model, const ds4_weights *weights) { - const uint32_t seed_tokens = g ? g->prefill_seed_tokens : 0; + const uint32_t seed_tokens = g ? g->scr->prefill_seed_tokens : 0; if (!metal_graph_streaming_prefill_cache_seed_enabled(g)) return true; - if (!model || !weights || !g->prefill_seed_router_selected || seed_tokens == 0) { + if (!model || !weights || !g->scr->prefill_seed_router_selected || seed_tokens == 0) { return false; } @@ -19696,7 +19796,7 @@ static bool metal_graph_seed_streaming_expert_cache_from_prefill( const uint64_t bytes = (uint64_t)DS4_N_LAYER * DS4_STREAMING_PREFILL_CACHE_SEED_MAX_TOKENS * DS4_N_EXPERT_USED * sizeof(selected[0]); - if (ds4_gpu_tensor_read(g->prefill_seed_router_selected, + if (ds4_gpu_tensor_read(g->scr->prefill_seed_router_selected, 0, selected, bytes) == 0) { @@ -19904,14 +20004,14 @@ static bool metal_graph_eval_token_raw_swa_top( if (ok) ok = metal_graph_encode_token_raw_swa(g, model, weights, token, pos, true, true); if (ok) { - ok = ds4_gpu_argmax_tensor(g->comp_selected, - g->logits, + ok = ds4_gpu_argmax_tensor(g->scr->comp_selected, + g->scr->logits, DS4_N_VOCAB) != 0; } if (ok) ok = ds4_gpu_end_commands() != 0; - if (ok) ok = ds4_gpu_tensor_read(g->comp_selected, 0, top_id, sizeof(*top_id)) != 0; + if (ok) ok = ds4_gpu_tensor_read(g->scr->comp_selected, 0, top_id, sizeof(*top_id)) != 0; if (ok && logits) { - ok = ds4_gpu_tensor_read(g->logits, 0, logits, (uint64_t)DS4_N_VOCAB * sizeof(float)) != 0; + ok = ds4_gpu_tensor_read(g->scr->logits, 0, logits, (uint64_t)DS4_N_VOCAB * sizeof(float)) != 0; } if (!ok) { if (ds4_gpu_synchronize() == 0) { @@ -19941,10 +20041,10 @@ static bool metal_graph_eval_mtp_draft_from_hc( if (n_raw > g->raw_window) n_raw = g->raw_window; if (n_raw > g->raw_cap) n_raw = g->raw_cap; - ds4_gpu_tensor *saved_cur = g->cur_hc; - ds4_gpu_tensor *saved_after = g->after_ffn_hc; + ds4_gpu_tensor *saved_cur = g->scr->cur_hc; + ds4_gpu_tensor *saved_after = g->scr->after_ffn_hc; bool ok = ds4_gpu_begin_commands() != 0; - if (ok) ok = ds4_gpu_embed_token_hc_tensor(g->mtp_embed, + if (ok) ok = ds4_gpu_embed_token_hc_tensor(g->scr->mtp_embed, base_model->map, base_model->size, base_weights->token_embd->abs_offset, @@ -19952,26 +20052,26 @@ static bool metal_graph_eval_mtp_draft_from_hc( (uint32_t)token, DS4_N_EMBD, 1) != 0; - if (ok) ok = ds4_gpu_rms_norm_weight_tensor(g->mtp_enorm, - g->mtp_embed, + if (ok) ok = ds4_gpu_rms_norm_weight_tensor(g->scr->mtp_enorm, + g->scr->mtp_embed, mtp_model->map, mtp_model->size, mtp->enorm->abs_offset, DS4_N_EMBD, DS4_RMS_EPS) != 0; - if (ok) ok = ds4_gpu_matmul_q8_0_tensor(g->mtp_eproj, + if (ok) ok = ds4_gpu_matmul_q8_0_tensor(g->scr->mtp_eproj, mtp_model->map, mtp_model->size, mtp->e_proj->abs_offset, DS4_N_EMBD, DS4_N_EMBD, - g->mtp_enorm, + g->scr->mtp_enorm, 1) != 0; - if (ok) ok = ds4_gpu_repeat_hc_tensor(g->mtp_eproj_hc, - g->mtp_eproj, + if (ok) ok = ds4_gpu_repeat_hc_tensor(g->scr->mtp_eproj_hc, + g->scr->mtp_eproj, DS4_N_EMBD, DS4_N_HC) != 0; - if (ok) ok = ds4_gpu_rms_norm_weight_rows_tensor(g->mtp_hnorm_hc, + if (ok) ok = ds4_gpu_rms_norm_weight_rows_tensor(g->scr->mtp_hnorm_hc, prev_hc, mtp_model->map, mtp_model->size, @@ -19979,21 +20079,21 @@ static bool metal_graph_eval_mtp_draft_from_hc( DS4_N_EMBD, DS4_N_HC, DS4_RMS_EPS) != 0; - if (ok) ok = ds4_gpu_matmul_q8_0_tensor(g->mtp_hproj_hc, + if (ok) ok = ds4_gpu_matmul_q8_0_tensor(g->scr->mtp_hproj_hc, mtp_model->map, mtp_model->size, mtp->h_proj->abs_offset, DS4_N_EMBD, DS4_N_EMBD, - g->mtp_hnorm_hc, + g->scr->mtp_hnorm_hc, DS4_N_HC) != 0; - if (ok) ok = ds4_gpu_add_tensor(g->mtp_input_hc, - g->mtp_eproj_hc, - g->mtp_hproj_hc, + if (ok) ok = ds4_gpu_add_tensor(g->scr->mtp_input_hc, + g->scr->mtp_eproj_hc, + g->scr->mtp_hproj_hc, (uint32_t)hc_dim) != 0; if (ok) { - g->cur_hc = g->mtp_input_hc; - g->after_ffn_hc = out_hc; + g->scr->cur_hc = g->scr->mtp_input_hc; + g->scr->after_ffn_hc = out_hc; ok = metal_graph_encode_decode_layer(g, mtp_model, &mtp->block, @@ -20005,7 +20105,7 @@ static bool metal_graph_eval_mtp_draft_from_hc( n_raw, token); } - if (ok) g->cur_hc = out_hc; + if (ok) g->scr->cur_hc = out_hc; if (ok) ok = metal_graph_encode_output_head_mtp(g, base_model, base_weights, @@ -20013,25 +20113,25 @@ static bool metal_graph_eval_mtp_draft_from_hc( mtp, base_weights->output->dim[1]); if (ok && top_id) { - ok = ds4_gpu_argmax_tensor(g->comp_selected, - g->logits, + ok = ds4_gpu_argmax_tensor(g->scr->comp_selected, + g->scr->logits, DS4_N_VOCAB) != 0; } if (ok) ok = ds4_gpu_end_commands() != 0; - g->cur_hc = saved_cur; - g->after_ffn_hc = saved_after; + g->scr->cur_hc = saved_cur; + g->scr->after_ffn_hc = saved_after; if (ok && logits) { - ok = ds4_gpu_tensor_read(g->logits, 0, logits, (uint64_t)DS4_N_VOCAB * sizeof(float)) != 0; + ok = ds4_gpu_tensor_read(g->scr->logits, 0, logits, (uint64_t)DS4_N_VOCAB * sizeof(float)) != 0; } if (ok && top_id) { - ok = ds4_gpu_tensor_read(g->comp_selected, 0, top_id, sizeof(*top_id)) != 0; + ok = ds4_gpu_tensor_read(g->scr->comp_selected, 0, top_id, sizeof(*top_id)) != 0; } if (ok && g->mtp_n_raw < g->raw_window) g->mtp_n_raw++; if (!ok) { (void)ds4_gpu_synchronize(); - g->cur_hc = saved_cur; - g->after_ffn_hc = saved_after; + g->scr->cur_hc = saved_cur; + g->scr->after_ffn_hc = saved_after; } return ok; } @@ -20051,8 +20151,8 @@ static bool metal_graph_eval_mtp_draft( base_weights, mtp_model, mtp, - g->cur_hc, - g->mtp_state_hc, + g->scr->cur_hc, + g->scr->mtp_state_hc, token, pos, logits, @@ -20138,14 +20238,14 @@ static bool imatrix_collect_layer_batch( const uint64_t norm_bytes = (uint64_t)n_tokens * DS4_N_EMBD * sizeof(float); const uint64_t mid_elems = (uint64_t)n_tokens * DS4_N_EXPERT_USED * DS4_N_FF_EXP; - const uint64_t mid_bytes = mid_elems * (g->batch_routed_mid_is_f16 ? sizeof(uint16_t) : sizeof(float)); + const uint64_t mid_bytes = mid_elems * (g->scr->batch_routed_mid_is_f16 ? sizeof(uint16_t) : sizeof(float)); const uint64_t sel_bytes = (uint64_t)n_tokens * DS4_N_EXPERT_USED * sizeof(int); - void *mid_dst = g->batch_routed_mid_is_f16 + void *mid_dst = g->scr->batch_routed_mid_is_f16 ? (void *)c->routed_mid_f16_buf : (void *)c->routed_mid_buf; - if (ds4_gpu_tensor_read(g->batch_ffn_norm, 0, c->ffn_norm_buf, norm_bytes) == 0 || - ds4_gpu_tensor_read(g->batch_routed_mid, 0, mid_dst, mid_bytes) == 0 || - ds4_gpu_tensor_read(g->batch_router_selected, 0, c->selected_buf, sel_bytes) == 0) + if (ds4_gpu_tensor_read(g->scr->batch_ffn_norm, 0, c->ffn_norm_buf, norm_bytes) == 0 || + ds4_gpu_tensor_read(g->scr->batch_routed_mid, 0, mid_dst, mid_bytes) == 0 || + ds4_gpu_tensor_read(g->scr->batch_router_selected, 0, c->selected_buf, sel_bytes) == 0) { return false; } @@ -20164,7 +20264,7 @@ static bool imatrix_collect_layer_batch( float *down = imatrix_down_ptr(c, il, (uint32_t)expert); const size_t mid_off = ((size_t)t * DS4_N_EXPERT_USED + slot) * DS4_N_FF_EXP; - if (g->batch_routed_mid_is_f16) { + if (g->scr->batch_routed_mid_is_f16) { const uint16_t *mid = c->routed_mid_f16_buf + mid_off; for (uint32_t i = 0; i < DS4_N_FF_EXP; i++) { const float v = f16_to_f32(mid[i]); @@ -20326,7 +20426,7 @@ static bool metal_graph_prefill_layer_major( if (display_progress) display_progress(display_progress_ud, "prefill_display", (int)start, prompt->len); - bool ok = metal_graph_upload_prompt_tokens(g->prefill_tokens, prompt, start, n_tokens); + bool ok = metal_graph_upload_prompt_tokens(g->scr->prefill_tokens, prompt, start, n_tokens); if (!ok) return false; #ifdef DS4_ROCM_BUILD @@ -20358,8 +20458,8 @@ static bool metal_graph_prefill_layer_major( double execute_s = 0.0; if (!split_commands) { - ok = metal_graph_upload_prompt_embeddings_hc(g->batch_cur_hc, - g->prefill_tokens, + ok = metal_graph_upload_prompt_embeddings_hc(g->scr->batch_cur_hc, + g->scr->prefill_tokens, model, weights, prompt, @@ -20396,22 +20496,22 @@ static bool metal_graph_prefill_layer_major( output_row = (uint32_t)v; } } - ds4_gpu_tensor *saved_cur = g->cur_hc; + ds4_gpu_tensor *saved_cur = g->scr->cur_hc; ds4_gpu_tensor *last_hc = NULL; if (ok && logits) { - last_hc = metal_graph_tensor_row_view(g->batch_cur_hc, output_row, hc_dim); + last_hc = metal_graph_tensor_row_view(g->scr->batch_cur_hc, output_row, hc_dim); ok = last_hc != NULL; } if (ok && logits) { - g->cur_hc = last_hc; + g->scr->cur_hc = last_hc; ok = metal_graph_encode_output_head(g, model, weights, weights->output->dim[1]); - g->cur_hc = saved_cur; + g->scr->cur_hc = saved_cur; } const double t_encoded = profile ? now_sec() : 0.0; if (ok) ok = ds4_gpu_end_commands() != 0; const double t_done = profile ? now_sec() : 0.0; - g->cur_hc = saved_cur; + g->scr->cur_hc = saved_cur; if (last_hc) ds4_gpu_tensor_free(last_hc); if (!ok) { if (ds4_gpu_synchronize() == 0) { @@ -20422,7 +20522,7 @@ static bool metal_graph_prefill_layer_major( const double t_before_read = profile ? now_sec() : 0.0; if (logits) { - ok = ds4_gpu_tensor_read(g->logits, 0, logits, (uint64_t)DS4_N_VOCAB * sizeof(float)) != 0; + ok = ds4_gpu_tensor_read(g->scr->logits, 0, logits, (uint64_t)DS4_N_VOCAB * sizeof(float)) != 0; } if (profile) { const double t_read = now_sec(); @@ -20505,8 +20605,8 @@ static bool metal_graph_prefill_layer_major( #endif double t_layer0 = (profile || throttle) ? now_sec() : 0.0; - ok = metal_graph_upload_prompt_embeddings_hc(g->batch_cur_hc, - g->prefill_tokens, + ok = metal_graph_upload_prompt_embeddings_hc(g->scr->batch_cur_hc, + g->scr->prefill_tokens, model, weights, prompt, @@ -20659,9 +20759,9 @@ static bool metal_graph_prefill_layer_major( fprintf(stderr, "ds4: gpu layer-major prefill layer %u ffn encode failed\n", il); } if (ok) { - ds4_gpu_tensor *tmp = g->batch_cur_hc; - g->batch_cur_hc = g->batch_next_hc; - g->batch_next_hc = tmp; + ds4_gpu_tensor *tmp = g->scr->batch_cur_hc; + g->scr->batch_cur_hc = g->scr->batch_next_hc; + g->scr->batch_next_hc = tmp; } if (ok) ok = metal_graph_capture_prefill_seed_router_selected(g, il, @@ -20828,12 +20928,12 @@ static bool metal_graph_prefill_layer_major( output_row = (uint32_t)v; } } - ds4_gpu_tensor *saved_cur = g->cur_hc; + ds4_gpu_tensor *saved_cur = g->scr->cur_hc; ds4_gpu_tensor *last_hc = NULL; const double t_head0 = profile ? now_sec() : 0.0; if (logits) { - last_hc = metal_graph_tensor_row_view(g->batch_cur_hc, + last_hc = metal_graph_tensor_row_view(g->scr->batch_cur_hc, output_row, hc_dim); ok = last_hc != NULL; @@ -20853,20 +20953,20 @@ static bool metal_graph_prefill_layer_major( } } if (ok && logits) { - g->cur_hc = last_hc; + g->scr->cur_hc = last_hc; ok = ds4_gpu_begin_commands() != 0; } if (ok && logits) ok = metal_graph_encode_output_head(g, model, weights, weights->output->dim[1]); const double t_head_encoded = profile ? now_sec() : 0.0; if (ok && logits) ok = ds4_gpu_end_commands() != 0; const double t_head_done = profile ? now_sec() : 0.0; - g->cur_hc = saved_cur; + g->scr->cur_hc = saved_cur; if (last_hc) ds4_gpu_tensor_free(last_hc); if (!ok) return false; const double t_before_read = profile ? now_sec() : 0.0; if (logits) { - ok = ds4_gpu_tensor_read(g->logits, 0, logits, (uint64_t)DS4_N_VOCAB * sizeof(float)) != 0; + ok = ds4_gpu_tensor_read(g->scr->logits, 0, logits, (uint64_t)DS4_N_VOCAB * sizeof(float)) != 0; } if (profile) { const double t_read = now_sec(); @@ -21124,14 +21224,14 @@ static bool metal_graph_verify_suffix_tops( bool capture_prefix1, int *row_tops, float *row_logits) { - if (n_tokens == 0 || n_tokens > g->prefill_cap || !g->spec_logits) return false; + if (n_tokens == 0 || n_tokens > g->prefill_cap || !g->scr->spec_logits) return false; if (start > (uint32_t)prompt->len || n_tokens > (uint32_t)prompt->len - start) return false; const uint32_t top_rows = n_tokens > 1 ? n_tokens - 1 : 0; if (top_rows && !row_tops) return false; - bool ok = metal_graph_upload_prompt_tokens(g->prefill_tokens, prompt, start, n_tokens); - if (ok) ok = metal_graph_upload_prompt_embeddings_hc(g->batch_cur_hc, - g->prefill_tokens, + bool ok = metal_graph_upload_prompt_tokens(g->scr->prefill_tokens, prompt, start, n_tokens); + if (ok) ok = metal_graph_upload_prompt_embeddings_hc(g->scr->batch_cur_hc, + g->scr->prefill_tokens, model, weights, prompt, @@ -21167,15 +21267,15 @@ static bool metal_graph_verify_suffix_tops( /* Common K=2 verify case: top_k=1 over n_vocab → use the dedicated * argmax kernel (single-block tree-reduce) instead of the legacy * indexer_topk_kernel's single-thread O(n_vocab * top_k) fall-through. */ - ok = ds4_gpu_argmax_tensor(g->comp_selected, - g->spec_logits, + ok = ds4_gpu_argmax_tensor(g->scr->comp_selected, + g->scr->spec_logits, DS4_N_VOCAB) != 0; } else if (top_rows) { /* top-1 of each of the top_rows rows: n_tokens=top_rows, top_k=1. * The order is transposed vs the indexer-score callers; a swap * silently scores row 0's runner-ups instead of each row. */ - ok = ds4_gpu_indexer_topk_tensor(g->comp_selected, - g->spec_logits, + ok = ds4_gpu_indexer_topk_tensor(g->scr->comp_selected, + g->scr->spec_logits, DS4_N_VOCAB, top_rows, 1) != 0; @@ -21184,13 +21284,13 @@ static bool metal_graph_verify_suffix_tops( if (ok) ok = ds4_gpu_end_commands() != 0; else (void)ds4_gpu_synchronize(); if (ok && top_rows) { - ok = ds4_gpu_tensor_read(g->comp_selected, + ok = ds4_gpu_tensor_read(g->scr->comp_selected, 0, row_tops, (uint64_t)top_rows * sizeof(row_tops[0])) != 0; } if (ok && row_logits) { - ok = ds4_gpu_tensor_read(g->spec_logits, + ok = ds4_gpu_tensor_read(g->scr->spec_logits, 0, row_logits, (uint64_t)n_tokens * DS4_N_VOCAB * sizeof(row_logits[0])) != 0; @@ -21199,9 +21299,9 @@ static bool metal_graph_verify_suffix_tops( } static bool metal_graph_read_spec_logits_row(ds4_gpu_graph *g, uint32_t row, float *logits) { - if (!g || !g->spec_logits || !logits || row >= g->prefill_cap) return false; + if (!g || !g->scr->spec_logits || !logits || row >= g->prefill_cap) return false; const uint64_t row_bytes = (uint64_t)DS4_N_VOCAB * sizeof(float); - return ds4_gpu_tensor_read(g->spec_logits, + return ds4_gpu_tensor_read(g->scr->spec_logits, (uint64_t)row * row_bytes, logits, row_bytes) != 0; @@ -21228,10 +21328,10 @@ static bool metal_graph_verify_decode2_exact( if (!g || !top0 || !logits1 || g->raw_cap == 0) return false; const uint64_t hc_dim = (uint64_t)DS4_N_HC * DS4_N_EMBD; - ds4_gpu_tensor *cur0 = metal_graph_tensor_row_view(g->batch_cur_hc, 0, hc_dim); - ds4_gpu_tensor *cur1 = metal_graph_tensor_row_view(g->batch_cur_hc, 1, hc_dim); - ds4_gpu_tensor *next0 = metal_graph_tensor_row_view(g->batch_next_hc, 0, hc_dim); - ds4_gpu_tensor *next1 = metal_graph_tensor_row_view(g->batch_next_hc, 1, hc_dim); + ds4_gpu_tensor *cur0 = metal_graph_tensor_row_view(g->scr->batch_cur_hc, 0, hc_dim); + ds4_gpu_tensor *cur1 = metal_graph_tensor_row_view(g->scr->batch_cur_hc, 1, hc_dim); + ds4_gpu_tensor *next0 = metal_graph_tensor_row_view(g->scr->batch_next_hc, 0, hc_dim); + ds4_gpu_tensor *next1 = metal_graph_tensor_row_view(g->scr->batch_next_hc, 1, hc_dim); bool ok = cur0 && cur1 && next0 && next1; if (ok) ok = ds4_gpu_embed_token_hc_tensor(cur0, @@ -21251,8 +21351,8 @@ static bool metal_graph_verify_decode2_exact( DS4_N_EMBD, DS4_N_HC) != 0; - ds4_gpu_tensor *saved_cur = g->cur_hc; - ds4_gpu_tensor *saved_after = g->after_ffn_hc; + ds4_gpu_tensor *saved_cur = g->scr->cur_hc; + ds4_gpu_tensor *saved_after = g->scr->after_ffn_hc; const bool saved_capture = g->spec_capture_prefix1; g->spec_capture_prefix1 = true; if (ok) ok = ds4_gpu_begin_commands() != 0; @@ -21260,8 +21360,8 @@ static bool metal_graph_verify_decode2_exact( const uint32_t pos0 = start; const uint32_t pos1 = start + 1u; - g->cur_hc = cur0; - g->after_ffn_hc = next0; + g->scr->cur_hc = cur0; + g->scr->after_ffn_hc = next0; ok = metal_graph_encode_decode_layer(g, model, &weights->layer[il], @@ -21277,8 +21377,8 @@ static bool metal_graph_verify_decode2_exact( metal_graph_capture_prefix1_index_state(g, il); if (!ok) break; - g->cur_hc = cur1; - g->after_ffn_hc = next1; + g->scr->cur_hc = cur1; + g->scr->after_ffn_hc = next1; ok = metal_graph_encode_decode_layer(g, model, &weights->layer[il], @@ -21297,22 +21397,22 @@ static bool metal_graph_verify_decode2_exact( if (ok) ok = ds4_gpu_end_commands() != 0; else (void)ds4_gpu_synchronize(); g->spec_capture_prefix1 = saved_capture; - g->cur_hc = saved_cur; - g->after_ffn_hc = saved_after; + g->scr->cur_hc = saved_cur; + g->scr->after_ffn_hc = saved_after; if (ok) { - g->cur_hc = cur0; + g->scr->cur_hc = cur0; ok = ds4_gpu_begin_commands() != 0; if (ok) ok = metal_graph_encode_output_head(g, model, weights, weights->output->dim[1]); - if (ok) ok = ds4_gpu_argmax_tensor(g->comp_selected, - g->logits, + if (ok) ok = ds4_gpu_argmax_tensor(g->scr->comp_selected, + g->scr->logits, DS4_N_VOCAB) != 0; if (ok) ok = ds4_gpu_end_commands() != 0; else (void)ds4_gpu_synchronize(); - g->cur_hc = saved_cur; - if (ok) ok = ds4_gpu_tensor_read(g->comp_selected, 0, top0, sizeof(*top0)) != 0; + g->scr->cur_hc = saved_cur; + if (ok) ok = ds4_gpu_tensor_read(g->scr->comp_selected, 0, top0, sizeof(*top0)) != 0; if (ok && logits0) { - ok = ds4_gpu_tensor_read(g->logits, + ok = ds4_gpu_tensor_read(g->scr->logits, 0, logits0, (uint64_t)DS4_N_VOCAB * sizeof(logits0[0])) != 0; @@ -21320,21 +21420,21 @@ static bool metal_graph_verify_decode2_exact( } if (ok) { - g->cur_hc = cur1; + g->scr->cur_hc = cur1; ok = ds4_gpu_begin_commands() != 0; if (ok) ok = metal_graph_encode_output_head(g, model, weights, weights->output->dim[1]); if (ok) ok = ds4_gpu_end_commands() != 0; else (void)ds4_gpu_synchronize(); - g->cur_hc = saved_cur; + g->scr->cur_hc = saved_cur; if (ok) { - ok = ds4_gpu_tensor_read(g->logits, + ok = ds4_gpu_tensor_read(g->scr->logits, 0, logits1, (uint64_t)DS4_N_VOCAB * sizeof(logits1[0])) != 0; } } - g->cur_hc = saved_cur; - g->after_ffn_hc = saved_after; + g->scr->cur_hc = saved_cur; + g->scr->after_ffn_hc = saved_after; g->spec_capture_prefix1 = saved_capture; ds4_gpu_tensor_free(next1); @@ -21509,10 +21609,13 @@ static int metal_graph_prompt_logits_test( const uint32_t raw_cap = metal_graph_raw_cap_for_context(ctx_size, (uint32_t)n_test); ds4_gpu_graph g; - bool ok = metal_graph_alloc_raw_cap(&g, weights, &weights->layer[0], + ds4_gpu_scratch scr; + memset(&scr, 0, sizeof(scr)); + bool ok = metal_graph_alloc_raw_cap(&g, &scr, weights, &weights->layer[0], raw_cap, (uint32_t)ctx_size, (uint32_t)n_test, false); if (!ok) { metal_graph_free(&g); + metal_scratch_free(&scr); fprintf(stderr, "ds4: failed to initialize Metal graph prompt test runtime\n"); return 1; } @@ -21664,6 +21767,7 @@ static int metal_graph_prompt_logits_test( free(oracle_logits); kv_cache_free(&cpu_cache); metal_graph_free(&g); + metal_scratch_free(&scr); return ok ? 0 : 1; } @@ -21811,6 +21915,12 @@ struct ds4_engine { ds4_vocab vocab; ds4_weights weights; ds4_mtp_weights mtp_weights; +#ifndef DS4_NO_GPU + /* Work buffers shared by every session of this engine (see + * ds4_gpu_scratch). Grown on demand by session creation, freed with the + * engine. */ + ds4_gpu_scratch gpu_scratch; +#endif ds4_backend backend; int mtp_draft_tokens; float mtp_margin; @@ -22955,7 +23065,9 @@ static int generate_metal_graph_raw_swa( prompt->len); } ds4_gpu_graph g; - bool ok = metal_graph_alloc_raw_cap(&g, weights, &weights->layer[0], + ds4_gpu_scratch scr; + memset(&scr, 0, sizeof(scr)); + bool ok = metal_graph_alloc_raw_cap(&g, &scr, weights, &weights->layer[0], raw_cap, (uint32_t)ctx_size, prefill_cap, false); if (!ok) { fprintf(stderr, "ds4: failed to allocate GPU graph runtime\n"); @@ -22971,6 +23083,7 @@ static int generate_metal_graph_raw_swa( directional_steering_attn, directional_steering_ffn)) { metal_graph_free(&g); + metal_scratch_free(&scr); return 1; } const bool memory_report = getenv("DS4_METAL_MEMORY_REPORT") != NULL; @@ -22999,6 +23112,7 @@ static int generate_metal_graph_raw_swa( if (!ok) { free(logits); metal_graph_free(&g); + metal_scratch_free(&scr); return 1; } const char *dump_prefill_logits = getenv("DS4_METAL_DUMP_PREFILL_LOGITS"); @@ -23006,6 +23120,7 @@ static int generate_metal_graph_raw_swa( if (!write_f32_binary_file(dump_prefill_logits, logits, DS4_N_VOCAB)) { free(logits); metal_graph_free(&g); + metal_scratch_free(&scr); return 1; } fprintf(stderr, "ds4: wrote GPU prefill logits to %s\n", dump_prefill_logits); @@ -23062,6 +23177,7 @@ static int generate_metal_graph_raw_swa( if (memory_report) ds4_gpu_print_memory_report("before graph free"); free(logits); metal_graph_free(&g); + metal_scratch_free(&scr); return ok ? 0 : 1; } #endif @@ -25022,7 +25138,9 @@ int ds4_engine_collect_imatrix(ds4_engine *e, const uint32_t raw_cap = metal_graph_raw_cap_for_context(ctx_size, prefill_cap); ds4_gpu_graph g; - bool ok = metal_graph_alloc_raw_cap(&g, weights, &weights->layer[0], + ds4_gpu_scratch scr; + memset(&scr, 0, sizeof(scr)); + bool ok = metal_graph_alloc_raw_cap(&g, &scr, weights, &weights->layer[0], raw_cap, (uint32_t)ctx_size, prefill_cap, false); if (!ok) { fprintf(stderr, "ds4: failed to allocate imatrix Metal graph runtime\n"); @@ -25039,6 +25157,7 @@ int ds4_engine_collect_imatrix(ds4_engine *e, if (!imatrix_collector_init(&collector, prefill_cap, dataset_path)) { fprintf(stderr, "ds4: failed to allocate imatrix collector\n"); metal_graph_free(&g); + metal_scratch_free(&scr); free(dataset); return 1; } @@ -25135,6 +25254,7 @@ int ds4_engine_collect_imatrix(ds4_engine *e, imatrix_collector_free(&collector); metal_graph_free(&g); + metal_scratch_free(&scr); free(dataset); return ok ? 0 : 1; #endif @@ -26027,6 +26147,12 @@ void ds4_engine_close(ds4_engine *e) { if (e->mtp_ready) model_close(&e->mtp_model); model_close(&e->model); #ifndef DS4_NO_GPU + if (e->gpu_scratch.refcount > 0) { + fprintf(stderr, + "ds4: warning: closing engine with %d live session(s)\n", + e->gpu_scratch.refcount); + } + metal_scratch_free(&e->gpu_scratch); ds4_gpu_cleanup(); #endif ds4_ssd_memory_lock_release(&e->simulated_memory); @@ -26071,7 +26197,7 @@ int ds4_session_create(ds4_session **out, ds4_engine *e, int ctx_size) { free(s); return 1; } - if (!metal_graph_alloc_raw_cap(&s->graph, &e->weights, shape_layer, + if (!metal_graph_alloc_raw_cap(&s->graph, &e->gpu_scratch, &e->weights, shape_layer, raw_cap, (uint32_t)ctx_size, s->prefill_cap, e->mtp_ready)) { free(s); @@ -26247,7 +26373,7 @@ int ds4_session_eval_output_head_from_hc(ds4_session *s, return 1; #else ds4_gpu_graph *g = &s->graph; - bool ok = ds4_gpu_tensor_write(g->cur_hc, + bool ok = ds4_gpu_tensor_write(g->scr->cur_hc, 0, last_hc, hc_dim * sizeof(float)) != 0; @@ -26257,7 +26383,7 @@ int ds4_session_eval_output_head_from_hc(ds4_session *s, &e->weights, e->weights.output->dim[1]); if (ok) ok = ds4_gpu_end_commands() != 0; - if (ok) ok = ds4_gpu_tensor_read(g->logits, + if (ok) ok = ds4_gpu_tensor_read(g->scr->logits, 0, logits, (uint64_t)DS4_N_VOCAB * sizeof(float)) != 0; @@ -26461,11 +26587,11 @@ int ds4_session_eval_layer_slice(ds4_session *s, ok = metal_graph_stream_map_token(&e->model, &e->weights); } if (input_hc) { - ok = ds4_gpu_tensor_write(g->cur_hc, 0, input_hc, hc_dim * sizeof(float)) != 0; + ok = ds4_gpu_tensor_write(g->scr->cur_hc, 0, input_hc, hc_dim * sizeof(float)) != 0; } if (ok) ok = ds4_gpu_begin_commands() != 0; if (ok && !input_hc) { - ok = ds4_gpu_embed_token_hc_tensor(g->cur_hc, + ok = ds4_gpu_embed_token_hc_tensor(g->scr->cur_hc, e->model.map, e->model.size, e->weights.token_embd->abs_offset, @@ -26495,9 +26621,9 @@ int ds4_session_eval_layer_slice(ds4_session *s, raw_row, n_raw, tokens[0]); - ds4_gpu_tensor *tmp = g->cur_hc; - g->cur_hc = g->after_ffn_hc; - g->after_ffn_hc = tmp; + ds4_gpu_tensor *tmp = g->scr->cur_hc; + g->scr->cur_hc = g->scr->after_ffn_hc; + g->scr->after_ffn_hc = tmp; } if (ok) ok = ds4_gpu_end_commands() != 0; } @@ -26520,9 +26646,9 @@ int ds4_session_eval_layer_slice(ds4_session *s, raw_row, n_raw, tokens[0]); - ds4_gpu_tensor *tmp = g->cur_hc; - g->cur_hc = g->after_ffn_hc; - g->after_ffn_hc = tmp; + ds4_gpu_tensor *tmp = g->scr->cur_hc; + g->scr->cur_hc = g->scr->after_ffn_hc; + g->scr->after_ffn_hc = tmp; encoded_layers++; if (ok && split_after_layers != 0 && @@ -26539,10 +26665,10 @@ int ds4_session_eval_layer_slice(ds4_session *s, } if (ok && !output_hc && !output_logits) ok = ds4_gpu_synchronize() != 0; if (ok && output_hc) { - ok = ds4_gpu_tensor_read(g->cur_hc, 0, output_hc, hc_dim * sizeof(float)) != 0; + ok = ds4_gpu_tensor_read(g->scr->cur_hc, 0, output_hc, hc_dim * sizeof(float)) != 0; } if (ok && output_logits) { - ok = ds4_gpu_tensor_read(g->logits, 0, logits, (uint64_t)DS4_N_VOCAB * sizeof(float)) != 0; + ok = ds4_gpu_tensor_read(g->scr->logits, 0, logits, (uint64_t)DS4_N_VOCAB * sizeof(float)) != 0; } if (!ok) { if (ds4_gpu_synchronize() == 0) { @@ -26569,12 +26695,12 @@ int ds4_session_eval_layer_slice(ds4_session *s, g->streaming_static_decode_map_current = false; ok = metal_graph_stream_map_token(&e->model, &e->weights); } - if (ok) ok = metal_graph_upload_prompt_tokens(g->prefill_tokens, &span, 0, n_tokens); + if (ok) ok = metal_graph_upload_prompt_tokens(g->scr->prefill_tokens, &span, 0, n_tokens); if (ok && input_hc) { - ok = ds4_gpu_tensor_write(g->batch_cur_hc, 0, input_hc, hc_bytes) != 0; + ok = ds4_gpu_tensor_write(g->scr->batch_cur_hc, 0, input_hc, hc_bytes) != 0; } else if (ok) { - ok = metal_graph_upload_prompt_embeddings_hc(g->batch_cur_hc, - g->prefill_tokens, + ok = metal_graph_upload_prompt_embeddings_hc(g->scr->batch_cur_hc, + g->scr->prefill_tokens, &e->model, &e->weights, &span, @@ -26618,31 +26744,31 @@ int ds4_session_eval_layer_slice(ds4_session *s, } } if (ok && output_logits) { - saved_cur = g->cur_hc; - last_hc = metal_graph_tensor_row_view(g->batch_cur_hc, n_tokens - 1u, hc_dim); + saved_cur = g->scr->cur_hc; + last_hc = metal_graph_tensor_row_view(g->scr->batch_cur_hc, n_tokens - 1u, hc_dim); ok = last_hc != NULL; if (ok && g->ssd_streaming) { g->streaming_static_decode_map_current = false; ok = metal_graph_stream_map_output(&e->model, &e->weights); } if (ok) { - g->cur_hc = last_hc; + g->scr->cur_hc = last_hc; if (g->ssd_streaming) ok = ds4_gpu_begin_commands() != 0; if (ok) ok = metal_graph_encode_output_head(g, &e->model, &e->weights, e->weights.output->dim[1]); if (ok && g->ssd_streaming) ok = ds4_gpu_end_commands() != 0; - g->cur_hc = saved_cur; + g->scr->cur_hc = saved_cur; } } if (ok && !g->ssd_streaming) ok = ds4_gpu_end_commands() != 0; - if (saved_cur) g->cur_hc = saved_cur; + if (saved_cur) g->scr->cur_hc = saved_cur; if (last_hc) ds4_gpu_tensor_free(last_hc); if (ok && !output_hc && !output_logits) ok = ds4_gpu_synchronize() != 0; if (ok && output_hc) { - ok = ds4_gpu_tensor_read(g->batch_cur_hc, 0, output_hc, hc_bytes) != 0; + ok = ds4_gpu_tensor_read(g->scr->batch_cur_hc, 0, output_hc, hc_bytes) != 0; } if (ok && output_logits) { - ok = ds4_gpu_tensor_read(g->logits, 0, logits, (uint64_t)DS4_N_VOCAB * sizeof(float)) != 0; + ok = ds4_gpu_tensor_read(g->scr->logits, 0, logits, (uint64_t)DS4_N_VOCAB * sizeof(float)) != 0; } if (!ok) { if (ds4_gpu_synchronize() == 0) { @@ -27262,8 +27388,8 @@ int ds4_session_eval_speculative_argmax(ds4_session *s, int first_token, } while (0) for (; draft_n < draft_cap; draft_n++) { - ds4_gpu_tensor *prev_hc = (draft_n & 1) ? s->graph.mtp_state_hc : s->graph.mtp_next_hc; - ds4_gpu_tensor *out_hc = (draft_n & 1) ? s->graph.mtp_next_hc : s->graph.mtp_state_hc; + ds4_gpu_tensor *prev_hc = (draft_n & 1) ? s->graph.scr->mtp_state_hc : s->graph.scr->mtp_next_hc; + ds4_gpu_tensor *out_hc = (draft_n & 1) ? s->graph.scr->mtp_next_hc : s->graph.scr->mtp_state_hc; int mtp_top = -1; if (!metal_graph_eval_mtp_draft_from_hc(&s->graph, &e->model, @@ -27724,7 +27850,7 @@ int ds4_session_eval_speculative_argmax(ds4_session *s, int first_token, if (drafts[i] == eos_token) break; } if (verified > 0 && !logits_on_host) { - if (ds4_gpu_tensor_read(s->graph.logits, + if (ds4_gpu_tensor_read(s->graph.scr->logits, 0, s->logits, (uint64_t)DS4_N_VOCAB * sizeof(s->logits[0])) == 0) diff --git a/ds4.h b/ds4.h index 9d040c92b..012ac3b21 100644 --- a/ds4.h +++ b/ds4.h @@ -217,6 +217,12 @@ int ds4_token_eos(ds4_engine *e); int ds4_token_user(ds4_engine *e); int ds4_token_assistant(ds4_engine *e); +/* Multiple sessions may share one engine: each session owns its live KV + * cache and timeline, while transient work buffers are engine-owned and + * shared. Sessions sharing an engine must therefore be driven from one + * thread at a time (create/sync/eval/sample/free); interleaving calls across + * sessions on that thread is fine, concurrent calls are not. This matches + * the server's single graph worker. */ int ds4_session_create(ds4_session **out, ds4_engine *e, int ctx_size); void ds4_session_free(ds4_session *s); int ds4_session_power(ds4_session *s); diff --git a/tests/two_sessions_smoke.c b/tests/two_sessions_smoke.c new file mode 100644 index 000000000..bb24346e3 --- /dev/null +++ b/tests/two_sessions_smoke.c @@ -0,0 +1,144 @@ +/* two_sessions_smoke.c — prove that two ds4_session objects sharing one + * ds4_engine produce the same greedy token streams when driven interleaved + * (one token each, alternating on the same thread) as when each session is + * driven alone. This is the ground assumption for multi-session serving: + * backend statics must not carry per-graph state across ds4_session_eval() + * calls of different sessions. + * + * Usage: tests/two_sessions_smoke [n_tokens] + */ +#include +#include +#include + +#include "../ds4.h" + +#define MAX_GEN 128 + +typedef struct { + int toks[MAX_GEN]; + int len; +} gen_result; + +static void die(const char *msg, const char *detail) { + fprintf(stderr, "two_sessions_smoke: %s%s%s\n", msg, + detail ? ": " : "", detail ? detail : ""); + exit(1); +} + +static void encode_prompt(ds4_engine *e, const char *system, const char *user, + ds4_tokens *out) { + memset(out, 0, sizeof(*out)); + ds4_encode_chat_prompt(e, system, user, DS4_THINK_NONE, out); +} + +/* One greedy step: argmax over current logits, then advance the session. + * Returns the sampled token, or -1 once the stream hit EOS. */ +static int greedy_step(ds4_session *s, int eos, char *err, size_t errlen) { + int tok = ds4_session_argmax(s); + if (tok < 0 || tok == eos) return -1; + if (ds4_session_eval(s, tok, err, errlen) != 0) die("eval failed", err); + return tok; +} + +static void run_isolated(ds4_engine *e, const ds4_tokens *prompt, int n, + int ctx_size, gen_result *out) { + ds4_session *s = NULL; + char err[256] = ""; + if (ds4_session_create(&s, e, ctx_size) != 0) die("session_create failed", NULL); + if (ds4_session_sync(s, prompt, err, sizeof(err)) != 0) die("sync failed", err); + int eos = ds4_token_eos(e); + out->len = 0; + for (int i = 0; i < n; i++) { + int tok = greedy_step(s, eos, err, sizeof(err)); + if (tok < 0) break; + out->toks[out->len++] = tok; + } + ds4_session_free(s); +} + +static void run_interleaved(ds4_engine *e, const ds4_tokens *pa, + const ds4_tokens *pb, int n, int ctx_size, + gen_result *ra, gen_result *rb) { + ds4_session *sa = NULL, *sb = NULL; + char err[256] = ""; + if (ds4_session_create(&sa, e, ctx_size) != 0) die("session_create A failed", NULL); + if (ds4_session_create(&sb, e, ctx_size) != 0) die("session_create B failed", NULL); + if (ds4_session_sync(sa, pa, err, sizeof(err)) != 0) die("sync A failed", err); + if (ds4_session_sync(sb, pb, err, sizeof(err)) != 0) die("sync B failed", err); + int eos = ds4_token_eos(e); + ra->len = rb->len = 0; + bool a_done = false, b_done = false; + for (int i = 0; i < n && (!a_done || !b_done); i++) { + if (!a_done) { + int tok = greedy_step(sa, eos, err, sizeof(err)); + if (tok < 0) a_done = true; else ra->toks[ra->len++] = tok; + } + if (!b_done) { + int tok = greedy_step(sb, eos, err, sizeof(err)); + if (tok < 0) b_done = true; else rb->toks[rb->len++] = tok; + } + } + ds4_session_free(sa); + ds4_session_free(sb); +} + +static void print_stream(ds4_engine *e, const char *label, const gen_result *r) { + printf("%s (%d tokens): ", label, r->len); + for (int i = 0; i < r->len; i++) { + size_t len = 0; + char *txt = ds4_token_text(e, r->toks[i], &len); + if (txt) fwrite(txt, 1, len, stdout); + } + printf("\n"); +} + +static bool same(const gen_result *x, const gen_result *y) { + if (x->len != y->len) return false; + return memcmp(x->toks, y->toks, sizeof(int) * (size_t)x->len) == 0; +} + +int main(int argc, char **argv) { + if (argc < 2) die("usage: two_sessions_smoke [n_tokens]", NULL); + int n = argc > 2 ? atoi(argv[2]) : 24; + if (n <= 0 || n > MAX_GEN) n = 24; + int ctx_size = 4096; + + ds4_engine_options opt; + memset(&opt, 0, sizeof(opt)); + opt.model_path = argv[1]; + opt.backend = DS4_BACKEND_CUDA; + opt.prefill_chunk = 256; + + ds4_engine *e = NULL; + if (ds4_engine_open(&e, &opt) != 0) die("engine_open failed", argv[1]); + + const char *sys = "You are a concise technical assistant."; + ds4_tokens pa, pb; + encode_prompt(e, sys, "Briefly explain what a compressor does in a refrigeration cycle.", &pa); + encode_prompt(e, sys, "List three common causes of high condensing pressure.", &pb); + + gen_result iso_a, iso_b, mix_a, mix_b; + printf("== isolated runs ==\n"); + run_isolated(e, &pa, n, ctx_size, &iso_a); + print_stream(e, "A/isolated", &iso_a); + run_isolated(e, &pb, n, ctx_size, &iso_b); + print_stream(e, "B/isolated", &iso_b); + + printf("== interleaved run ==\n"); + run_interleaved(e, &pa, &pb, n, ctx_size, &mix_a, &mix_b); + print_stream(e, "A/interleaved", &mix_a); + print_stream(e, "B/interleaved", &mix_b); + + bool ok_a = same(&iso_a, &mix_a); + bool ok_b = same(&iso_b, &mix_b); + printf("A: %s, B: %s\n", ok_a ? "MATCH" : "MISMATCH", ok_b ? "MATCH" : "MISMATCH"); + + ds4_tokens_free(&pa); + ds4_tokens_free(&pb); + ds4_engine_close(e); + + if (!ok_a || !ok_b) { printf("FAILED\n"); return 1; } + printf("OK\n"); + return 0; +}