diff --git a/common/arg.cpp b/common/arg.cpp index 63f3a7b06adc..94dead3ff3ba 100644 --- a/common/arg.cpp +++ b/common/arg.cpp @@ -4118,6 +4118,16 @@ common_params_context common_params_parser_init(common_params & params, llama_ex params.speculative.draft.n_min = value; } ).set_spec().set_examples({LLAMA_EXAMPLE_SPECULATIVE, LLAMA_EXAMPLE_LOOKUP, LLAMA_EXAMPLE_SERVER, LLAMA_EXAMPLE_CLI}).set_env("LLAMA_ARG_SPEC_DRAFT_N_MIN")); + add_opt(common_arg( + {"--spec-draft-depth-max"}, "N", + string_format("stop drafting once the sequence is longer than N tokens, 0 = never (default: %d)", params.speculative.draft.n_depth_max), + [](common_params & params, int value) { + if (value < 0) { + throw std::invalid_argument("invalid value"); + } + params.speculative.draft.n_depth_max = value; + } + ).set_spec().set_examples({LLAMA_EXAMPLE_SERVER}).set_env("LLAMA_ARG_SPEC_DRAFT_DEPTH_MAX")); add_opt(common_arg( {"--spec-draft-p-split", "--draft-p-split"}, "P", diff --git a/common/common.h b/common/common.h index eb50a6cbc520..0ab2a3ffb3d1 100644 --- a/common/common.h +++ b/common/common.h @@ -329,6 +329,12 @@ struct common_params_speculative_draft { float p_split = 0.1f; // speculative decoding split probability float p_min = 0.0f; // minimum speculative decoding probability (greedy) + // stop drafting once the sequence is this long (0 = never). Deep in the context a step is bound + // by reading the KV cache, and the draft passes plus the multi-column verify add to that read + // without shortening it: on a 4070 with Bonsai 2 27B the draft is +85% at zero depth, breaks + // even near 24k tokens and costs 30% at 64k. Past the cutoff the slot decodes one token per step. + int32_t n_depth_max = 0; + bool backend_sampling = true; // offload draft sampling to the backend (default: on) common_params_model mparams; diff --git a/common/sampling.cpp b/common/sampling.cpp index 06dea1e1ccea..dd5b14823c76 100644 --- a/common/sampling.cpp +++ b/common/sampling.cpp @@ -9,6 +9,7 @@ #include #include +#include #include #include #include @@ -614,6 +615,21 @@ llama_token common_sampler_sample(struct common_sampler * gsmpl, struct llama_co if (id != LLAMA_TOKEN_NULL) { LOG_DBG("%s: Backend sampler selected token: '%d'. Will not run any CPU samplers\n", __func__, id); + { + const int32_t n_vocab = llama_vocab_n_tokens(llama_model_get_vocab(llama_get_model(ctx))); + if (id < 0 || id >= n_vocab) { + // diagnostic: a backend-sampled id outside the vocab would otherwise surface much later as a + // std::vector::at() throw in the tokenizer ("invalid vector subscript"). log the origin and + // fall back to the CPU chain on the (already fetched) logits/candidates. + LOG_ERR("%s: backend sampler returned out-of-vocab token id %d (idx=%d, n_vocab=%d, cur_p.size=%zu, cur_p[0].id=%d) - falling back to CPU sampling\n", + __func__, id, idx, n_vocab, cur_p.size, cur_p.size > 0 ? cur_p.data[0].id : -1); + id = LLAMA_TOKEN_NULL; + } + } + } + + if (id != LLAMA_TOKEN_NULL) { + GGML_ASSERT(!gsmpl->grmr && "using grammar in combination with backend sampling is not supported"); GGML_ASSERT(!gsmpl->rbudget && "using reasoning budget in combination with backend sampling is not supported"); @@ -639,6 +655,14 @@ llama_token common_sampler_sample(struct common_sampler * gsmpl, struct llama_co id = cur_p.data[cur_p.selected].id; + { + const int32_t n_vocab = llama_vocab_n_tokens(llama_model_get_vocab(llama_get_model(ctx))); + if (id < 0 || id >= n_vocab) { + LOG_ERR("%s: CPU chain selected out-of-vocab token id %d (idx=%d, n_vocab=%d, selected=%" PRId64 ", cur_p.size=%zu) - candidate ids from the backend are corrupt\n", + __func__, id, idx, n_vocab, cur_p.selected, cur_p.size); + } + } + if (grammar_first || !grammar_should_apply(gsmpl)) { return id; } diff --git a/common/speculative.cpp b/common/speculative.cpp index 9b28b3188d65..7e3e8f5bee53 100644 --- a/common/speculative.cpp +++ b/common/speculative.cpp @@ -2069,6 +2069,53 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl { std::vector i_last; std::vector> chain_h; + // Deferred catch-up (single head, separate draft memory). process() used to run the draft + // over every row of the target's verify batch right away, only to fill the draft's memory: + // one extra graph launch per step, spent partly on rows the target then rejects. Instead the + // rows are kept here and decoded at the start of draft() in the same llama_decode as the + // first draft row (positions are contiguous, the draft row attends to them in-batch). + // accept() trims n_valid to the accepted prefix; without an accept() every row is real. + struct deferred_rows { + std::vector tokens; + std::vector pos; + std::vector h; // n_rows * n_embd, shifted target rows as process() builds them + int32_t n_valid = 0; + bool pending = false; + bool catchup_failed = false; // flush_deferred decode failed; do not draft this seq + }; + std::vector deferred; + + bool can_defer() const { + return !is_mem_shared && !chain_heads && !getenv("LLAMA_MTP_EAGER_CATCHUP"); + } + + // decode the deferred rows on their own (fallback for prompts split into tiny ubatches, or a + // sequence that stops drafting). Returns false on decode failure. + bool flush_deferred(llama_seq_id seq_id) { + auto & d = deferred[seq_id]; + if (!d.pending) { + return true; + } + if (d.n_valid <= 0) { + d.pending = false; + return true; + } + const size_t row_bytes = (size_t) n_embd * sizeof(float); + common_batch_clear(batch); + for (int32_t k = 0; k < d.n_valid; ++k) { + common_batch_add(batch, d.tokens[k], d.pos[k], { seq_id }, false); + std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd, d.h.data() + (size_t) k * n_embd, row_bytes); + } + const int32_t rc = llama_decode(params.ctx_dft, batch); + if (rc != 0) { + SPC_ERR("llama_decode(ctx_dft) deferred catch-up failed rc=%d (pos=%d)\n", (int) rc, (int) d.pos[0]); + d.catchup_failed = true; + return false; + } + d.pending = false; + return true; + } + common_speculative_impl_draft_mtp(const common_params_speculative & params, uint32_t n_seq) : common_speculative_impl(COMMON_SPECULATIVE_TYPE_DRAFT_MTP, n_seq) , params(params.draft) @@ -2146,6 +2193,8 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl { verify_h.assign(n_seq, {}); verify_h_rows.assign(n_seq, 0); + + deferred.assign(n_seq, {}); } ~common_speculative_impl_draft_mtp() override { @@ -2175,6 +2224,32 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl { } auto * ctx_dft = this->params.ctx_dft; + if (seq_id >= 0 && seq_id < (llama_seq_id) n_seq) { + // A failed catch-up used to stick for the life of the slot. This task is a new prompt. + deferred[seq_id].catchup_failed = false; + } + if (seq_id >= 0 && seq_id < (llama_seq_id) n_seq && deferred[seq_id].pending) { + // rows left over from the previous task on this slot. They are only worth decoding when they + // are the tail of a prefix the new prompt reuses: the same tokens at the same positions, and + // positions that directly continue what ctx_dft holds. Otherwise (client cancelled a task, + // a different conversation landed on the slot) the server has already trimmed ctx_dft past + // them, and decoding them would put stale positions into the recurrent draft state. + auto & d = deferred[seq_id]; + const llama_pos pos_max_dft = llama_memory_seq_pos_max(llama_get_memory(ctx_dft), seq_id); + bool reuse = d.n_valid > 0 && d.pos[0] == pos_max_dft + 1; + for (int32_t k = 0; reuse && k < d.n_valid; ++k) { + reuse = d.pos[k] < N && prompt[d.pos[k]] == d.tokens[k]; + } + if (reuse) { + if (!flush_deferred(seq_id)) { + // ctx_dft is not caught up; begin() cannot draft this prompt + return; + } + } else { + d.pending = false; + d.n_valid = 0; + } + } const llama_pos pos_max = llama_memory_seq_pos_max(llama_get_memory(ctx_dft), seq_id); if (pos_max < N - 1 && !is_mem_shared) { @@ -2243,14 +2318,73 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl { std::memcpy(batch.embd + (size_t) idx * n_embd, h_row, row_bytes); }; + int n_seq_in_batch = 0; for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) { if (i_batch_beg[seq_id] < 0) { continue; } + n_seq_in_batch++; set_h(i_batch_beg[seq_id], pending_h[seq_id].data()); } + // verify-sized batch of one sequence: hold the rows for draft() instead of decoding now. + // A stash that is still pending here belongs to an earlier batch: decode it first if this + // batch continues it (prompt fed in small ubatches), drop it if the positions overlap + // (the server rolled the sequence back and is replaying). + if (can_defer() && n_seq_in_batch == 1 && n_tokens <= this->params.n_max + 1) { + const llama_seq_id seq_id = batch_in.seq_id[0][0]; + auto & d = deferred[seq_id]; + if (d.pending) { + const bool continues = !d.pos.empty() && batch_in.pos[0] == d.pos[d.n_valid - 1] + 1; + if (continues) { + // flush_deferred() reuses `batch`, so the rows built above are rebuilt below + std::vector tok_save(batch.token, batch.token + n_tokens); + std::vector pos_save(batch.pos, batch.pos + n_tokens); + std::vector h_save(batch.embd, batch.embd + (size_t) n_tokens * n_embd); + if (!flush_deferred(seq_id)) { + return false; + } + d.tokens = std::move(tok_save); + d.pos = std::move(pos_save); + d.h = std::move(h_save); + } else { + d.pending = false; + d.tokens.assign(batch.token, batch.token + n_tokens); + d.pos.assign(batch.pos, batch.pos + n_tokens); + d.h.assign(batch.embd, batch.embd + (size_t) n_tokens * n_embd); + } + } else { + d.tokens.assign(batch.token, batch.token + n_tokens); + d.pos.assign(batch.pos, batch.pos + n_tokens); + d.h.assign(batch.embd, batch.embd + (size_t) n_tokens * n_embd); + } + d.n_valid = n_tokens; + d.pending = true; + } else { + for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) { + if (i_batch_beg[seq_id] >= 0 && deferred[seq_id].pending) { + // a large batch after a pending stash: prompt continuation (flush) or rollback replay (drop) + auto & d = deferred[seq_id]; + const bool continues = !d.pos.empty() && batch_in.pos[i_batch_beg[seq_id]] == d.pos[d.n_valid - 1] + 1; + if (continues) { + std::vector tok_save(batch.token, batch.token + n_tokens); + std::vector pos_save(batch.pos, batch.pos + n_tokens); + std::vector h_save(batch.embd, batch.embd + (size_t) n_tokens * n_embd); + if (!flush_deferred(seq_id)) { + return false; + } + common_batch_clear(batch); + for (int k = 0; k < n_tokens; ++k) { + common_batch_add(batch, tok_save[k], pos_save[k], { batch_in.seq_id[k][0] }, 0); + } + std::memcpy(batch.embd, h_save.data(), row_bytes * n_tokens); + } else { + d.pending = false; + } + } + } + auto * mem_dft = llama_get_memory(ctx_dft); bool ok = true; @@ -2281,6 +2415,7 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl { if (!ok) { return false; } + } // eager catch-up } for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) { @@ -2307,13 +2442,66 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl { void draft(common_speculative_draft_params_vec & dparams) override { auto & ctx_dft = params.ctx_dft; + const size_t row_bytes = (size_t) n_embd * sizeof(float); + + // deferred catch-up rows that do not lead straight into this draft position are decoded + // on their own first (uses `batch`, so before it is built) + for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) { + auto & d = deferred[seq_id]; + auto & dp = dparams[seq_id]; + if (d.catchup_failed) { + dp.drafting = false; + continue; + } + if (!d.pending) { + continue; + } + const bool leads_in = dp.drafting && d.n_valid > 0 && d.pos[d.n_valid - 1] + 1 == dp.n_past; + if (!leads_in) { + if (!flush_deferred(seq_id)) { + dp.drafting = false; + } + } + } + + // Catch-up that leads into this draft rides in the same decode as the anchors. batch is + // allocated at exactly llama_n_batch(ctx_dft); a full stash (n_tokens == n_max+1 == n_batch) + // plus the anchor is a one-row heap overflow, and llama_decode enforces the same limit. + // Flush first when the combined count would not fit. Growing only the host allocation + // would still abort in decode. + { + const int32_t n_b = (int32_t) llama_n_batch(ctx_dft); + int32_t catchup_rows = 0; + int32_t n_anchors = 0; + for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) { + auto & dp = dparams[seq_id]; + auto & d = deferred[seq_id]; + if (d.catchup_failed || !dp.drafting) { + continue; + } + if (d.pending) { + catchup_rows += d.n_valid; + } + n_anchors++; + } + if (!common_speculative_mtp_first_decode_fits(n_b, catchup_rows, n_anchors)) { + for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) { + if (!dparams[seq_id].drafting || !deferred[seq_id].pending) { + continue; + } + if (!flush_deferred(seq_id)) { + dparams[seq_id].drafting = false; + } + } + } + } + common_batch_clear(batch); // keep track of which sequences are still drafting int n_drafting = 0; std::vector drafting(n_seq); - - const size_t row_bytes = (size_t) n_embd * sizeof(float); + std::vector catchup_in_batch(n_seq, 0); for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) { auto & dp = dparams[seq_id]; @@ -2326,6 +2514,18 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl { drafting[seq_id] = true; common_sampler_reset(smpls[seq_id].get()); + // catch-up rows ride in the same decode as the first draft row (no logits) + auto & d = deferred[seq_id]; + if (d.pending) { + for (int32_t k = 0; k < d.n_valid; ++k) { + common_batch_add(batch, d.tokens[k], d.pos[k], { seq_id }, false); + std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd, d.h.data() + (size_t) k * n_embd, row_bytes); + } + // Leave pending set until this decode succeeds. Clearing it here dropped the + // rows on a failed llama_decode: they were in neither the stash nor ctx_dft. + catchup_in_batch[seq_id] = 1; + } + common_batch_add(batch, dp.id_last, dp.n_past, { seq_id }, true); std::memcpy(batch.embd + (size_t) (batch.n_tokens - 1) * n_embd, pending_h[seq_id].data(), row_bytes); @@ -2358,8 +2558,22 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl { int ret = llama_decode(ctx_dft, batch); if (ret != 0) { SPC_ERR("llama_decode[%d] returned %d\n", i, ret); + if (i == 0) { + for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) { + if (catchup_in_batch[seq_id]) { + deferred[seq_id].catchup_failed = true; + } + } + } break; } + if (i == 0) { + for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) n_seq; ++seq_id) { + if (catchup_in_batch[seq_id]) { + deferred[seq_id].pending = false; + } + } + } // rebuild the batch for the next step: the growing-KV paths re-add only the // new token (the KV already holds the prefix), while chained heads re-add the @@ -2387,6 +2601,14 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl { // add drafted token for each sequence const llama_token id = cur_p->data[0].id; + if (id < 0 || id >= llama_vocab_n_tokens(llama_model_get_vocab(llama_get_model(ctx_dft)))) { + SPC_ERR("draft candidate id %d out of vocab (seq_id=%d, step=%d, i_last=%d, cur_p.size=%zu) - dropping draft\n", + id, (int) seq_id, i, i_last[seq_id], (size_t) cur_p->size); + drafting[seq_id] = false; + n_drafting--; + continue; + } + // only collect very high-confidence draft tokens if (cur_p->data[0].p < params.p_min) { drafting[seq_id] = false; @@ -2468,6 +2690,12 @@ struct common_speculative_impl_draft_mtp : public common_speculative_impl { const int32_t i_h = std::min(n_accepted, n_rows - 1); const size_t row_bytes = (size_t) n_embd * sizeof(float); std::memcpy(pending_h[seq_id].data(), verify_h[seq_id].data() + (size_t) i_h * n_embd, row_bytes); + + // only the accepted prefix of the verify batch needs to enter the draft's memory + auto & d = deferred[seq_id]; + if (d.pending) { + d.n_valid = std::min(d.n_valid, i_h + 1); + } } }; @@ -3217,6 +3445,13 @@ common_speculative_init_result_ptr common_speculative_init_from_params(common_pa return std::make_unique(params, model_tgt, ctx_tgt); } +bool common_speculative_mtp_first_decode_fits(int32_t n_batch, int32_t catchup_rows, int32_t n_anchors) { + if (n_batch <= 0 || catchup_rows < 0 || n_anchors < 0) { + return false; + } + return (int64_t) catchup_rows + (int64_t) n_anchors <= (int64_t) n_batch; +} + common_speculative_output_limits common_speculative_get_output_limits( int32_t n_batch, int32_t n_parallel, int32_t n_draft) { const int64_t per_seq = 1 + (int64_t) std::max(0, n_draft); diff --git a/common/speculative.h b/common/speculative.h index d1a724d519bb..0193af142b5c 100644 --- a/common/speculative.h +++ b/common/speculative.h @@ -37,6 +37,10 @@ struct common_speculative_output_limits { common_speculative_output_limits common_speculative_get_output_limits( int32_t n_batch, int32_t n_parallel, int32_t n_draft); +// True if deferred catch-up rows plus one first-draft anchor per sequence fit in one llama_decode. +// Used by draft-mtp so a full-prefill stash (n_tokens == n_batch) does not add a 33rd row. +bool common_speculative_mtp_first_decode_fits(int32_t n_batch, int32_t catchup_rows, int32_t n_anchors); + common_speculative * common_speculative_init(common_params_speculative & params, uint32_t n_seq); void common_speculative_free(common_speculative * spec); diff --git a/ggml/src/ggml-cuda/common.cuh b/ggml/src/ggml-cuda/common.cuh index ef929d3d7842..7e133e3d6398 100644 --- a/ggml/src/ggml-cuda/common.cuh +++ b/ggml/src/ggml-cuda/common.cuh @@ -173,6 +173,14 @@ static int ggml_cuda_highest_compiled_arch(const int arch) { // --------------------------------------------------------------------------------------------------------- +// GGML_CUDA_BATCH_INVARIANT=1: prefer kernels whose per-column arithmetic does not depend on the +// number of columns in the batch (1 to 8), so that a token decoded alone and a token verified +// inside a speculative batch see the same logits bit for bit. Costs some throughput at 2 to 8 columns. +static inline bool ggml_cuda_batch_invariant() { + static const bool enabled = getenv("GGML_CUDA_BATCH_INVARIANT") != nullptr; + return enabled; +} + #define MATRIX_ROW_PADDING 512 // last row of quant. matrices is a multiple of this to avoid out-of-bounds memory accesses #define GGML_CUDA_MAX_STREAMS 8 @@ -741,6 +749,20 @@ static __device__ __forceinline__ int ggml_cuda_dp4a(const int a, const int b, i #endif // defined(GGML_USE_HIP) } +// c += dot(a as 4 unsigned bytes, b as 4 signed bytes). Used by the ternary paths that keep the raw +// digits {0,1,2} and subtract the exact integer activation sum once per block instead of biasing +// every word (two SIMD ops per 4 weights). PTX dp4a takes mixed .u32.s32 operand types directly. +static __device__ __forceinline__ int ggml_cuda_dp4a_us(const unsigned int a, const int b, int c) { +#if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) && __CUDA_ARCH__ >= GGML_CUDA_CC_DP4A + asm("dp4a.u32.s32 %0, %1, %2, %3;" : "=r"(c) : "r"(a), "r"(b), "r"(c)); + return c; +#else + const uint8_t * a8 = (const uint8_t *) &a; + const int8_t * b8 = (const int8_t *) &b; + return c + (int) a8[0]*b8[0] + (int) a8[1]*b8[1] + (int) a8[2]*b8[2] + (int) a8[3]*b8[3]; +#endif +} + static __device__ __forceinline__ void ggml_cuda_mad(float & acc, const float v, const float u) { acc += v*u; } @@ -1026,6 +1048,58 @@ struct ggml_cuda_type_traits { static constexpr int qi = QI_PTQ1_0; }; +// Activation (src1) q8_1 layouts produced by quantize_row_q8_1_cuda and consumed by the MMVQ kernels. +// Every layout keeps block_q8_1's bytes per row, so launcher strides in block_q8_1 units are valid +// for all three; only the byte order inside a column differs. +enum ggml_cuda_q8_1_layout : int { + GGML_CUDA_Q8_1_AOS = 0, // plain block_q8_1 array (every type except the PTQ1_0 cases below) + GGML_CUDA_Q8_1_SOA_ISUM = 1, // PTQ1_0, one column: warp-transposed, exact int sums (ggml_cuda_ptq1_q8_word) + GGML_CUDA_Q8_1_PT = 2, // PTQ1_0, 2-8 columns or MoE ids: planar-transposed (mmvq-ptq1_0.cuh) +}; + +// Column-count helper, not the layout decision. 2-8 columns and MoE (ids) take the planar +// kernel. One column returns the Ada default (SOA_ISUM). Ampere and GGML_CUDA_BATCH_INVARIANT +// override that to PT in ggml_cuda_q8_1_layout_host, which is what the quantizer and the kernel +// switch both call. This helper does not consult the compute capability. +static constexpr __host__ __device__ ggml_cuda_q8_1_layout ggml_cuda_q8_1_layout_for(ggml_type type_src0, int ncols_dst, bool has_ids) { +#if defined(GGML_USE_HIP) + GGML_UNUSED(type_src0); GGML_UNUSED(ncols_dst); GGML_UNUSED(has_ids); + return GGML_CUDA_Q8_1_AOS; +#else + if (type_src0 != GGML_TYPE_PTQ1_0) { + return GGML_CUDA_Q8_1_AOS; + } + return (ncols_dst == 1 && !has_ids) ? GGML_CUDA_Q8_1_SOA_ISUM : GGML_CUDA_Q8_1_PT; +#endif +} + +// ggml_cuda_q8_1_layout_host is defined below ggml_cuda_info(); it reads the device +// compute capability, which is not declared yet at this point in the header. + +// Warp-transposed (SoA) q8_1 activation layout for the ternary MMVQ. +// +// One PTQ1_0 K-block (128 weights) consumes 4 block_q8_1 = 36 words (32 qs + 4 ds). In the +// small-K MMVQ geometry each lane owns one K-block, so with the plain AoS layout a warp-wide +// load of "word w" touches 32 lines 144 B apart: ~36 L1 wavefronts per instruction, ~1300 per +// K-iteration against ~250 for the weights themselves. That LSU traffic, not GDDR, capped the +// PTQ1 GEMV near 370 GB/s on Ada. Here K-blocks are grouped by 32 and word w of the group is +// stored contiguously, so the same load is 32 consecutive words = 1 wavefront. +// Bytes per column are unchanged when K is padded to a multiple of 32*128 = 4096. +#define GGML_CUDA_PTQ1_Q8_GROUP_KB 32 +#define GGML_CUDA_PTQ1_Q8_WORDS_PER_KB 36 +#define GGML_CUDA_PTQ1_Q8_GROUP_WORDS (GGML_CUDA_PTQ1_Q8_GROUP_KB * GGML_CUDA_PTQ1_Q8_WORDS_PER_KB) +#define GGML_CUDA_PTQ1_K_PAD (GGML_CUDA_PTQ1_Q8_GROUP_KB * QK_PTQ1_0) + +// Word offset (within one activation column) of word w (0..7 = qs words, 8 = ds) of block_q8_1 ib. +static constexpr __host__ __device__ int ggml_cuda_ptq1_q8_word(int ib, int w) { + const int kb = ib >> 2; + const int sub = ib & 3; + const int g = kb >> 5; + const int lane = kb & 31; + const int ww = w < 8 ? sub * 8 + w : 32 + sub; + return g * GGML_CUDA_PTQ1_Q8_GROUP_WORDS + ww * GGML_CUDA_PTQ1_Q8_GROUP_KB + lane; +} + template<> struct ggml_cuda_type_traits { static constexpr int qk = QK4_0; @@ -1205,6 +1279,26 @@ const ggml_cuda_device_info & ggml_cuda_info(); void ggml_cuda_set_device(int device); int ggml_cuda_get_device(); +// Host-side wrapper used by both the quantizer call and the kernel switch. +// Under GGML_CUDA_BATCH_INVARIANT the one-column case must run the same arithmetic as 2-8 +// columns, so it takes the planar layout (the SoA vec-dot sums in a different order). +// Ampere (sm_80/86, including the 3060/3090/170HX): the #218 PT kernel wins at one column +// too (+5.9% tg128 vs SoA on a 3060). Ada and newer keep SOA_ISUM at one column (4070 win). +// Lives here, after ggml_cuda_info(), because the body reads the current device's cc. +static inline ggml_cuda_q8_1_layout ggml_cuda_q8_1_layout_host(ggml_type type_src0, int ncols_dst, bool has_ids) { + const ggml_cuda_q8_1_layout l = ggml_cuda_q8_1_layout_for(type_src0, ncols_dst, has_ids); + if (l == GGML_CUDA_Q8_1_SOA_ISUM && ggml_cuda_batch_invariant()) { + return GGML_CUDA_Q8_1_PT; + } + if (l == GGML_CUDA_Q8_1_SOA_ISUM) { + const int cc = ggml_cuda_info().devices[ggml_cuda_get_device()].cc; + if (GGML_CUDA_CC_IS_NVIDIA(cc) && cc >= GGML_CUDA_CC_AMPERE && cc < GGML_CUDA_CC_ADA_LOVELACE) { + return GGML_CUDA_Q8_1_PT; + } + } + return l; +} + struct ggml_cuda_pool { virtual ~ggml_cuda_pool() = default; @@ -1453,6 +1547,79 @@ struct ggml_cuda_stream_context { } }; +// Fused recurrent-state gather for GATED_DELTA_NET. build_rs materialises GET_ROWS(cache, s_copy) +// into a temp per layer that only the GDN kernel reads; when the graph evaluator can prove that, it +// skips the GET_ROWS and records the gather here so the kernel reads cache row ids[seq] directly. +struct ggml_cuda_gated_delta_net_gather { + const float * base = nullptr; // cache rows [row_stride floats each] + const int32_t * ids = nullptr; // per-seq row index + int64_t row_stride = 0; // in floats +}; + +// Owned by the backend context that evaluates the graph: registrations are keyed by node pointer, +// so they are only meaningful for the evaluation that made them. Reset at the start of every +// graph evaluation/capture; never shared between contexts or threads. +struct ggml_cuda_gdn_gather_context { + std::unordered_map gathers; + + void reset() { + gathers.clear(); + } + + void set(const ggml_tensor * gdn, const ggml_cuda_gated_delta_net_gather & gather) { + gathers[gdn] = gather; + } + + const ggml_cuda_gated_delta_net_gather * find(const ggml_tensor * gdn) const { + const auto it = gathers.find(gdn); + return it == gathers.end() ? nullptr : &it->second; + } +}; + +// A Hadamard transform (MUL_MAT with GGML_HINT_SRC0_IS_HADAMARD, optionally preceded by the sign +// flip) whose only consumers are PTQ1_0 mat-vecs writes the q8_1-quantized activation into its own +// output buffer instead of the F32 result (the quantized rows are 9/8 bytes per element, so they +// fit in the F32 allocation), and the mat-vecs skip their quantize launch. The output buffer's +// lifetime is exactly the consumers' lifetime, so nothing about allocation changes. Keyed by tensor +// identity (the transform output and every reshape view of it that a mat-vec consumes), never by +// data pointer: ggml-alloc recycles a dead output's block for later tensors in the same graph, and a +// later PTQ1_0 mat-vec whose src1 landed there must not mistake its F32 rows for q8_1. Valid for one +// graph evaluation. +struct ggml_cuda_fwht_q8 { + ggml_cuda_q8_1_layout layout = GGML_CUDA_Q8_1_AOS; + int64_t ne0 = 0; // padded row width the quantizer wrote (what the mat-vec expects) + int64_t ncols = 0; // rows quantized (src1->ne[1] of every consumer) + const void * data = nullptr; // the q8_1 rows: the transform's own output buffer, or a pool block + // held for the rest of the graph when that buffer aliases the input +}; + +struct ggml_cuda_fwht_q8_context { + std::unordered_map entries; + std::vector>> held; // released at reset(), newest first (VMM pool is LIFO) + + // the implicit destructor would release `held` oldest first, which the VMM pool asserts on. The + // owning context also calls reset() before its pools go away (this member is declared before them). + ~ggml_cuda_fwht_q8_context() { + reset(); + } + + void reset() { + entries.clear(); + while (!held.empty()) { + held.pop_back(); + } + } + + void set(const ggml_tensor * out, const ggml_cuda_fwht_q8 & e) { + entries[out] = e; + } + + const ggml_cuda_fwht_q8 * find(const ggml_tensor * out) const { + const auto it = entries.find(out); + return it == entries.end() ? nullptr : &it->second; + } +}; + struct ggml_backend_cuda_context { int device; std::string name; @@ -1523,6 +1690,8 @@ struct ggml_backend_cuda_context { } ggml_cuda_stream_context concurrent_stream_context; + ggml_cuda_gdn_gather_context gdn_gather_context; + ggml_cuda_fwht_q8_context fwht_q8_context; ~ggml_backend_cuda_context(); @@ -1538,6 +1707,10 @@ struct ggml_backend_cuda_context { ggml_cuda_stream_context & stream_context() { return concurrent_stream_context; } + ggml_cuda_gdn_gather_context & gdn_gathers() { return gdn_gather_context; } + + ggml_cuda_fwht_q8_context & fwht_q8() { return fwht_q8_context; } + cublasHandle_t cublas_handle() { if (cublas_handles[device][curr_stream_no] == nullptr) { ggml_cuda_set_device(device); diff --git a/ggml/src/ggml-cuda/fattn-common.cuh b/ggml/src/ggml-cuda/fattn-common.cuh index e67cc7fdf784..ae4318cbc881 100644 --- a/ggml/src/ggml-cuda/fattn-common.cuh +++ b/ggml/src/ggml-cuda/fattn-common.cuh @@ -1154,11 +1154,14 @@ void launch_fattn( // If ntiles_total % blocks_per_wave != 0 then some efficiency is lost due to tail effects. // Test whether parallel_blocks can be set to a higher value for better efficiency. + // Batch-invariant mode: size the KV split as for a single query tile, so the order in which the + // partial softmax results are combined does not depend on how many queries are in the batch. + const int ntiles_dst_eff = ggml_cuda_batch_invariant() ? ntiles_dst / ntiles_x : ntiles_dst; const int blocks_per_wave = nsm * max_blocks_per_sm; int nwaves_best = 0; int efficiency_percent_best = 0; for (int parallel_blocks_test = parallel_blocks; parallel_blocks_test <= ntiles_KV; ++parallel_blocks_test) { - const int nblocks_total = ntiles_dst * parallel_blocks_test; + const int nblocks_total = ntiles_dst_eff * parallel_blocks_test; const int nwaves = (nblocks_total + blocks_per_wave - 1) / blocks_per_wave; const int efficiency_percent = 100 * nblocks_total / (nwaves*blocks_per_wave); diff --git a/ggml/src/ggml-cuda/fattn-mma-f16.cuh b/ggml/src/ggml-cuda/fattn-mma-f16.cuh index 7f4cfd5511ff..0783e0d86329 100644 --- a/ggml/src/ggml-cuda/fattn-mma-f16.cuh +++ b/ggml/src/ggml-cuda/fattn-mma-f16.cuh @@ -358,6 +358,22 @@ static constexpr __device__ int ggml_cuda_fattn_mma_get_nstages(const int DKQ, c #endif // CP_ASYNC_AVAILABLE } +// K/V read straight from a quantized cache. The tile loader dequantizes into the F16 shared tile, so there is no +// F16 copy of the whole cache in the compute buffer (n_ctx*D*n_head_kv*2 bytes for K and again for V, reserved +// for the worst case at graph allocation) and no full-cache conversion pass per micro-batch. cp.async cannot +// dequantize, so these variants always run the single-stage register loader. +static constexpr __host__ __device__ bool ggml_cuda_fattn_mma_kv_native(const ggml_type type_KV) { + return type_KV == GGML_TYPE_Q4_0 || type_KV == GGML_TYPE_Q8_0; +} + +static __host__ int ggml_cuda_fattn_mma_get_nstages(const int DKQ, const int DV, const int ncols1, const int ncols2, const int cc, const ggml_type type_KV) { + return type_KV == GGML_TYPE_F16 ? ggml_cuda_fattn_mma_get_nstages(DKQ, DV, ncols1, ncols2, cc) : 0; +} + +static constexpr __device__ int ggml_cuda_fattn_mma_get_nstages(const int DKQ, const int DV, const int ncols1, const int ncols2, const ggml_type type_KV) { + return type_KV == GGML_TYPE_F16 ? ggml_cuda_fattn_mma_get_nstages(DKQ, DV, ncols1, ncols2) : 0; +} + // ------------------------------------------------------------------------------------------------------------------ template @@ -447,6 +463,96 @@ static __device__ __forceinline__ void flash_attn_ext_f16_load_tile( } } +// Dequantize one q4_0 block (32 weights) given its scale and the 16 nibble bytes as 4 words, store 16 half2. +// Same arithmetic as dequantize_q4_0 -> half, so the tile is bit-identical to a to_fp16 pass over the cache. +static __device__ __forceinline__ void flash_attn_ext_f16_store_q4_0_block( + half2 * const __restrict__ dst, const float d, const uint32_t q0, const uint32_t q1, const uint32_t q2, const uint32_t q3) { + const uint32_t q[4] = {q0, q1, q2, q3}; +#pragma unroll + for (int m = 0; m < 4; ++m) { + const uint32_t v = q[m]; + // low nibbles are weights 4m..4m+3, high nibbles are weights 16+4m..16+4m+3 + dst[ 2*m ] = __floats2half2_rn(d*(float(int((v ) & 0xF)) - 8.0f), d*(float(int((v >> 8) & 0xF)) - 8.0f)); + dst[ 2*m + 1] = __floats2half2_rn(d*(float(int((v >> 16) & 0xF)) - 8.0f), d*(float(int((v >> 24) & 0xF)) - 8.0f)); + dst[8 + 2*m ] = __floats2half2_rn(d*(float(int((v >> 4) & 0xF)) - 8.0f), d*(float(int((v >> 12) & 0xF)) - 8.0f)); + dst[8 + 2*m + 1] = __floats2half2_rn(d*(float(int((v >> 20) & 0xF)) - 8.0f), d*(float(int((v >> 28) & 0xF)) - 8.0f)); + } +} + +// Dequantize one q8_0 block (32 weights) given its scale and the 32 int8 bytes as 8 words, store 16 half2. +static __device__ __forceinline__ void flash_attn_ext_f16_store_q8_0_block( + half2 * const __restrict__ dst, const float d, const uint32_t * const __restrict__ q) { +#pragma unroll + for (int m = 0; m < 8; ++m) { + const uint32_t v = q[m]; + dst[2*m ] = __floats2half2_rn(d*float(int(int8_t(v & 0xFF))), d*float(int(int8_t((v >> 8) & 0xFF)))); + dst[2*m + 1] = __floats2half2_rn(d*float(int(int8_t((v >> 16) & 0xFF))), d*float(int(int8_t((v >> 24) & 0xFF)))); + } +} + +// Loads a [nbatch_fa x ne] tile of a q4_0 / q8_0 K or V cache and dequantizes it into the F16 shared tile, in the +// same layout the F16 loader produces. KV points at byte 0 of tile row 0, e0 / ne are the element offset and count +// within a row (multiples of 64), stride_KV is the row stride in bytes. Each thread handles two adjacent blocks +// (64 weights): a q4_0 pair is 36 bytes and a q8_0 pair 68 bytes, both 4-byte aligned at even block offsets, so the +// pair is read as aligned 32-bit words and the odd block's bytes are realigned with byte permutes. +template +static __device__ __forceinline__ void flash_attn_ext_f16_load_tile_q( + const char * const __restrict__ KV, half2 * const __restrict__ tile_KV, + const int e0, const int ne, const int stride_KV, const int i_sup) { + static_assert(type_KV == GGML_TYPE_Q4_0 || type_KV == GGML_TYPE_Q8_0, "unsupported quantized K/V type"); + constexpr int warp_size = ggml_cuda_get_physical_warp_size(); + constexpr int nthreads = nwarps*warp_size; + constexpr int bs = 32; // weights per block, both types + constexpr int ts = type_KV == GGML_TYPE_Q4_0 ? 18 : 34; // bytes per block + constexpr int nwords = 2*ts/4; // 32-bit words per block pair: 9 or 17 + static_assert((2*ts) % 4 == 0, "block pair is not word aligned"); + + const int pairs_per_row = ne / (2*bs); + const int ntotal = nbatch_fa * pairs_per_row; + + for (int t = threadIdx.y*warp_size + threadIdx.x; t < ntotal; t += nthreads) { + const int i = t / pairs_per_row; + const int p = t - i*pairs_per_row; + half2 * dst = tile_KV + i*stride_tile + p*bs; // 2 blocks = 64 weights = 32 half2 per pair + + if (oob_check && i >= i_sup) { +#pragma unroll + for (int k = 0; k < bs; ++k) { + dst[k] = make_half2(0.0f, 0.0f); + } + continue; + } + + const uint32_t * src = (const uint32_t *) (KV + int64_t(i)*stride_KV + ((e0 + p*2*bs)/bs)*ts); + uint32_t w[nwords]; +#pragma unroll + for (int k = 0; k < nwords; ++k) { + w[k] = __ldg(src + k); + } + + if constexpr (type_KV == GGML_TYPE_Q4_0) { + // block 0: d = bytes 0-1, qs = bytes 2-17. block 1: d = bytes 18-19, qs = bytes 20-35 (= w5..w8). + const float d0 = __half2float(__ushort_as_half((unsigned short) (w[0] & 0xFFFF))); + const float d1 = __half2float(__ushort_as_half((unsigned short) (w[4] >> 16))); + flash_attn_ext_f16_store_q4_0_block(dst, d0, + __byte_perm(w[0], w[1], 0x5432), __byte_perm(w[1], w[2], 0x5432), + __byte_perm(w[2], w[3], 0x5432), __byte_perm(w[3], w[4], 0x5432)); + flash_attn_ext_f16_store_q4_0_block(dst + bs/2, d1, w[5], w[6], w[7], w[8]); + } else { + // block 0: d = bytes 0-1, qs = bytes 2-33. block 1: d = bytes 34-35, qs = bytes 36-67 (= w9..w16). + const float d0 = __half2float(__ushort_as_half((unsigned short) (w[0] & 0xFFFF))); + const float d1 = __half2float(__ushort_as_half((unsigned short) (w[8] >> 16))); + uint32_t q0[8]; +#pragma unroll + for (int k = 0; k < 8; ++k) { + q0[k] = __byte_perm(w[k], w[k + 1], 0x5432); + } + flash_attn_ext_f16_store_q8_0_block(dst, d0, q0); + flash_attn_ext_f16_store_q8_0_block(dst + bs/2, d1, w + 9); + } + } +} + template static __device__ __forceinline__ void flash_attn_ext_f16_load_mask( const half * const __restrict__ mask_h, half * const __restrict__ tile_mask, @@ -529,7 +635,8 @@ static __device__ __forceinline__ void flash_attn_ext_f16_load_mask( template + typename T_A_KQ, typename T_B_KQ, typename T_C_KQ, typename T_A_VKQ, typename T_B_VKQ, typename T_C_VKQ, + ggml_type type_KV> static __device__ __forceinline__ void flash_attn_ext_f16_iter( const float2 * const __restrict__ Q_f2, const half2 * const __restrict__ K_h2, @@ -566,12 +673,16 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( constexpr int nbatch_K2 = ggml_cuda_fattn_mma_get_nbatch_K2(DKQ, DV, ncols); constexpr int nbatch_V2 = ggml_cuda_fattn_mma_get_nbatch_V2(DKQ, DV, ncols); constexpr bool Q_in_reg = ggml_cuda_fattn_mma_get_Q_in_reg (DKQ, DV, ncols); - constexpr int nstages = ggml_cuda_fattn_mma_get_nstages (DKQ, DV, ncols1, ncols2); + constexpr int nstages = ggml_cuda_fattn_mma_get_nstages (DKQ, DV, ncols1, ncols2, type_KV); constexpr int stride_tile_K = nbatch_K2 + 4; constexpr int stride_tile_V = V_is_K_view ? stride_tile_K : nbatch_V2 + 4; + // For quantized K/V, stride_K / stride_V are byte strides and the loader dequantizes block pairs (64 weights). + static_assert(type_KV == GGML_TYPE_F16 || ((2*nbatch_K2) % 64 == 0 && (2*nbatch_V2) % 64 == 0 && !V_is_K_view), + "quantized K/V loading needs whole block pairs per pass"); + const int k_VKQ_0 = kb0 * nbatch_fa; #if defined(TURING_MMA_AVAILABLE) T_C_KQ KQ_C[nbatch_fa/(np*(cols_per_warp == 8 ? T_C_KQ::I : T_C_KQ::J))]; @@ -606,11 +717,16 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( if constexpr (nstages <= 1) { const int k0_diff = k0_stop - k0_start; - constexpr bool use_cp_async = nstages == 1; - flash_attn_ext_f16_load_tile - (K_h2 + int64_t(k_VKQ_0)*stride_K + k0_start, tile_K, k0_diff, stride_K, k_VKQ_sup); - if (use_cp_async) { - cp_async_wait_all(); + if constexpr (type_KV == GGML_TYPE_F16) { + constexpr bool use_cp_async = nstages == 1; + flash_attn_ext_f16_load_tile + (K_h2 + int64_t(k_VKQ_0)*stride_K + k0_start, tile_K, k0_diff, stride_K, k_VKQ_sup); + if (use_cp_async) { + cp_async_wait_all(); + } + } else { + flash_attn_ext_f16_load_tile_q + ((const char *) K_h2 + int64_t(k_VKQ_0)*stride_K, tile_K, 2*k0_start, 2*k0_diff, stride_K, k_VKQ_sup); } __syncthreads(); } @@ -958,11 +1074,16 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter( if constexpr (nstages <= 1) { const int i0_diff = i0_stop - i0_start; if (!V_is_K_view || i0_stop > 2*nbatch_K2) { - constexpr bool use_cp_async = nstages == 1; - flash_attn_ext_f16_load_tile - (V_h2 + int64_t(k_VKQ_0)*stride_V + i0_start/2, tile_V, i0_diff/2, stride_V, k_VKQ_sup); - if (use_cp_async) { - cp_async_wait_all(); + if constexpr (type_KV == GGML_TYPE_F16) { + constexpr bool use_cp_async = nstages == 1; + flash_attn_ext_f16_load_tile + (V_h2 + int64_t(k_VKQ_0)*stride_V + i0_start/2, tile_V, i0_diff/2, stride_V, k_VKQ_sup); + if (use_cp_async) { + cp_async_wait_all(); + } + } else { + flash_attn_ext_f16_load_tile_q + ((const char *) V_h2 + int64_t(k_VKQ_0)*stride_V, tile_V, i0_start, i0_diff, stride_V, k_VKQ_sup); } __syncthreads(); } @@ -1113,7 +1234,7 @@ template struct mma_tile_sizes { }; #endif // defined(TURING_MMA_AVAILABLE) -template +template static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( const float2 * const __restrict__ Q_f2, const half2 * const __restrict__ K_h2, @@ -1158,7 +1279,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( constexpr int nbatch_V2 = ggml_cuda_fattn_mma_get_nbatch_V2 (DKQ, DV, ncols); constexpr int nbatch_combine = ggml_cuda_fattn_mma_get_nbatch_combine(DKQ, DV, ncols); constexpr bool Q_in_reg = ggml_cuda_fattn_mma_get_Q_in_reg (DKQ, DV, ncols); - constexpr int nstages = ggml_cuda_fattn_mma_get_nstages (DKQ, DV, ncols1, ncols2); + constexpr int nstages = ggml_cuda_fattn_mma_get_nstages (DKQ, DV, ncols1, ncols2, type_KV); if (cols_per_warp > ncols) { NO_DEVICE_CODE; @@ -1277,7 +1398,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( constexpr int k_VKQ_sup = nbatch_fa; flash_attn_ext_f16_iter + T_A_KQ, T_B_KQ, T_C_KQ, T_A_VKQ, T_B_VKQ, T_C_VKQ, type_KV> (Q_f2, K_h2, V_h2, mask_h, dstk, dstk_fixup, scale, slope, logit_softcap, ne01, ne02, stride_K, stride_V, stride_mask, tile_Q, tile_K, tile_V, tile_mask, Q_B, VKQ_C, KQ_max, KQ_rowsum, jt, kb0, k_VKQ_sup); @@ -1286,7 +1407,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( const int k_VKQ_sup = ne11 - kb0*nbatch_fa; flash_attn_ext_f16_iter + T_A_KQ, T_B_KQ, T_C_KQ, T_A_VKQ, T_B_VKQ, T_C_VKQ, type_KV> (Q_f2, K_h2, V_h2, mask_h, dstk, dstk_fixup, scale, slope, logit_softcap, ne01, ne02, stride_K, stride_V, stride_mask, tile_Q, tile_K, tile_V, tile_mask, Q_B, VKQ_C, KQ_max, KQ_rowsum, jt, kb0, k_VKQ_sup); @@ -1297,7 +1418,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( constexpr int k_VKQ_sup = nbatch_fa; flash_attn_ext_f16_iter + T_A_KQ, T_B_KQ, T_C_KQ, T_A_VKQ, T_B_VKQ, T_C_VKQ, type_KV> (Q_f2, K_h2, V_h2, mask_h, dstk, dstk_fixup, scale, slope, logit_softcap, ne01, ne02, stride_K, stride_V, stride_mask, tile_Q, tile_K, tile_V, tile_mask, Q_B, VKQ_C, KQ_max, KQ_rowsum, jt, kb0, k_VKQ_sup); @@ -1306,7 +1427,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( constexpr int k_VKQ_sup = nbatch_fa; flash_attn_ext_f16_iter + T_A_KQ, T_B_KQ, T_C_KQ, T_A_VKQ, T_B_VKQ, T_C_VKQ, type_KV> (Q_f2, K_h2, V_h2, mask_h, dstk, dstk_fixup, scale, slope, logit_softcap, ne01, ne02, stride_K, stride_V, stride_mask, tile_Q, tile_K, tile_V, tile_mask, Q_B, VKQ_C, KQ_max, KQ_rowsum, jt, kb0, k_VKQ_sup); @@ -1700,7 +1821,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile( #endif // defined(VOLTA_MMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE) || defined(AMD_MFMA_AVAILABLE) } -template +template __launch_bounds__(ggml_cuda_fattn_mma_get_nthreads(DKQ, DV, ncols1*ncols2), ggml_cuda_fattn_mma_get_occupancy(DKQ, DV, ncols1*ncols2)) static __global__ void flash_attn_ext_f16( const char * Q_ptr, @@ -1782,10 +1903,11 @@ static __global__ void flash_attn_ext_f16( const int stride_Q1 = nb01 / sizeof(float2); const int stride_Q2 = nb02 / sizeof(float2); - const int stride_K = nb11 / sizeof(half2); + // Quantized K/V: byte strides, the tile loader dequantizes block pairs straight out of the cache. + const int stride_K = type_KV == GGML_TYPE_F16 ? nb11 / sizeof(half2) : nb11; const int stride_mask = nb31 / sizeof(half); - const int stride_V = V_is_K_view ? stride_K : nb21 / sizeof(half2); + const int stride_V = V_is_K_view ? stride_K : (type_KV == GGML_TYPE_F16 ? nb21 / sizeof(half2) : nb21); const int iter_k = (ne11 + (nbatch_fa - 1)) / nbatch_fa; const int iter_j = (ne01.z + (ncols1 - 1)) / ncols1; @@ -1829,12 +1951,12 @@ static __global__ void flash_attn_ext_f16( constexpr bool is_fixup = false; // All but (potentially) the last iterations write their data to dst rather than the fixup buffer. if (kb0_start == 0) { constexpr bool needs_fixup = false; // CUDA block is working on an entire tile. - flash_attn_ext_f16_process_tile + flash_attn_ext_f16_process_tile (Q_f2, K_h2, V_h2, mask_h, sinks_f, dstk, dst_meta, scale, slope, logit_softcap, ne01, ne02, gqa_ratio, ne11, stride_Q1, stride_Q2, stride_K, stride_V, stride_mask, jt, zt_gqa, kb0_start, kb0_stop); } else { constexpr bool needs_fixup = true; // CUDA block is missing the beginning of a tile. - flash_attn_ext_f16_process_tile + flash_attn_ext_f16_process_tile (Q_f2, K_h2, V_h2, mask_h, sinks_f, dstk, dst_meta, scale, slope, logit_softcap, ne01, ne02, gqa_ratio, ne11, stride_Q1, stride_Q2, stride_K, stride_V, stride_mask, jt, zt_gqa, kb0_start, kb0_stop); } @@ -1875,7 +1997,7 @@ static __global__ void flash_attn_ext_f16( constexpr bool is_fixup = true; // Last index writes its data to fixup buffer to avoid data races with other blocks. constexpr bool needs_fixup = false; - flash_attn_ext_f16_process_tile + flash_attn_ext_f16_process_tile (Q_f2, K_h2, V_h2, mask_h, sinks_f, dstk, dst_meta, scale, slope, logit_softcap, ne01, ne02, gqa_ratio, ne11, stride_Q1, stride_Q2, stride_K, stride_V, stride_mask, jt, zt_gqa, kb0_start, kb0_stop); #else @@ -1892,7 +2014,7 @@ static __global__ void flash_attn_ext_f16( #endif // defined(FLASH_ATTN_AVAILABLE) && (defined(VOLTA_MMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE) || defined(AMD_MFMA_AVAILABLE)) } -template +template void ggml_cuda_flash_attn_ext_mma_f16_case(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { const ggml_tensor * KQV = dst; const int id = ggml_cuda_get_device(); @@ -1900,13 +2022,22 @@ void ggml_cuda_flash_attn_ext_mma_f16_case(ggml_backend_cuda_context & ctx, ggml constexpr int ncols = ncols1 * ncols2; + if constexpr (type_KV != GGML_TYPE_F16) { + // The kernel reads the quantized cache in place: same K/V type, whole block pairs per row, word-aligned rows. + const ggml_tensor * K = dst->src[1]; + const ggml_tensor * V = dst->src[2]; + GGML_ASSERT(K->type == type_KV && V->type == type_KV); + GGML_ASSERT(K->nb[1] % 4 == 0 && K->nb[2] % 4 == 0 && K->nb[3] % 4 == 0); + GGML_ASSERT(V->nb[1] % 4 == 0 && V->nb[2] % 4 == 0 && V->nb[3] % 4 == 0); + } + const int nthreads = ggml_cuda_fattn_mma_get_nthreads (DKQ, DV, ncols, cc); const int nbatch_fa = ggml_cuda_fattn_mma_get_nbatch_fa (DKQ, DV, ncols, cc); const int nbatch_K2 = ggml_cuda_fattn_mma_get_nbatch_K2 (DKQ, DV, ncols, cc); const int nbatch_V2 = ggml_cuda_fattn_mma_get_nbatch_V2 (DKQ, DV, ncols, cc); const int nbatch_combine = ggml_cuda_fattn_mma_get_nbatch_combine(DKQ, DV, ncols, cc); const bool Q_in_reg = ggml_cuda_fattn_mma_get_Q_in_reg (DKQ, DV, ncols, cc); - const int nstages = ggml_cuda_fattn_mma_get_nstages (DKQ, DV, ncols1, ncols2, cc); + const int nstages = ggml_cuda_fattn_mma_get_nstages (DKQ, DV, ncols1, ncols2, cc, type_KV); const int cols_per_warp = std::min(ncols, get_cols_per_warp(cc)); const int warp_size_host = ggml_cuda_info().devices[ctx.device].warp_size; @@ -1937,7 +2068,7 @@ void ggml_cuda_flash_attn_ext_mma_f16_case(ggml_backend_cuda_context & ctx, ggml fattn_kernel_t fattn_kernel; if (logit_softcap == 0.0f) { constexpr bool use_logit_softcap = false; - fattn_kernel = flash_attn_ext_f16; + fattn_kernel = flash_attn_ext_f16; #if !defined(GGML_USE_MUSA) static bool shared_memory_limit_raised[GGML_CUDA_MAX_DEVICES] = {false}; @@ -1948,7 +2079,7 @@ void ggml_cuda_flash_attn_ext_mma_f16_case(ggml_backend_cuda_context & ctx, ggml #endif // !defined(GGML_USE_MUSA) } else { constexpr bool use_logit_softcap = true; - fattn_kernel = flash_attn_ext_f16; + fattn_kernel = flash_attn_ext_f16; #if !defined(GGML_USE_MUSA) static bool shared_memory_limit_raised[GGML_CUDA_MAX_DEVICES] = {false}; @@ -1959,8 +2090,9 @@ void ggml_cuda_flash_attn_ext_mma_f16_case(ggml_backend_cuda_context & ctx, ggml #endif // !defined(GGML_USE_MUSA) } + constexpr bool need_f16_KV = type_KV == GGML_TYPE_F16; // quantized variants read the cache in place launch_fattn - (ctx, dst, fattn_kernel, nwarps, nbytes_shared_total, nbatch_fa, true, true, true, warp_size_host); + (ctx, dst, fattn_kernel, nwarps, nbytes_shared_total, nbatch_fa, need_f16_KV, need_f16_KV, true, warp_size_host); } @@ -1968,6 +2100,58 @@ void ggml_cuda_flash_attn_ext_mma_f16_case(ggml_backend_cuda_context & ctx, ggml template void ggml_cuda_flash_attn_ext_mma_f16_case \ (ggml_backend_cuda_context & ctx, ggml_tensor * dst) \ +// In-place quantized K/V variants, instantiated for the head sizes and GQA tilings a quantized cache is served with. +#define DECL_FATTN_MMA_F16_CASE_KV(DKQ, DV, ncols1, ncols2, type_KV) \ + template void ggml_cuda_flash_attn_ext_mma_f16_case \ + (ggml_backend_cuda_context & ctx, ggml_tensor * dst) \ + +#define DECL_FATTN_MMA_F16_CASE_KV_ALL_NCOLS2(DKQ, DV, ncols, type_KV) \ + extern DECL_FATTN_MMA_F16_CASE_KV(DKQ, DV, (ncols)/ 1, 1, type_KV); \ + extern DECL_FATTN_MMA_F16_CASE_KV(DKQ, DV, (ncols)/ 2, 2, type_KV); \ + extern DECL_FATTN_MMA_F16_CASE_KV(DKQ, DV, (ncols)/ 4, 4, type_KV); \ + extern DECL_FATTN_MMA_F16_CASE_KV(DKQ, DV, (ncols)/ 8, 8, type_KV); \ + +DECL_FATTN_MMA_F16_CASE_KV_ALL_NCOLS2(128, 128, 8, GGML_TYPE_Q4_0) +DECL_FATTN_MMA_F16_CASE_KV_ALL_NCOLS2(128, 128, 16, GGML_TYPE_Q4_0) +DECL_FATTN_MMA_F16_CASE_KV_ALL_NCOLS2(128, 128, 32, GGML_TYPE_Q4_0) +DECL_FATTN_MMA_F16_CASE_KV_ALL_NCOLS2(128, 128, 64, GGML_TYPE_Q4_0) +DECL_FATTN_MMA_F16_CASE_KV_ALL_NCOLS2(256, 256, 8, GGML_TYPE_Q4_0) +DECL_FATTN_MMA_F16_CASE_KV_ALL_NCOLS2(256, 256, 16, GGML_TYPE_Q4_0) +DECL_FATTN_MMA_F16_CASE_KV_ALL_NCOLS2(256, 256, 32, GGML_TYPE_Q4_0) +DECL_FATTN_MMA_F16_CASE_KV_ALL_NCOLS2(256, 256, 64, GGML_TYPE_Q4_0) + +DECL_FATTN_MMA_F16_CASE_KV_ALL_NCOLS2(128, 128, 8, GGML_TYPE_Q8_0) +DECL_FATTN_MMA_F16_CASE_KV_ALL_NCOLS2(128, 128, 16, GGML_TYPE_Q8_0) +DECL_FATTN_MMA_F16_CASE_KV_ALL_NCOLS2(128, 128, 32, GGML_TYPE_Q8_0) +DECL_FATTN_MMA_F16_CASE_KV_ALL_NCOLS2(128, 128, 64, GGML_TYPE_Q8_0) +DECL_FATTN_MMA_F16_CASE_KV_ALL_NCOLS2(256, 256, 8, GGML_TYPE_Q8_0) +DECL_FATTN_MMA_F16_CASE_KV_ALL_NCOLS2(256, 256, 16, GGML_TYPE_Q8_0) +DECL_FATTN_MMA_F16_CASE_KV_ALL_NCOLS2(256, 256, 32, GGML_TYPE_Q8_0) +DECL_FATTN_MMA_F16_CASE_KV_ALL_NCOLS2(256, 256, 64, GGML_TYPE_Q8_0) + +// Head sizes with in-place quantized K/V kernels (matches the instantiations above). +static inline bool ggml_cuda_fattn_mma_kv_native_supported(const ggml_tensor * dst) { +#if defined(GGML_USE_HIP) || defined(GGML_USE_MUSA) + GGML_UNUSED(dst); + return false; +#else + const int cc = ggml_cuda_info().devices[ggml_cuda_get_device()].cc; + if (!GGML_CUDA_CC_IS_NVIDIA(cc)) { + return false; + } + const ggml_tensor * Q = dst->src[0]; + const ggml_tensor * K = dst->src[1]; + const ggml_tensor * V = dst->src[2]; + if (K->type != V->type || !ggml_cuda_fattn_mma_kv_native(K->type)) { + return false; + } + if (!(Q->ne[0] == 128 || Q->ne[0] == 256) || V->ne[0] != Q->ne[0]) { + return false; + } + return ggml_cuda_is_aligned(K, 4) && ggml_cuda_is_aligned(V, 4); +#endif +} + #define DECL_FATTN_MMA_F16_CASE_ALL_NCOLS2(DKQ, DV, ncols) \ extern DECL_FATTN_MMA_F16_CASE(DKQ, DV, (ncols)/ 1, 1); \ extern DECL_FATTN_MMA_F16_CASE(DKQ, DV, (ncols)/ 2, 2); \ diff --git a/ggml/src/ggml-cuda/fattn.cu b/ggml/src/ggml-cuda/fattn.cu index ab7a3b297c07..59e6c9ba74b0 100644 --- a/ggml/src/ggml-cuda/fattn.cu +++ b/ggml/src/ggml-cuda/fattn.cu @@ -5,35 +5,35 @@ #include "fattn-vec.cuh" #include "fattn.cuh" -template +template static void ggml_cuda_flash_attn_ext_mma_f16_switch_ncols1(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { const int cc = ggml_cuda_info().devices[ggml_cuda_get_device()].cc; const ggml_tensor * Q = dst->src[0]; if constexpr (ncols2 <= 8) { if (turing_mma_available(cc) && Q->ne[1] <= 8/ncols2) { - ggml_cuda_flash_attn_ext_mma_f16_case(ctx, dst); + ggml_cuda_flash_attn_ext_mma_f16_case(ctx, dst); return; } } if constexpr (ncols2 <= 16) { if (Q->ne[1] <= 16/ncols2) { - ggml_cuda_flash_attn_ext_mma_f16_case(ctx, dst); + ggml_cuda_flash_attn_ext_mma_f16_case(ctx, dst); return; } } if (Q->ne[1] <= 32/ncols2 || (GGML_CUDA_CC_IS_NVIDIA(cc) && ggml_cuda_highest_compiled_arch(cc) == GGML_CUDA_CC_TURING) || (GGML_CUDA_CC_IS_AMD(cc) && DKQ > 256)) { - ggml_cuda_flash_attn_ext_mma_f16_case(ctx, dst); + ggml_cuda_flash_attn_ext_mma_f16_case(ctx, dst); return; } - ggml_cuda_flash_attn_ext_mma_f16_case(ctx, dst); + ggml_cuda_flash_attn_ext_mma_f16_case(ctx, dst); } -template +template static void ggml_cuda_flash_attn_ext_mma_f16_switch_ncols2(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { const int cc = ggml_cuda_info().devices[ggml_cuda_get_device()].cc; const ggml_tensor * KQV = dst; @@ -66,22 +66,22 @@ static void ggml_cuda_flash_attn_ext_mma_f16_switch_ncols2(ggml_backend_cuda_con // On Volta the GQA optimizations aren't as impactful vs. minimizing wasted compute: if (cc == GGML_CUDA_CC_VOLTA) { if (use_gqa_opt && gqa_ratio % 8 == 0) { - ggml_cuda_flash_attn_ext_mma_f16_switch_ncols1(ctx, dst); + ggml_cuda_flash_attn_ext_mma_f16_switch_ncols1(ctx, dst); return; } if (use_gqa_opt && gqa_ratio % 4 == 0) { - ggml_cuda_flash_attn_ext_mma_f16_switch_ncols1(ctx, dst); + ggml_cuda_flash_attn_ext_mma_f16_switch_ncols1(ctx, dst); return; } if constexpr (DKQ <= 256) { if (use_gqa_opt && gqa_ratio % 2 == 0) { - ggml_cuda_flash_attn_ext_mma_f16_switch_ncols1(ctx, dst); + ggml_cuda_flash_attn_ext_mma_f16_switch_ncols1(ctx, dst); return; } - ggml_cuda_flash_attn_ext_mma_f16_switch_ncols1(ctx, dst); + ggml_cuda_flash_attn_ext_mma_f16_switch_ncols1(ctx, dst); return; } else { GGML_ABORT("fatal error"); @@ -89,22 +89,22 @@ static void ggml_cuda_flash_attn_ext_mma_f16_switch_ncols2(ggml_backend_cuda_con } if (use_gqa_opt && gqa_ratio > 4) { - ggml_cuda_flash_attn_ext_mma_f16_switch_ncols1(ctx, dst); + ggml_cuda_flash_attn_ext_mma_f16_switch_ncols1(ctx, dst); return; } if (use_gqa_opt && gqa_ratio > 2) { - ggml_cuda_flash_attn_ext_mma_f16_switch_ncols1(ctx, dst); + ggml_cuda_flash_attn_ext_mma_f16_switch_ncols1(ctx, dst); return; } if (use_gqa_opt && gqa_ratio > 1) { - ggml_cuda_flash_attn_ext_mma_f16_switch_ncols1(ctx, dst); + ggml_cuda_flash_attn_ext_mma_f16_switch_ncols1(ctx, dst); return; } if constexpr (DKQ <= 256) { - ggml_cuda_flash_attn_ext_mma_f16_switch_ncols1(ctx, dst); + ggml_cuda_flash_attn_ext_mma_f16_switch_ncols1(ctx, dst); } else { GGML_ABORT("fatal error"); } @@ -137,6 +137,14 @@ static void ggml_cuda_flash_attn_ext_mma_f16(ggml_backend_cuda_context & ctx, gg break; case 128: GGML_ASSERT(V->ne[0] == 128); + if (ggml_cuda_fattn_mma_kv_native_supported(dst)) { + if (K->type == GGML_TYPE_Q4_0) { + ggml_cuda_flash_attn_ext_mma_f16_switch_ncols2<128, 128, GGML_TYPE_Q4_0>(ctx, dst); + } else { + ggml_cuda_flash_attn_ext_mma_f16_switch_ncols2<128, 128, GGML_TYPE_Q8_0>(ctx, dst); + } + break; + } ggml_cuda_flash_attn_ext_mma_f16_switch_ncols2<128, 128>(ctx, dst); break; case 192: { @@ -157,6 +165,14 @@ static void ggml_cuda_flash_attn_ext_mma_f16(ggml_backend_cuda_context & ctx, gg } break; case 256: GGML_ASSERT(V->ne[0] == 256); + if (ggml_cuda_fattn_mma_kv_native_supported(dst)) { + if (K->type == GGML_TYPE_Q4_0) { + ggml_cuda_flash_attn_ext_mma_f16_switch_ncols2<256, 256, GGML_TYPE_Q4_0>(ctx, dst); + } else { + ggml_cuda_flash_attn_ext_mma_f16_switch_ncols2<256, 256, GGML_TYPE_Q8_0>(ctx, dst); + } + break; + } ggml_cuda_flash_attn_ext_mma_f16_switch_ncols2<256, 256>(ctx, dst); break; case 320: @@ -460,6 +476,11 @@ static best_fattn_kernel ggml_cuda_get_best_fattn_kernel(const int device, const // If Turing tensor cores are available, use them: if (turing_mma_available(cc) && Q->ne[0] != 40 && Q->ne[0] != 72) { if (can_use_vector_kernel) { + // batch-invariant mode: the same (vector) kernel for 1 to 8 queries, so a token verified in a + // speculative batch attends with the same arithmetic as a token decoded alone + if (ggml_cuda_batch_invariant() && Q->ne[1] <= 8 && Q->ne[3] == 1) { + return BEST_FATTN_KERNEL_VEC; + } if (!ggml_is_quantized(K->type) && !ggml_is_quantized(V->type)) { if (cc >= GGML_CUDA_CC_ADA_LOVELACE && Q->ne[1] == 1 && Q->ne[3] == 1 && !(gqa_ratio > 4 && K->ne[1] >= 8192)) { return BEST_FATTN_KERNEL_VEC; @@ -548,8 +569,15 @@ size_t ggml_cuda_flash_attn_ext_get_alloc_size(int device, const ggml_tensor * d bool need_f16_V = false; switch (kernel) { - case BEST_FATTN_KERNEL_TILE: case BEST_FATTN_KERNEL_MMA_F16: + if (ggml_cuda_fattn_mma_kv_native_supported(dst)) { + // In-place quantized K/V kernel: nothing to reserve beyond dst. + break; + } + need_f16_K = true; + need_f16_V = true; + break; + case BEST_FATTN_KERNEL_TILE: need_f16_K = true; need_f16_V = true; break; diff --git a/ggml/src/ggml-cuda/fwht.cu b/ggml/src/ggml-cuda/fwht.cu index 467a84a926ed..6ff8b0bc046b 100644 --- a/ggml/src/ggml-cuda/fwht.cu +++ b/ggml/src/ggml-cuda/fwht.cu @@ -167,49 +167,7 @@ __global__ void fwht_cuda_block(const T * src, float * dst, const int64_t n_rows } } - // stages within a warp: partner differs in the lane bits -#pragma unroll - for (int h = 1; h < warp_size; h *= 2) { -#pragma unroll - for (int j = 0; j < NE; j++) { - const float val = reg[j]; - const float val2 = __shfl_xor_sync(0xFFFFFFFF, val, h, warp_size); - reg[j] = (lane & h) == 0 ? val + val2 : val2 - val; - } - } - - // stages across warps: partner differs in the thread-index bits above the lane -#pragma unroll - for (int h = warp_size; h < NT; h *= 2) { -#pragma unroll - for (int j = 0; j < NE; j++) { - s[j * NT + tid] = reg[j]; - } - __syncthreads(); -#pragma unroll - for (int j = 0; j < NE; j++) { - const float val = reg[j]; - const float val2 = s[j * NT + (tid ^ h)]; - reg[j] = (tid & h) == 0 ? val + val2 : val2 - val; - } - __syncthreads(); - } - - // stages above the block width: partner is another register of the same thread -#pragma unroll - for (int h = NT; h < N; h *= 2) { - const int step = h / NT; -#pragma unroll - for (int j = 0; j < NE; j += 2 * step) { -#pragma unroll - for (int k = 0; k < step; k++) { - const float x = reg[j + k]; - const float y = reg[j + k + step]; - reg[j + k] = x + y; - reg[j + k + step] = x - y; - } - } - } + ggml_cuda_fwht_block_butterfly(reg, s, tid, lane); #pragma unroll for (int i = 0; i < NE; ++i) { diff --git a/ggml/src/ggml-cuda/fwht.cuh b/ggml/src/ggml-cuda/fwht.cuh index 62b2f288dab9..dad45aa810c6 100644 --- a/ggml/src/ggml-cuda/fwht.cuh +++ b/ggml/src/ggml-cuda/fwht.cuh @@ -1,5 +1,59 @@ #include "common.cuh" +// In-register block FWHT over N floats held as NE = N/NT per thread, element i*NT + tid in reg[i]. +// Shared by the standalone kernel (fwht_cuda_block) and the fused FWHT + q8_1 quantizer so both +// produce bit-identical transforms. `s` must hold N floats. +template +__device__ __forceinline__ void ggml_cuda_fwht_block_butterfly(float * reg, float * s, const int tid, const int lane) { + constexpr int warp_size = ggml_cuda_get_physical_warp_size(); + constexpr int NE = N / NT; + static_assert(NE >= 1 && N % NT == 0 && NT % warp_size == 0, "bad FWHT block shape"); + + // stages within a warp: partner differs in the lane bits +#pragma unroll + for (int h = 1; h < warp_size; h *= 2) { +#pragma unroll + for (int j = 0; j < NE; j++) { + const float val = reg[j]; + const float val2 = __shfl_xor_sync(0xFFFFFFFF, val, h, warp_size); + reg[j] = (lane & h) == 0 ? val + val2 : val2 - val; + } + } + + // stages across warps: partner differs in the thread-index bits above the lane +#pragma unroll + for (int h = warp_size; h < NT; h *= 2) { +#pragma unroll + for (int j = 0; j < NE; j++) { + s[j * NT + tid] = reg[j]; + } + __syncthreads(); +#pragma unroll + for (int j = 0; j < NE; j++) { + const float val = reg[j]; + const float val2 = s[j * NT + (tid ^ h)]; + reg[j] = (tid & h) == 0 ? val + val2 : val2 - val; + } + __syncthreads(); + } + + // stages above the block width: partner is another register of the same thread +#pragma unroll + for (int h = NT; h < N; h *= 2) { + const int step = h / NT; +#pragma unroll + for (int j = 0; j < NE; j += 2 * step) { +#pragma unroll + for (int k = 0; k < step; k++) { + const float x = reg[j + k]; + const float y = reg[j + k + step]; + reg[j + k] = x + y; + reg[j + k + step] = x - y; + } + } + } +} + // Returns whether the Fast Walsh-Hadamard transform could be used. bool ggml_cuda_op_fwht(ggml_backend_cuda_context & ctx, const ggml_tensor * src, ggml_tensor * dst); bool ggml_cuda_op_fwht_signed(ggml_backend_cuda_context & ctx, const ggml_tensor * src, diff --git a/ggml/src/ggml-cuda/gated_delta_net.cu b/ggml/src/ggml-cuda/gated_delta_net.cu index f0d08152acb0..40a639dd4601 100644 --- a/ggml/src/ggml-cuda/gated_delta_net.cu +++ b/ggml/src/ggml-cuda/gated_delta_net.cu @@ -49,7 +49,9 @@ gated_delta_net_cuda(const float * q, const uint3 rq3_magic, float scale, int64_t state_slot_stride, - int K) { + int K, + const int32_t * s_ids, + int64_t s_row_stride) { const uint32_t h_idx = blockIdx.x; const uint32_t sequence = blockIdx.y; // Each warp owns one or more columns, using warp-level primitives to reduce across rows. @@ -66,9 +68,15 @@ gated_delta_net_cuda(const float * q, float * attn_data = dst; + // Wait before any s_ids / state load. inp_s_copy is a host upload today, so this is a no-op + // on that path; ggml_cuda_try_gdn_gather_skip accepts any I32 ids tensor. + ggml_cuda_pdl_sync(); + // input state holds s0 only: [S_v, S_v, H, n_seqs] — seq stride is D = H * S_v * S_v. // output state layout (per-slot D * n_seqs) — same per-(seq,head) offset as before. - const int64_t state_in_offset = sequence * H * S_v * S_v + h_idx * S_v * S_v; + // fused gather: read this sequence's live state straight out of the cache row s_ids[sequence] + const int64_t state_in_offset = (s_ids ? (int64_t) s_ids[sequence] * s_row_stride : sequence * H * S_v * S_v) + + h_idx * S_v * S_v; const int64_t state_out_offset = (sequence * H + h_idx) * S_v * S_v; state += state_out_offset; curr_state += state_in_offset; @@ -80,7 +88,6 @@ gated_delta_net_cuda(const float * q, float s_shard[cols_per_warp][rows_per_lane]; // state is stored transposed: M[col][i] = S[i][col], row col is contiguous - ggml_cuda_pdl_sync(); #pragma unroll for (int c = 0; c < cols_per_warp; ++c) { #pragma unroll @@ -222,7 +229,8 @@ static void launch_gated_delta_net( int64_t sv1, int64_t sv2, int64_t sv3, int64_t sb1, int64_t sb2, int64_t sb3, int64_t neqk1, int64_t rq3, - float scale, int64_t state_slot_stride, int K, cudaStream_t stream) { + float scale, int64_t state_slot_stride, int K, + const int32_t * s_ids, int64_t s_row_stride, cudaStream_t stream) { //TODO: Add chunked kernel for even faster pre-fill const int warp_size = ggml_cuda_info().devices[ggml_cuda_get_device()].warp_size; const int cc = ggml_cuda_info().devices[ggml_cuda_get_device()].cc; @@ -241,26 +249,26 @@ static void launch_gated_delta_net( ggml_cuda_kernel_launch(gated_delta_net_cuda<16, KDA, keep_rs_t, RAW, G_PRECOMPUTED>, launch_params, q_d, k_d, v_d, g_d, b_d, rb_d, ra_d, s_d, dst_d, state_d, H, n_tokens, n_seqs, sq1, sq2, sq3, sv1, sv2, sv3, - sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, state_slot_stride, K); + sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, state_slot_stride, K, s_ids, s_row_stride); break; case 32: ggml_cuda_kernel_launch(gated_delta_net_cuda<32, KDA, keep_rs_t, RAW, G_PRECOMPUTED>, launch_params, q_d, k_d, v_d, g_d, b_d, rb_d, ra_d, s_d, dst_d, state_d, H, n_tokens, n_seqs, sq1, sq2, sq3, sv1, sv2, sv3, - sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, state_slot_stride, K); + sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, state_slot_stride, K, s_ids, s_row_stride); break; case 64: { ggml_cuda_kernel_launch(gated_delta_net_cuda<64, KDA, keep_rs_t, RAW, G_PRECOMPUTED>, launch_params, q_d, k_d, v_d, g_d, b_d, rb_d, ra_d, s_d, dst_d, state_d, H, n_tokens, n_seqs, sq1, sq2, sq3, sv1, sv2, sv3, - sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, state_slot_stride, K); + sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, state_slot_stride, K, s_ids, s_row_stride); break; } case 128: { ggml_cuda_kernel_launch(gated_delta_net_cuda<128, KDA, keep_rs_t, RAW, G_PRECOMPUTED>, launch_params, q_d, k_d, v_d, g_d, b_d, rb_d, ra_d, s_d, dst_d, state_d, H, n_tokens, n_seqs, sq1, sq2, sq3, sv1, sv2, sv3, - sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, state_slot_stride, K); + sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, state_slot_stride, K, s_ids, s_row_stride); break; } default: @@ -307,6 +315,15 @@ static void ggml_cuda_op_gated_delta_net_impl( const float * s_d = (const float *) src_state->data; float * dst_d = (float *) dst->data; + // fused state gather, registered for this node by this context's graph evaluator (ggml_cuda_try_gdn_gather_skip) + const int32_t * s_ids = nullptr; + int64_t s_row_stride = 0; + if (const ggml_cuda_gated_delta_net_gather * gather = ctx.gdn_gathers().find(dst)) { + s_d = gather->base; + s_ids = gather->ids; + s_row_stride = gather->row_stride; + } + GGML_ASSERT(ggml_is_contiguous_rows(src_q)); GGML_ASSERT(ggml_is_contiguous_rows(src_k)); GGML_ASSERT(ggml_is_contiguous_rows(src_v)); @@ -366,7 +383,7 @@ static void ggml_cuda_op_gated_delta_net_impl( #define GDN_LAUNCH(KDA_, KEEP_, RAW_, PRE_) \ launch_gated_delta_net(q_d, k_d, v_d, g_d, b_d, rb_d, ra_d, s_d, dst_d, state_d, \ S_v, H, n_tokens, n_seqs, sq1, sq2, sq3, sv1, sv2, sv3, \ - sb1, sb2, sb3, neqk1, rq3, scale, state_slot_stride, K, stream) + sb1, sb2, sb3, neqk1, rq3, scale, state_slot_stride, K, s_ids, s_row_stride, stream) if (kda) { if (keep_rs) { GDN_LAUNCH(true, true, false, false); } else { GDN_LAUNCH(true, false, false, false); } diff --git a/ggml/src/ggml-cuda/gated_delta_net.cuh b/ggml/src/ggml-cuda/gated_delta_net.cuh index f9bf43706789..bde78c42daf6 100644 --- a/ggml/src/ggml-cuda/gated_delta_net.cuh +++ b/ggml/src/ggml-cuda/gated_delta_net.cuh @@ -7,6 +7,8 @@ struct ggml_cuda_gated_delta_net_fused_cache { int64_t slot_stride; // between rollback slots (0 when K==1) }; +// The fused recurrent-state gather (ggml_cuda_gated_delta_net_gather, common.cuh) is looked up in +// ctx.gdn_gathers() by node pointer; the graph evaluator registers it when it skips the GET_ROWS. void ggml_cuda_op_gated_delta_net(ggml_backend_cuda_context & ctx, ggml_tensor * dst); // same op, but writes the snapshot(s) into the cache instead of dst (see ggml_cuda_try_gdn_cache_fusion) diff --git a/ggml/src/ggml-cuda/ggml-cuda.cu b/ggml/src/ggml-cuda/ggml-cuda.cu index 7536a5d2a094..41b3be6317d8 100644 --- a/ggml/src/ggml-cuda/ggml-cuda.cu +++ b/ggml/src/ggml-cuda/ggml-cuda.cu @@ -704,6 +704,9 @@ ggml_backend_cuda_context::~ggml_backend_cuda_context() { std::unique_lock lock(ggml_cuda_lock); ggml_cuda_lock_cv.wait(lock, []{ return ggml_cuda_lock_counter.load(std::memory_order_relaxed) == 0; }); + // pool blocks held across a graph go back to the pool before the pools are destroyed below + fwht_q8_context.reset(); + if (copy_event != nullptr) { CUDA_CHECK(cudaEventDestroy(copy_event)); } @@ -2747,6 +2750,238 @@ static bool ggml_cuda_should_fuse_rms_norm_mul_rope(const ggml_tensor * rms_norm return true; } +// A Hadamard rotation of an activation ([MUL signs, RESHAPE,] MUL_MAT with GGML_HINT_SRC0_IS_HADAMARD) +// whose every use is a PTQ1_0 mat-vec on the mmvq path does not need to exist in F32: each of those +// mat-vecs quantizes it to q8_1 straight away, so the transform kernel quantizes as it goes and writes +// the q8_1 rows (in the layout the mat-vecs expect) into its own output buffer; the mat-vecs then skip +// their quantize launch (ggml_cuda_fwht_q8). On Bonsai 2 that removes ~390 launches per decode step. +// The check walks every later node, so any other use of the output (or of a view of it) keeps the +// plain transform, and use counts guard against consumers outside this graph. Returns the number of +// nodes consumed (3 or 1) or 0 to fall through. GGML_CUDA_FWHT_FUSION=0 disables it. +static int ggml_cuda_try_fwht_q8(ggml_backend_cuda_context & ctx, const ggml_cgraph * cgraph, int node_idx) { + static const bool disabled = getenv("GGML_CUDA_FWHT_FUSION") != nullptr && + atoi(getenv("GGML_CUDA_FWHT_FUSION")) == 0; + if (disabled) { + return 0; + } + const int cc = ggml_cuda_info().devices[ctx.device].cc; + if (!GGML_CUDA_CC_IS_NVIDIA(cc) || cc < GGML_CUDA_CC_TURING) { + return 0; // the PTQ1_0 mmvq predicate below is only meaningful there + } + + const ggml_tensor * x = nullptr; + const ggml_tensor * signs = nullptr; + ggml_tensor * mm = nullptr; + int consumed = 0; + int i_mm = -1; + + if (ggml_can_fuse_subgraph(cgraph, node_idx, { GGML_OP_MUL, GGML_OP_RESHAPE, GGML_OP_MUL_MAT }, { node_idx + 2 })) { + const ggml_tensor * mul = cgraph->nodes[node_idx]; + const ggml_tensor * reshape = cgraph->nodes[node_idx + 1]; + mm = cgraph->nodes[node_idx + 2]; + x = mul->src[0]; + signs = mul->src[1]; + const bool pattern_ok = ggml_get_op_params_i32(mm, 1) == GGML_HINT_SRC0_IS_HADAMARD && + mm->src[1] == reshape && reshape->src[0] == mul && + signs->ne[1] == 1 && signs->ne[2] == 1 && signs->ne[3] == 1 && + signs->type == GGML_TYPE_F32 && mul->type == x->type && + ggml_is_contiguous(x) && ggml_is_contiguous(signs) && + signs->ne[0] == x->ne[0] && signs->ne[0] % mm->src[0]->ne[0] == 0; + if (!pattern_ok) { + return 0; + } + consumed = 3; + i_mm = node_idx + 2; + } else { + mm = cgraph->nodes[node_idx]; + if (mm->op != GGML_OP_MUL_MAT || ggml_get_op_params_i32(mm, 1) != GGML_HINT_SRC0_IS_HADAMARD) { + return 0; + } + x = mm->src[1]; + if (!ggml_is_contiguous(x) || !ggml_are_same_shape(x, mm)) { + return 0; + } + consumed = 1; + i_mm = node_idx; + } + + const int n = (int) mm->src[0]->ne[0]; + if (mm->type != GGML_TYPE_F32 || !ggml_is_contiguous(mm) || (mm->flags & GGML_TENSOR_FLAG_OUTPUT) || + (x->type != GGML_TYPE_F32 && x->type != GGML_TYPE_F16) || ggml_nelements(mm) != ggml_nelements(x)) { + return 0; + } + + // The fused kernel reads x (and signs) in the same launch that writes its output. The separate transform + // and quantize launches never needed those apart, and ggml-alloc does hand mm the block of x when x is a + // reshape view (no in-place reuse through a view parent, so the viewed tensor is released at the MUL and + // mm takes it two nodes later - the output-projection input on every Bonsai layer). Blocks storing q8_1 + // rows early would then overwrite f32 input other blocks have not loaded yet: sporadic garbage that on a + // hybrid model lands in the recurrent state and poisons the rest of the sequence. In that case the rows + // go to a pool block held until the next graph evaluation instead of mm's buffer. + auto overlaps = [](const ggml_tensor * a, const ggml_tensor * b) { + const char * a0 = (const char *) a->data; + const char * a1 = a0 + ggml_nbytes(a); + const char * b0 = (const char *) b->data; + const char * b1 = b0 + ggml_nbytes(b); + return a0 < b1 && b0 < a1; + }; + const bool out_aliases_in = overlaps(x, mm) || (signs && overlaps(signs, mm)); + + // every use of the transform output (directly or through a reshape view of the whole tensor) + // must be src1 of a PTQ1_0 MUL_MAT that ggml_cuda_mul_mat routes to ggml_cuda_mul_mat_vec_q. + // The consumers' src1 shape [K, ncols] is what gets quantized (x is only known to be + // contiguous with the same element count; in the unsigned pattern it is the [n, rows] reshape). + const ggml_tensor * aliases[8] = { mm }; + int n_aliases = 1; + int uses_expected = ggml_node_get_use_count(cgraph, i_mm); + int uses_found = 0; + int64_t K = 0; + int64_t ncols = 0; + int64_t ne0_padded = 0; + ggml_cuda_q8_1_layout layout = GGML_CUDA_Q8_1_AOS; + for (int j = i_mm + 1; j < cgraph->n_nodes; ++j) { + ggml_tensor * t = cgraph->nodes[j]; + for (int s = 0; s < GGML_MAX_SRC; ++s) { + const ggml_tensor * src = t->src[s]; + if (!src) { + continue; + } + bool is_alias = false; + for (int a = 0; a < n_aliases; ++a) { + if (src == aliases[a]) { + is_alias = true; + break; + } + } + if (!is_alias) { + continue; + } + uses_found++; + if (t->op == GGML_OP_RESHAPE && t->view_src == mm && t->data == mm->data && + ggml_nelements(t) == ggml_nelements(mm) && ggml_is_contiguous(t)) { + if (n_aliases == 8) { + return 0; + } + aliases[n_aliases++] = t; + uses_expected += ggml_node_get_use_count(cgraph, j); + continue; + } + if (t->op != GGML_OP_MUL_MAT || s != 1 || !t->src[0] || t->src[0]->type != GGML_TYPE_PTQ1_0 || + ggml_get_op_params_i32(t, 1) == GGML_HINT_SRC0_IS_HADAMARD || + t->type != GGML_TYPE_F32 || src->type != GGML_TYPE_F32 || !ggml_is_contiguous(src) || + src->ne[2] != 1 || src->ne[3] != 1 || + (K != 0 && (src->ne[0] != K || src->ne[1] != ncols)) || + !ggml_cuda_should_use_mmvq(GGML_TYPE_PTQ1_0, cc, src->ne[1])) { + return 0; + } + K = src->ne[0]; + ncols = src->ne[1]; + // ggml_cuda_mul_mat sends a padded compute-buffer view to cuBLAS instead + const ggml_tensor * w = t->src[0]; + if (ggml_backend_buffer_get_usage(w->buffer) == GGML_BACKEND_BUFFER_USAGE_COMPUTE && + ggml_nbytes(w) != ggml_backend_buffer_get_alloc_size(w->buffer, w) && w->view_src) { + return 0; + } + } + } + if (uses_found == 0 || uses_found != uses_expected || K == 0) { + return 0; + } + if (K % n != 0 || !ggml_cuda_fwht_quantize_supported(n, K) || (signs && signs->ne[0] != K)) { + return 0; + } + + // same (type, ncols, ids) -> layout and padding as ggml_cuda_mul_mat_vec_q + layout = ggml_cuda_q8_1_layout_host(GGML_TYPE_PTQ1_0, (int) ncols, false); + ne0_padded = GGML_PAD(K, MATRIX_ROW_PADDING); + if (layout == GGML_CUDA_Q8_1_SOA_ISUM) { + ne0_padded = GGML_PAD(ne0_padded, GGML_CUDA_PTQ1_K_PAD); + } + if (ne0_padded % n != 0) { + return 0; // the fused quantizer writes pad blocks in transform-width units + } + const size_t q8_bytes = (size_t) (ncols * ne0_padded) * sizeof(block_q8_1) / QK8_1; + void * out = mm->data; + if (out_aliases_in) { + auto blk = std::make_unique>(ctx.pool(), q8_bytes); + out = blk->get(); + ctx.fwht_q8().held.push_back(std::move(blk)); + } else if ((int64_t) q8_bytes > (int64_t) ggml_nbytes(mm)) { + return 0; // the quantized rows must fit in the F32 output allocation + } + + fwht_quantize_row_q8_1_cuda(x->data, x->type, signs ? (const float *) signs->data : nullptr, n, + out, layout, K, ne0_padded, ncols, ctx.stream()); + + ggml_cuda_fwht_q8 e; + e.layout = layout; + e.ne0 = ne0_padded; + e.ncols = ncols; + e.data = out; + for (int a = 0; a < n_aliases; ++a) { + ctx.fwht_q8().set(aliases[a], e); + } + return consumed; +} + +// GET_ROWS and let the kernel index the cache row directly. Single-sequence only: with several +// sequences a gathered row may alias a row another sequence writes in the same op. +// GGML_CUDA_GDN_GATHER_FUSION=0 disables it. The registration lives in the evaluating context +// (ctx.gdn_gathers()), so concurrent contexts or threads never see each other's skips. +static bool ggml_cuda_try_gdn_gather_skip(ggml_backend_cuda_context & ctx, const ggml_cgraph * cgraph, int node_idx) { + static const bool disabled = getenv("GGML_CUDA_GDN_GATHER_FUSION") != nullptr && + atoi(getenv("GGML_CUDA_GDN_GATHER_FUSION")) == 0; + if (disabled) { + return false; + } + const ggml_tensor * gr = cgraph->nodes[node_idx]; + if (gr->op != GGML_OP_GET_ROWS || gr->type != GGML_TYPE_F32 || (gr->flags & GGML_TENSOR_FLAG_OUTPUT) || + !ggml_is_contiguous(gr)) { + return false; + } + const ggml_tensor * cache = gr->src[0]; + const ggml_tensor * ids = gr->src[1]; + if (cache->type != GGML_TYPE_F32 || ids->type != GGML_TYPE_I32 || cache->data == nullptr || ids->data == nullptr || + cache->nb[0] != sizeof(float) || cache->nb[1] % sizeof(float) != 0 || !ggml_is_contiguous(ids) || + ids->ne[0] != 1 || ids->ne[1] != 1 || ids->ne[2] != 1 || ids->ne[3] != 1 || + gr->ne[1] != 1 || gr->ne[2] != 1 || gr->ne[3] != 1 || gr->ne[0] != cache->ne[0]) { + return false; + } + if (ggml_node_get_use_count(cgraph, node_idx) != 1) { + return false; + } + const ggml_tensor * cur = gr; + for (int j = node_idx + 1; j < cgraph->n_nodes; ++j) { + const ggml_tensor * n = cgraph->nodes[j]; + if (n->op == GGML_OP_GATED_DELTA_NET && n->src[5] == cur) { + const ggml_tensor * v = n->src[2]; + const int64_t D = v->ne[0] * v->ne[0] * v->ne[1]; + if (gr->ne[0] != D || v->ne[3] != 1 || ggml_nelements(cur) != D) { + return false; + } + ggml_cuda_gated_delta_net_gather gather; + gather.base = (const float *) cache->data; + gather.ids = (const int32_t *) ids->data; + gather.row_stride = (int64_t) (cache->nb[1] / sizeof(float)); + ctx.gdn_gathers().set(n, gather); + return true; + } + if (n->op == GGML_OP_RESHAPE && n->src[0] == cur) { + if (ggml_node_get_use_count(cgraph, j) != 1) { + return false; + } + cur = n; + continue; + } + for (int s = 0; s < GGML_MAX_SRC; ++s) { + if (n->src[s] == cur || (n->view_src != nullptr && n->view_src == gr)) { + return false; + } + } + } + return false; +} + // match gated_delta_net + the strided cpy that scatters its state snapshots into the cache // (slot i -> rollback group i, slot 0 newest), so the kernel can write them and skip the cpy. static int ggml_cuda_try_gdn_cache_fusion( @@ -4281,6 +4516,9 @@ static void ggml_cuda_graph_evaluate_and_capture(ggml_backend_cuda_context * cud stream_ctx.concurrent_events.clear(); } + cuda_ctx->gdn_gathers().reset(); + cuda_ctx->fwht_q8().reset(); + for (int i = 0; i < cgraph->n_nodes; i++) { ggml_tensor * node = cgraph->nodes[i]; if (is_concurrent_event_active) { @@ -4323,6 +4561,21 @@ static void ggml_cuda_graph_evaluate_and_capture(ggml_backend_cuda_context * cud continue; } + // recurrent-state gather folded into the GDN kernel (no temp, no kernel) + if (node->op == GGML_OP_GET_ROWS && !is_concurrent_event_active && + ggml_cuda_try_gdn_gather_skip(*cuda_ctx, cgraph, i)) { + continue; + } + + // Hadamard rotation that quantizes its own output; its PTQ1_0 mat-vecs skip quantize + if ((node->op == GGML_OP_MUL || node->op == GGML_OP_MUL_MAT) && !is_concurrent_event_active) { + const int consumed = ggml_cuda_try_fwht_q8(*cuda_ctx, cgraph, i); + if (consumed > 0) { + i += consumed - 1; + continue; + } + } + // The normalized pre-attention residual is consumed only by a // group of low-bit projections. Preserve residual + one scale per // row and let their shared Q8 quantizer apply the norm weight. diff --git a/ggml/src/ggml-cuda/mmvf.cu b/ggml/src/ggml-cuda/mmvf.cu index d7dbc8b99282..b12245e71f63 100644 --- a/ggml/src/ggml-cuda/mmvf.cu +++ b/ggml/src/ggml-cuda/mmvf.cu @@ -821,6 +821,9 @@ bool ggml_cuda_should_use_mmvf(enum ggml_type type, int cc, const int64_t * src0 if (GGML_CUDA_CC_IS_NVIDIA(cc)) { const bool src0_small = (src0_ne[1] <= 512 || src0_ne[2]*src0_ne[3] == 1); if (ampere_mma_available(cc)) { + if (ggml_cuda_batch_invariant()) { + return src0_small && ne11 <= MMVF_MAX_BATCH_SIZE; + } return src0_small && ne11 == 1; } if (cc >= GGML_CUDA_CC_ADA_LOVELACE) { @@ -847,6 +850,11 @@ bool ggml_cuda_should_use_mmvf(enum ggml_type type, int cc, const int64_t * src0 if (GGML_CUDA_CC_IS_NVIDIA(cc)) { const bool src0_small = (src0_ne[1] <= 512 || src0_ne[2]*src0_ne[3] == 1); if (ampere_mma_available(cc)) { + // a few dozen rows (the qwen35 gated-delta-net gate projections) run faster as a + // mat-vec than through the tensor-core path at 2 to 8 columns: 3.4 vs 10.5 us on an RTX 3060 + if (ggml_cuda_batch_invariant() || src0_ne[1] <= 64) { + return src0_small && ne11 <= MMVF_MAX_BATCH_SIZE; + } return src0_small && ne11 == 1; } if (cc >= GGML_CUDA_CC_ADA_LOVELACE) { diff --git a/ggml/src/ggml-cuda/mmvq-ptq1_0.cuh b/ggml/src/ggml-cuda/mmvq-ptq1_0.cuh new file mode 100644 index 000000000000..4755d7a86d00 --- /dev/null +++ b/ggml/src/ggml-cuda/mmvq-ptq1_0.cuh @@ -0,0 +1,551 @@ +// PTQ1_0 mat-vec inner loop on a planar-transposed Q8_1 activation layout. +// +// Why: the stock mmvq path hands each thread one 128-weight PTQ1_0 block and +// walks the activations as 32 scattered 4-byte loads per column out of +// 36-byte block_q8_1 structs. The weight decode is shared across columns but +// the activation traffic is not, so every extra column costs a full pass +// (measured 2.2x at 2 columns, 3.3x at 3 on an RTX 3060). Here the activations +// are stored so that the 128 quants a thread needs are 8 aligned 16-byte +// pieces, one per plane, and the 4 (d, s) scales are one more 16-byte piece. +// Adjacent threads read adjacent 16-byte pieces of the same plane, so a warp +// load touches 4 cache lines instead of 32, and a thread reuses each piece +// for every row it owns. GGML_CUDA_BATCH_INVARIANT=1 keeps the older +// warp-reduce epilogue so a column verified in a 2-4 wide MTP batch matches +// the same column decoded alone. The faster four-accumulator epilogue is the +// default; it can flip a late near-tie (5080: "coastal waters" vs "coastal areas"). +// +// PT layout, per activation column (all sizes for the padded row length): +// plane t (t = 0..7): nblk * 16 bytes, byte b of block kb is the quant of +// element kb*128 + t*16 + b +// plane 8: nblk * 16 bytes, block kb holds 4 half2 (d, s), one +// per 32-element sub-block +// Column stride is 9 * nblk * 16 = padded_row * 9/8 bytes, exactly the +// block_q8_1 stride (padded_row/32 blocks of 36 bytes), so every stride the +// mmvq launcher computes in block_q8_1 units stays valid. +#pragma once + +#include "common.cuh" +#include "unary.cuh" +#include "vecdotq.cuh" + +#define PTQ1_0_PT_PLANES 9 + +// dedicated 2D kernel geometry, see mul_mat_vec_ptq1_0_pt below +#define PTQ1_0_PT_THREADS 128 +#define PTQ1_0_PT_MAX_ROWS 16 +#define PTQ1_0_PT_MAX_COLS 8 // equals MMVQ_MAX_BATCH_SIZE, checked in mmvq.cu +#define PTQ1_0_PT_SMEM_FLOATS 4096 // 16 KiB target when choosing rows per CTA; the launch may request more for one item + +// the PT path is CUDA only; HIP keeps the block_q8_1 layout and the old vec_dot +static constexpr __host__ __device__ bool ptq1_0_pt_enabled() { +#if defined(GGML_USE_HIP) + return false; +#else + return true; +#endif +} + +// rows of the weight matrix one thread handles per K block: more rows reuse +// each activation piece more often, fewer rows keep the register count down +static constexpr __host__ __device__ int ptq1_0_pt_rows_per_block(const int ncols_dst) { + return ncols_dst == 1 ? 1 : 2; +} + +// number of 128-element blocks in a row padded to MATRIX_ROW_PADDING +static __host__ __device__ __forceinline__ int ptq1_0_pt_nblk(const int ncols_x) { + return ((ncols_x + MATRIX_ROW_PADDING - 1) / MATRIX_ROW_PADDING) * (MATRIX_ROW_PADDING / QK_PTQ1_0); +} + +static __device__ __forceinline__ int int4_at(const int4 & v, const int k) { + switch (k & 3) { + case 0: return v.x; + case 1: return v.y; + case 2: return v.z; + default: return v.w; + } +} + +// one base-3 digit step on four bytes held as two 16-bit-lane words: +// returns the raw digits (0, 1, 2) as four unsigned bytes, advances the remainders. +// The digit bias is not applied here: the dot product uses the digits as-is with a mixed-sign +// dp4a and subtracts the exact integer activation sum of the 32-block once (see ptq1_0_pt_block_dot), +// which is the same integer as the biased sum, two ALU ops per 4 weights cheaper. +static __device__ __forceinline__ uint32_t ptq1_0_trit_step(uint32_t & vlo, uint32_t & vhi) { + const uint32_t wlo = vlo * 3; + const uint32_t whi = vhi * 3; + vlo = wlo & 0x00FF00FF; + vhi = whi & 0x00FF00FF; + return __byte_perm(wlo, whi, 0x7531); +} + +// Dot products of nrows PTQ1_0 blocks with the same block index of ncols +// activation columns. bq[i] points at the weight block of row i, ycol[j] at +// the PT column base of column j, kbx is the block index along K. +// +// The integer sum of each 32-element sub-block k is folded into the fp32 +// accumulator as soon as the sub-block is complete, in the order k = 0..3, +// which is the expression acc = sum_k d8_k * sumi_k of the block_q8_1 kernel. +// sumi_k is accumulated from the raw digits {0,1,2} and corrected by the exact +// integer activation sum stored in the layout: sum((q-1)*a) = sum(q*a) - sum(a), +// exact in int32, so the result is bit-identical to the biased-weight form. +template +static __device__ __forceinline__ void ptq1_0_pt_block_dot( + const block_ptq1_0 * const (&bq)[nrows], + const char * const (&ycol)[ncols], + const int kbx, const int nblk, + float (&result)[ncols][nrows]) { + int4 dsraw[ncols]; +#pragma unroll + for (int j = 0; j < ncols; ++j) { + dsraw[j] = *((const int4 *) ycol[j] + 8*nblk + kbx); + } + + int sumi[ncols][nrows]; + float acc[ncols][nrows]; +#pragma unroll + for (int j = 0; j < ncols; ++j) { +#pragma unroll + for (int i = 0; i < nrows; ++i) { + sumi[j][i] = 0; + acc[j][i] = 0.0f; + } + } + + auto fold = [&](const int k) { +#pragma unroll + for (int j = 0; j < ncols; ++j) { + const half2 ds = ((const half2 *) &dsraw[j])[k]; + const float d8 = __low2float(ds); + const int isum = __half_as_short(__high2half(ds)); // exact sum of the 32 quantized activations +#pragma unroll + for (int i = 0; i < nrows; ++i) { + acc[j][i] = __fmaf_rn(d8, (float) (sumi[j][i] - isum), acc[j][i]); // one FFMA in every instantiation + sumi[j][i] = 0; + } + } + }; + + // qs[0..15]: four groups of four bytes, five trits each: element 16*t + 4*g + b, + // plane t holds words 4*t .. 4*t+3 + uint32_t vlo[nrows][4]; + uint32_t vhi[nrows][4]; +#pragma unroll + for (int i = 0; i < nrows; ++i) { +#pragma unroll + for (int g = 0; g < 4; ++g) { + const uint32_t packed = get_int_b4(bq[i]->qs, g); + vlo[i][g] = __byte_perm(packed, 0, 0x4140); + vhi[i][g] = __byte_perm(packed, 0, 0x4342); + } + } +#pragma unroll + for (int t = 0; t < 5; ++t) { + int4 u[ncols]; +#pragma unroll + for (int j = 0; j < ncols; ++j) { + u[j] = *((const int4 *) ycol[j] + t*nblk + kbx); + } +#pragma unroll + for (int i = 0; i < nrows; ++i) { +#pragma unroll + for (int g = 0; g < 4; ++g) { + const uint32_t q = ptq1_0_trit_step(vlo[i][g], vhi[i][g]); +#pragma unroll + for (int j = 0; j < ncols; ++j) { + sumi[j][i] = ggml_cuda_dp4a_us(q, int4_at(u[j], g), sumi[j][i]); + } + } + } + if (t == 1) { + fold(0); // elements 0..31 done + } + if (t == 3) { + fold(1); // elements 32..63 done + } + } + + // qs[16..23]: two groups of four bytes, five trits each: element 80 + 8*t + 4*g + b, + // words 20..29 live in planes 5, 6 and the lower half of 7 + uint32_t vlo2[nrows][2]; + uint32_t vhi2[nrows][2]; +#pragma unroll + for (int i = 0; i < nrows; ++i) { +#pragma unroll + for (int g = 0; g < 2; ++g) { + const uint32_t packed = get_int_b4(bq[i]->qs + 16, g); + vlo2[i][g] = __byte_perm(packed, 0, 0x4140); + vhi2[i][g] = __byte_perm(packed, 0, 0x4342); + } + } + int4 u2[ncols]; +#pragma unroll + for (int t = 0; t < 5; ++t) { + if (t % 2 == 0) { +#pragma unroll + for (int j = 0; j < ncols; ++j) { + u2[j] = *((const int4 *) ycol[j] + (5 + t/2)*nblk + kbx); + } + } +#pragma unroll + for (int i = 0; i < nrows; ++i) { +#pragma unroll + for (int g = 0; g < 2; ++g) { + const uint32_t q = ptq1_0_trit_step(vlo2[i][g], vhi2[i][g]); + const int w = 20 + 2*t + g; // word index within the 128-element block +#pragma unroll + for (int j = 0; j < ncols; ++j) { + sumi[j][i] = ggml_cuda_dp4a_us(q, int4_at(u2[j], w & 3), sumi[j][i]); + } + } + } + if (t == 1) { + fold(2); // elements 64..95 done (words 16..23) + } + } + + // qh: two bytes, four trits each, interleaved: element 120 + 2*t + h -> words 30, 31, + // the upper half of plane 7 that u2 still holds +#pragma unroll + for (int i = 0; i < nrows; ++i) { + uint32_t v = (uint32_t) bq[i]->qh[0] | ((uint32_t) bq[i]->qh[1] << 16); +#pragma unroll + for (int t = 0; t < 4; t += 2) { + const uint32_t w0 = v * 3; + v = w0 & 0x00FF00FF; + const uint32_t w1 = v * 3; + v = w1 & 0x00FF00FF; + const uint32_t q = __byte_perm(w0, w1, 0x7531); +#pragma unroll + for (int j = 0; j < ncols; ++j) { + sumi[j][i] = ggml_cuda_dp4a_us(q, int4_at(u2[j], 2 + t/2), sumi[j][i]); + } + } + } + fold(3); // elements 96..127 done + +#pragma unroll + for (int j = 0; j < ncols; ++j) { +#pragma unroll + for (int i = 0; i < nrows; ++i) { + result[j][i] = __fmul_rn((float) bq[i]->d, acc[j][i]); + } + } +} + +// --------------------------------------------------------------------------- +// Dedicated PTQ1_0 mat-vec for plain 2D MUL_MAT (no batch dims, no expert ids). +// +// The generic mmvq kernel gives every thread of a 128-thread block one +// 128-weight K block of the same row, so a K = 5120 projection (40 blocks per +// row) keeps 40 of 128 threads busy and K = 17408 (136 blocks) keeps 53%. Here +// the work items are (row group, K block) pairs of `rows_per_cta` rows +// flattened into one index space, with rows_per_cta chosen on the host so that +// the items fill whole 128-thread iterations where possible. A thread handles +// ROWS adjacent rows per item so that each activation piece it loads serves +// ROWS rows. Each thread writes one fp32 partial per (row, column, K block) to +// shared memory and one warp per (row, column) sums them in a fixed order +// (lane-strided sequential, then a butterfly). That order depends only on the +// weight shape, so the result for a column is the same bits for every column +// count. +// --------------------------------------------------------------------------- + +// rows_per_cta: fill whole 128-thread iterations where possible, within the shared memory budget +static __host__ int ptq1_0_pt_rows_per_cta(const int blocks_per_row, const int ncols_dst, const int nrows_x, const int rows_per_item) { + int rmax = PTQ1_0_PT_SMEM_FLOATS / (ncols_dst * (blocks_per_row + 1)); + rmax = rmax < rows_per_item ? rows_per_item : (rmax > PTQ1_0_PT_MAX_ROWS ? PTQ1_0_PT_MAX_ROWS : rmax); + rmax -= rmax % rows_per_item; + int best = rows_per_item; + double best_util = 0.0; + for (int r = rows_per_item; r <= rmax; r += rows_per_item) { + const int items = (r / rows_per_item) * blocks_per_row; + const int iters = (items + PTQ1_0_PT_THREADS - 1) / PTQ1_0_PT_THREADS; + const double util = (double) items / (double) (iters * PTQ1_0_PT_THREADS); + if (util > best_util + 1e-9) { + best_util = util; + best = r; + } + if (util > 0.999) { + break; + } + } + GGML_UNUSED(nrows_x); + return best; +} + +#ifndef PTQ1_0_PT_MINB_34 +#define PTQ1_0_PT_MINB_34 3 +#endif +#ifndef PTQ1_0_PT_ROWS_34 +#define PTQ1_0_PT_ROWS_34 4 +#endif + +// Same row-per-item rule the launcher instantiates. 4 up to 4 columns, 2 beyond (8 spills). +static constexpr __host__ __device__ int ptq1_0_pt_rows_per_item(const int ncols_dst) { + return ncols_dst <= 2 ? 4 : (ncols_dst <= 4 ? PTQ1_0_PT_ROWS_34 : 2); +} + +// Bytes the launch requests: one fp32 partial per (column, row, K block + 1 pad) for the CTA, +// twice that with gate fusion. The +1 is the odd stride the epilogue uses to avoid bank conflicts. +// Guard and launcher both call this so a shape that would exceed smpb falls back to the generic kernel. +static __host__ size_t ptq1_0_pt_smem_bytes(const int blocks_per_row, const int ncols_dst, const int nrows_x, const bool has_gate) { + const int rows_per_cta = ptq1_0_pt_rows_per_cta(blocks_per_row, ncols_dst, nrows_x, ptq1_0_pt_rows_per_item(ncols_dst)); + return (size_t) ncols_dst * rows_per_cta * (blocks_per_row + 1) * sizeof(float) * (has_gate ? 2 : 1); +} + +template +__launch_bounds__(PTQ1_0_PT_THREADS, (ncols <= 2 ? 4 : (ncols <= 4 ? PTQ1_0_PT_MINB_34 : 2))) +static __global__ void mul_mat_vec_ptq1_0_pt( + const void * vx_, const void * vy_, const ggml_cuda_mm_fusion_args_device fusion, + float * dst_, + const int ncols_x, const int nrows_x, const int stride_row_x, const int stride_col_y, const int stride_col_dst, + const int rows_per_cta, const uint3 bpr_fd, const uint3 rpc_fd, const bool invariant) { + // GGML_CUDA_RESTRICT stays off the formal parameters: cudafe's host stub drops __restrict + // from the explicit specialization and MSVC/GCC then reject it (C2912 / "does not match + // any template declaration") when compiling sm_90/sm_120. Same pattern as mul_mat_vec_q. + const void * GGML_CUDA_RESTRICT vx = vx_; + const void * GGML_CUDA_RESTRICT vy = vy_; + float * GGML_CUDA_RESTRICT dst = dst_; + // launched through ggml_cuda_kernel_launch, which opts into PDL on Hopper and newer: wait for the kernel + // that wrote vy (the q8_1 activation quantization) before reading it + ggml_cuda_pdl_sync(); + extern __shared__ float partials[]; // [ncols][rows_per_cta][bprp], then partials_gate + const int bpr = ncols_x / QK_PTQ1_0; // K blocks per row + const int bprp = bpr + 1; // partials row stride: odd, so the per-pair epilogue reads are bank-conflict-free + [[maybe_unused]] float * partials_gate = partials + ncols*rows_per_cta*bprp; + + const int nblk = ptq1_0_pt_nblk(ncols_x); // plane stride of the PT layout + const int row0 = rows_per_cta * blockIdx.x; + const int tid = threadIdx.x; + + // rows this CTA really owns (the last CTA may be short); clamped item rows read the last real row + const int n_rows_cta = min(rows_per_cta, nrows_x - row0); + + const char * ycol[ncols]; +#pragma unroll + for (int j = 0; j < ncols; ++j) { + ycol[j] = (const char *) ((const block_q8_1 *) vy + j*stride_col_y); + } + + const int n_items = (rows_per_cta / ROWS) * bpr; + for (int idx = tid; idx < n_items; idx += PTQ1_0_PT_THREADS) { + const int rg = fastdiv((uint32_t) idx, bpr_fd); // row group within the CTA + const int kbx = idx - rg*bpr; + + const block_ptq1_0 * bq[ROWS]; +#pragma unroll + for (int i = 0; i < ROWS; ++i) { + int r = rg*ROWS + i; + r = r < n_rows_cta ? r : n_rows_cta - 1; // clamp the tail, that result is not written + bq[i] = (const block_ptq1_0 *) vx + (int64_t) (row0 + r)*stride_row_x + kbx; + } + float dots[ncols][ROWS]; + ptq1_0_pt_block_dot(bq, ycol, kbx, nblk, dots); +#pragma unroll + for (int j = 0; j < ncols; ++j) { +#pragma unroll + for (int i = 0; i < ROWS; ++i) { + partials[(j*rows_per_cta + rg*ROWS + i)*bprp + kbx] = dots[j][i]; + } + } + if constexpr (has_gate) { + const block_ptq1_0 * bg[ROWS]; +#pragma unroll + for (int i = 0; i < ROWS; ++i) { + int r = rg*ROWS + i; + r = r < n_rows_cta ? r : n_rows_cta - 1; + bg[i] = (const block_ptq1_0 *) fusion.gate + (int64_t) (row0 + r)*stride_row_x + kbx; + } + ptq1_0_pt_block_dot(bg, ycol, kbx, nblk, dots); +#pragma unroll + for (int j = 0; j < ncols; ++j) { +#pragma unroll + for (int i = 0; i < ROWS; ++i) { + partials_gate[(j*rows_per_cta + rg*ROWS + i)*bprp + kbx] = dots[j][i]; + } + } + } + } + + __syncthreads(); + + if (invariant) { + // Warp-per-(row, column) lane-strided sum + butterfly. Same bits at 1 column and at + // 2-4 columns, which is what BATCH_INVARIANT needs for MTP-on == MTP-off. The four- + // accumulator path below is faster and can diverge (5080 bisect: 5300cd1). + const int warp = tid / WARP_SIZE; + const int lane = tid % WARP_SIZE; + for (int w = warp; w < rows_per_cta*ncols; w += PTQ1_0_PT_THREADS / WARP_SIZE) { + const int j = w / rows_per_cta; + const int r = w - j*rows_per_cta; + const int row = row0 + r; + + float sum = 0.0f; + [[maybe_unused]] float sum_gate = 0.0f; + for (int kbx = lane; kbx < bpr; kbx += WARP_SIZE) { + sum += partials[(j*rows_per_cta + r)*bprp + kbx]; + if constexpr (has_gate) { + sum_gate += partials_gate[(j*rows_per_cta + r)*bprp + kbx]; + } + } + sum = warp_reduce_sum(sum); + if constexpr (has_gate) { + sum_gate = warp_reduce_sum(sum_gate); + } + + if (lane == 0 && row < nrows_x) { + float result = sum; + if constexpr (has_fusion) { + if (fusion.x_bias) { + result += ((const float *) fusion.x_bias)[j*stride_col_dst + row]; + } + if constexpr (has_gate) { + float gate_value = sum_gate; + if (fusion.gate_bias) { + gate_value += ((const float *) fusion.gate_bias)[j*stride_col_dst + row]; + } + switch (fusion.glu_op) { + case GGML_GLU_OP_SWIGLU: + result *= ggml_cuda_op_silu_single(gate_value); + break; + case GGML_GLU_OP_GEGLU: + result *= ggml_cuda_op_gelu_single(gate_value); + break; + case GGML_GLU_OP_SWIGLU_OAI: + result = ggml_cuda_op_swiglu_oai_single(gate_value, result); + break; + default: + result = result * gate_value; + break; + } + } + } + dst[j*stride_col_dst + row] = result; + } + } + return; + } + + // One thread per (row, column): four interleaved accumulators (k mod 4), then (s0+s1)+(s2+s3). + // Faster than the warp path (~3% decode). Association depends on column count enough to break + // MTP-on vs MTP-off identity. GGML_CUDA_BATCH_INVARIANT takes the warp path above. + for (int p = tid; p < rows_per_cta*ncols; p += PTQ1_0_PT_THREADS) { + const int j = fastdiv((uint32_t) p, rpc_fd); // column + const int r = p - j*rows_per_cta; // row within the CTA + const int row = row0 + r; + const float * src = partials + (j*rows_per_cta + r)*bprp; + + float s0 = 0.0f, s1 = 0.0f, s2 = 0.0f, s3 = 0.0f; + [[maybe_unused]] float g0 = 0.0f, g1 = 0.0f, g2 = 0.0f, g3 = 0.0f; + int kbx = 0; + for (; kbx + 4 <= bpr; kbx += 4) { + s0 += src[kbx + 0]; s1 += src[kbx + 1]; s2 += src[kbx + 2]; s3 += src[kbx + 3]; + if constexpr (has_gate) { + const float * sg = partials_gate + (j*rows_per_cta + r)*bprp; + g0 += sg[kbx + 0]; g1 += sg[kbx + 1]; g2 += sg[kbx + 2]; g3 += sg[kbx + 3]; + } + } + for (; kbx < bpr; ++kbx) { + s0 += src[kbx]; + if constexpr (has_gate) { + g0 += partials_gate[(j*rows_per_cta + r)*bprp + kbx]; + } + } + const float sum = (s0 + s1) + (s2 + s3); + [[maybe_unused]] const float sum_gate = (g0 + g1) + (g2 + g3); + + if (row < nrows_x) { + float result = sum; + if constexpr (has_fusion) { + if (fusion.x_bias) { + result += ((const float *) fusion.x_bias)[j*stride_col_dst + row]; + } + if constexpr (has_gate) { + float gate_value = sum_gate; + if (fusion.gate_bias) { + gate_value += ((const float *) fusion.gate_bias)[j*stride_col_dst + row]; + } + switch (fusion.glu_op) { + case GGML_GLU_OP_SWIGLU: + result *= ggml_cuda_op_silu_single(gate_value); + break; + case GGML_GLU_OP_GEGLU: + result *= ggml_cuda_op_gelu_single(gate_value); + break; + case GGML_GLU_OP_SWIGLU_OAI: + result = ggml_cuda_op_swiglu_oai_single(gate_value, result); + break; + default: + result = result * gate_value; + break; + } + } + } + dst[j*stride_col_dst + row] = result; + } + } +} + +template +static void mul_mat_vec_ptq1_0_pt_launch( + const void * vx, const void * vy, const ggml_cuda_mm_fusion_args_device & fusion, float * dst, + const int ncols_x, const int nrows_x, const int stride_row_x, const int stride_col_y, const int stride_col_dst, + cudaStream_t stream) { + constexpr int ROWS = ptq1_0_pt_rows_per_item(ncols); + const int bpr = ncols_x / QK_PTQ1_0; + const bool has_fusion = fusion.gate != nullptr || fusion.x_bias != nullptr || fusion.gate_bias != nullptr; + const bool has_gate = fusion.gate != nullptr; + + const int rows_per_cta = ptq1_0_pt_rows_per_cta(bpr, ncols, nrows_x, ROWS); + const uint3 bpr_fd = init_fastdiv_values((uint32_t) bpr); + const uint3 rpc_fd = init_fastdiv_values((uint32_t) rows_per_cta); + const dim3 block_nums((nrows_x + rows_per_cta - 1) / rows_per_cta, 1, 1); + const dim3 block_dims(PTQ1_0_PT_THREADS, 1, 1); + + const size_t smem = ptq1_0_pt_smem_bytes(bpr, ncols, nrows_x, has_gate); + const ggml_cuda_kernel_launch_params lp = ggml_cuda_kernel_launch_params(block_nums, block_dims, smem, stream); + +#define PTQ1_0_PT_LAUNCH(FUS, GATE) \ + ggml_cuda_kernel_launch(mul_mat_vec_ptq1_0_pt, lp, \ + vx, vy, fusion, dst, ncols_x, nrows_x, stride_row_x, stride_col_y, stride_col_dst, rows_per_cta, bpr_fd, rpc_fd, \ + ggml_cuda_batch_invariant()) + + if (has_fusion) { + GGML_ASSERT(ncols == 1 && "fusion only supported for ncols_dst=1"); + if (has_gate) { + PTQ1_0_PT_LAUNCH(true, true); + } else { + PTQ1_0_PT_LAUNCH(true, false); + } + return; + } + PTQ1_0_PT_LAUNCH(false, false); +#undef PTQ1_0_PT_LAUNCH +} + +// true when the dedicated kernel handles this call (plain 2D, K a multiple of 128, up to 8 columns) +static bool mul_mat_vec_ptq1_0_pt_switch( + const void * vx, const void * vy, const ggml_cuda_mm_fusion_args_device & fusion, float * dst, + const int ncols_x, const int nrows_x, const int ncols_dst, + const int stride_row_x, const int stride_col_y, const int stride_col_dst, + const int nchannels_dst, const int nsamples_dst, cudaStream_t stream) { + if (!ptq1_0_pt_enabled() || nchannels_dst != 1 || nsamples_dst != 1 || ncols_x % QK_PTQ1_0 != 0 || + ncols_dst < 1 || ncols_dst > PTQ1_0_PT_MAX_COLS) { + return false; + } + const size_t smem = ptq1_0_pt_smem_bytes(ncols_x / QK_PTQ1_0, ncols_dst, nrows_x, fusion.gate != nullptr); + if (smem > ggml_cuda_info().devices[ggml_cuda_get_device()].smpb) { + return false; + } + switch (ncols_dst) { + case 1: mul_mat_vec_ptq1_0_pt_launch<1>(vx, vy, fusion, dst, ncols_x, nrows_x, stride_row_x, stride_col_y, stride_col_dst, stream); break; + case 2: mul_mat_vec_ptq1_0_pt_launch<2>(vx, vy, fusion, dst, ncols_x, nrows_x, stride_row_x, stride_col_y, stride_col_dst, stream); break; + case 3: mul_mat_vec_ptq1_0_pt_launch<3>(vx, vy, fusion, dst, ncols_x, nrows_x, stride_row_x, stride_col_y, stride_col_dst, stream); break; + case 4: mul_mat_vec_ptq1_0_pt_launch<4>(vx, vy, fusion, dst, ncols_x, nrows_x, stride_row_x, stride_col_y, stride_col_dst, stream); break; + case 5: mul_mat_vec_ptq1_0_pt_launch<5>(vx, vy, fusion, dst, ncols_x, nrows_x, stride_row_x, stride_col_y, stride_col_dst, stream); break; + case 6: mul_mat_vec_ptq1_0_pt_launch<6>(vx, vy, fusion, dst, ncols_x, nrows_x, stride_row_x, stride_col_y, stride_col_dst, stream); break; + case 7: mul_mat_vec_ptq1_0_pt_launch<7>(vx, vy, fusion, dst, ncols_x, nrows_x, stride_row_x, stride_col_y, stride_col_dst, stream); break; + case 8: mul_mat_vec_ptq1_0_pt_launch<8>(vx, vy, fusion, dst, ncols_x, nrows_x, stride_row_x, stride_col_y, stride_col_dst, stream); break; + default: return false; + } + return true; +} diff --git a/ggml/src/ggml-cuda/mmvq.cu b/ggml/src/ggml-cuda/mmvq.cu index bf51b61e17b1..b8279e40abb5 100644 --- a/ggml/src/ggml-cuda/mmvq.cu +++ b/ggml/src/ggml-cuda/mmvq.cu @@ -1,4 +1,5 @@ #include "mmvq.cuh" +#include "mmvq-ptq1_0.cuh" #include "quantize.cuh" #include "unary.cuh" #include "vecdotq.cuh" @@ -296,7 +297,10 @@ bool ggml_cuda_should_use_mmvq(enum ggml_type type, int cc, int64_t ne11) { } #if !defined(GGML_USE_HIP) if (type == GGML_TYPE_PTQ1_0 && GGML_CUDA_CC_IS_NVIDIA(cc) && cc >= GGML_CUDA_CC_TURING) { - return ne11 <= 7; + // The PT mat-vec path shares the weight decode across columns; with the branch-free PTQ1_0 + // MMQ tile loader the tile path overtakes it at 5+ columns on Ada (RTX 4070, Bonsai 2 27B: + // mat-vec 155 t/s vs MMQ 244 t/s at n=8). + return ne11 <= 4; } #endif // k-quants cost more to decode and mvq redoes that per column, so MMQ wins sooner. @@ -404,6 +408,14 @@ static constexpr __device__ int get_mmvq_mmid_max_batch_for_device() { } static constexpr __host__ __device__ int calc_nwarps(ggml_type type, int ncols_dst, mmvq_parameter_table_id table_id, bool small_k = false, bool halve_iters = false) { + if (ptq1_0_pt_enabled() && type == GGML_TYPE_PTQ1_0 && ncols_dst > 1 && + (table_id == MMVQ_PARAMETERS_GENERIC || table_id == MMVQ_PARAMETERS_TURING)) { + // one K partition for every column count, so the fp32 sums of a column + // do not depend on how many columns share the launch (batch invariance). + // ncols_dst == 1 keeps the stock table: the SoA path relies on rows_per_block == nwarps + // for its warp-per-row small-K geometry. + return ncols_dst <= MMVQ_MAX_BATCH_SIZE ? 4 : 1; + } if (table_id == MMVQ_PARAMETERS_GENERIC) { switch (ncols_dst) { case 1: @@ -534,7 +546,11 @@ static constexpr __host__ __device__ int calc_nwarps(ggml_type type, int ncols_d return 1; } -static constexpr __host__ __device__ int calc_rows_per_block(int ncols_dst, int table_id, bool small_k = false, int nwarps = 1) { +static constexpr __host__ __device__ int calc_rows_per_block(ggml_type type, int ncols_dst, int table_id, bool small_k = false, int nwarps = 1) { + if (ptq1_0_pt_enabled() && type == GGML_TYPE_PTQ1_0 && ncols_dst > 1 && + (table_id == MMVQ_PARAMETERS_GENERIC || table_id == MMVQ_PARAMETERS_TURING)) { + return ptq1_0_pt_rows_per_block(ncols_dst); + } if (table_id == MMVQ_PARAMETERS_GENERIC || table_id == MMVQ_PARAMETERS_GCN || table_id == MMVQ_PARAMETERS_TURING || table_id == MMVQ_PARAMETERS_GB10) { switch (ncols_dst) { case 1: @@ -554,7 +570,9 @@ static constexpr __host__ __device__ int calc_rows_per_block(int ncols_dst, int return 1; } -template +// y_soa: PTQ1_0 one-column activations are in the warp-transposed exact-isum layout (see +// ggml_cuda_q8_1_layout_for). Only ever instantiated true for type == PTQ1_0 && ncols_dst == 1. +template __launch_bounds__(calc_nwarps(type, ncols_dst, get_device_table_id(), small_k, halve_iters)*ggml_cuda_get_physical_warp_size(), 1) static __global__ void mul_mat_vec_q( const void * vx_ptr, const void * vy_ptr, const int32_t * ids_ptr, const ggml_cuda_mm_fusion_args_device fusion, float * dst_ptr, @@ -573,7 +591,7 @@ static __global__ void mul_mat_vec_q( constexpr int vdr = get_vdr_mmvq(type); constexpr mmvq_parameter_table_id table_id = get_device_table_id(); constexpr int nwarps = calc_nwarps(type, ncols_dst, table_id, small_k, halve_iters); - constexpr int rows_per_cuda_block = calc_rows_per_block(ncols_dst, table_id, small_k, nwarps); + constexpr int rows_per_cuda_block = calc_rows_per_block(type, ncols_dst, table_id, small_k, nwarps); constexpr int warp_size = ggml_cuda_get_physical_warp_size(); constexpr vec_dot_q_cuda_t vec_dot_q_cuda = get_vec_dot_q_cuda(type); @@ -667,6 +685,91 @@ static __global__ void mul_mat_vec_q( const block_q8_1 * y = ((const block_q8_1 *) vy) + sample_y*stride_sample_y + channel_y*stride_channel_y; const int kbx_offset = sample_x*stride_sample_x + channel_x*stride_channel_x + row0*stride_row_x; +#if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) + if constexpr (small_k && type == GGML_TYPE_PTQ1_0 && ncols_dst == 1 && y_soa && nwarps > 1 && rows_per_cuda_block == nwarps) { + { + // Warp-per-row geometry for small K (SoA activation layout only). + // The generic small_k loop strides K-blocks across the whole block (kbx = tid; kbx += 128), so + // for K=5120 (40 blocks) only threads 0..39 ever load anything: warp 0 is full, warp 1 has + // 8 lanes, warps 2..3 only wait at __syncthreads while holding SM warp slots. Measured + // in-graph with CUPTI on an RTX 4070: 350-390 GB/s for every K=5120 GEMV vs 460 GB/s for the + // K=17408 geometry where all four warps stream. Here warp w owns row row0+w and its 32 lanes + // stride that row's K-blocks: every warp issues loads, the reduction is a shuffle, and + // there is no shared memory or block barrier. Activations are re-read per warp but they + // are 20 KB and L1/L2 resident. + const int warp = threadIdx.y; + const int lane = threadIdx.x; + const bool row_ok = uint32_t(row0 + warp) < stride_col_dst; + + float acc[ncols_dst] = { 0.0f }; + float acc_gate[ncols_dst] = { 0.0f }; + if (row_ok) { + const int kbx_row = kbx_offset + warp * stride_row_x; + for (int kbx = lane; kbx < blocks_per_row_x; kbx += warp_size) { + float dots[ncols_dst]; + vec_dot_ptq1_0_q8_1_multi(vx, y, kbx_row + kbx, kbx, stride_col_y, dots); +#pragma unroll + for (int j = 0; j < ncols_dst; ++j) { + acc[j] += dots[j]; + } + if constexpr (has_fusion && has_gate) { + vec_dot_ptq1_0_q8_1_multi(vgate, y, kbx_row + kbx, kbx, stride_col_y, dots); +#pragma unroll + for (int j = 0; j < ncols_dst; ++j) { + acc_gate[j] += dots[j]; + } + } + } + } +#pragma unroll + for (int j = 0; j < ncols_dst; ++j) { + acc[j] = warp_reduce_sum(acc[j]); + if constexpr (has_fusion && has_gate) { + acc_gate[j] = warp_reduce_sum(acc_gate[j]); + } + } + if (lane == 0 && row_ok) { + float * dst_row = dst + sample_dst*stride_sample_dst + channel_dst*stride_channel_dst + row0 + warp; + [[maybe_unused]] const uint32_t channel_bias = ids ? channel_x : channel_dst; + [[maybe_unused]] const int64_t bias_off = sample_dst*stride_sample_dst + channel_bias*stride_channel_dst + row0 + warp; +#pragma unroll + for (int j = 0; j < ncols_dst; ++j) { + float result = acc[j]; + if constexpr (has_fusion) { + if (use_bias) { + result += ((const float *) fusion.x_bias)[bias_off + j*stride_col_dst]; + } + if constexpr (has_gate) { + float gate_value = acc_gate[j]; + if (use_gate_bias) { + gate_value += ((const float *) fusion.gate_bias)[bias_off + j*stride_col_dst]; + } + switch (active_glu) { + case GGML_GLU_OP_SWIGLU: + result *= ggml_cuda_op_silu_single(gate_value); + break; + case GGML_GLU_OP_GEGLU: + result *= ggml_cuda_op_gelu_single(gate_value); + break; + case GGML_GLU_OP_SWIGLU_OAI: + result = ggml_cuda_op_swiglu_oai_single(gate_value, result); + break; + default: + result = result * gate_value; + break; + } + } + } + dst_row[j*stride_col_dst] = result; + } + } + GGML_UNUSED_VARS(use_gate, use_scale, use_gate_scale, gate_bias, x_bias, x_scale, gate_scale, x_scales, + gate_scales, x_biases, gate_biases, tmp, tmp_gate, vec_dot_q_cuda, blocks_per_iter); + return; + } + } +#endif // !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) + if constexpr ((type == GGML_TYPE_Q1_0 || type == GGML_TYPE_Q2_0 || type == GGML_TYPE_PQ2_0) && table_id == MMVQ_PARAMETERS_GB10) { using block_t = std::conditional_t 1 && ncols_dst <= 3) { + if constexpr (type == GGML_TYPE_PTQ1_0) { + GGML_UNUSED(kqs); + GGML_UNUSED(kby); + // Two activation layouts, one decision (ggml_cuda_q8_1_layout_host): Ada one-column + // reads the warp-transposed exact-isum layout through vec_dot_ptq1_0_q8_1_multi; + // Ampere one-column and every 2-8 column / MoE path take the planar kernel. + // everything else (2-8 columns, MoE ids) reads the planar layout of mmvq-ptq1_0.cuh. + // y_soa is the host's copy of that decision, baked in as a template parameter. + if constexpr (ncols_dst == 1 && y_soa) { # pragma unroll - for (int i = 0; i < rows_per_cuda_block; ++i) { - float dots[ncols_dst]; - vec_dot_ptq1_0_q8_1_multi(vx, &y[kby], kbx_offset + i * stride_row_x + kbx, kqs, - stride_col_y, dots); + for (int i = 0; i < rows_per_cuda_block; ++i) { + float dots[ncols_dst]; + vec_dot_ptq1_0_q8_1_multi(vx, y, kbx_offset + i * stride_row_x + kbx, kbx, + stride_col_y, dots); +# pragma unroll + for (int j = 0; j < ncols_dst; ++j) { + tmp[j][i] += dots[j]; + } + if constexpr (has_fusion) { + if constexpr (has_gate) { + vec_dot_ptq1_0_q8_1_multi(vgate, y, kbx_offset + i * stride_row_x + kbx, kbx, + stride_col_y, dots); +# pragma unroll + for (int j = 0; j < ncols_dst; ++j) { + tmp_gate[j][i] += dots[j]; + } + } + } + } + } else { + // planar layout: every column count runs this same code, one thread per 128-weight block + const int nblk = ptq1_0_pt_nblk(ncols_x); + const char * ycol[ncols_dst]; # pragma unroll for (int j = 0; j < ncols_dst; ++j) { - tmp[j][i] += dots[j]; + ycol[j] = (const char *) (y + j*stride_col_y); + } + const block_ptq1_0 * bq[rows_per_cuda_block]; +# pragma unroll + for (int i = 0; i < rows_per_cuda_block; ++i) { + bq[i] = (const block_ptq1_0 *) vx + kbx_offset + i*stride_row_x + kbx; + } + float dots[ncols_dst][rows_per_cuda_block]; + ptq1_0_pt_block_dot(bq, ycol, kbx, nblk, dots); +# pragma unroll + for (int j = 0; j < ncols_dst; ++j) { +# pragma unroll + for (int i = 0; i < rows_per_cuda_block; ++i) { + tmp[j][i] += dots[j][i]; + } } if constexpr (has_fusion) { if constexpr (has_gate) { - vec_dot_ptq1_0_q8_1_multi(vgate, &y[kby], kbx_offset + i * stride_row_x + kbx, kqs, - stride_col_y, dots); + const block_ptq1_0 * bg[rows_per_cuda_block]; +# pragma unroll + for (int i = 0; i < rows_per_cuda_block; ++i) { + bg[i] = (const block_ptq1_0 *) vgate + kbx_offset + i*stride_row_x + kbx; + } + ptq1_0_pt_block_dot(bg, ycol, kbx, nblk, dots); # pragma unroll for (int j = 0; j < ncols_dst; ++j) { - tmp_gate[j][i] += dots[j]; +# pragma unroll + for (int i = 0; i < rows_per_cuda_block; ++i) { + tmp_gate[j][i] += dots[j][i]; + } } } } @@ -897,9 +1062,31 @@ static __global__ void mul_mat_vec_q_moe( const int kby = kbx * (qk/QK8_1); const int kqs = vdr * (threadIdx.x % (qi/vdr)); +#if !defined(GGML_USE_HIP) + if constexpr (type == GGML_TYPE_PTQ1_0) { + // MoE always has ids, so ggml_cuda_q8_1_layout_for picked the planar layout (mmvq-ptq1_0.cuh) + const int nblk = ptq1_0_pt_nblk(ncols_x); + const char * ycol[1] = { (const char *) y }; + const block_ptq1_0 * bq[c_rows_per_block]; +#pragma unroll + for (int i = 0; i < c_rows_per_block; ++i) { + bq[i] = (const block_ptq1_0 *) vx + kbx_offset + i*stride_row_x + kbx; + } + float dots[1][c_rows_per_block]; + ptq1_0_pt_block_dot<1, c_rows_per_block>(bq, ycol, kbx, nblk, dots); #pragma unroll - for (int i = 0; i < c_rows_per_block; ++i) { - tmp[i] += vec_dot_q_cuda(vx, &y[kby], kbx_offset + i*stride_row_x + kbx, kqs); + for (int i = 0; i < c_rows_per_block; ++i) { + tmp[i] += dots[0][i]; + } + GGML_UNUSED(kqs); + GGML_UNUSED(kby); + } else +#endif + { +#pragma unroll + for (int i = 0; i < c_rows_per_block; ++i) { + tmp[i] += vec_dot_q_cuda(vx, &y[kby], kbx_offset + i*stride_row_x + kbx, kqs); + } } } @@ -922,13 +1109,30 @@ static std::pair calc_launch_params( const int ncols_dst, const int nrows_x, const int nchannels_dst, const int nsamples_or_ntokens, const int warp_size, const mmvq_parameter_table_id table_id, const bool small_k = false, const bool halve_iters = false) { const int nwarps = calc_nwarps(type, ncols_dst, table_id, small_k, halve_iters); - const int rpb = calc_rows_per_block(ncols_dst, table_id, small_k, nwarps); + const int rpb = calc_rows_per_block(type, ncols_dst, table_id, small_k, nwarps); const int64_t nblocks = (nrows_x + rpb - 1) / rpb; const dim3 block_nums(nblocks, nchannels_dst, nsamples_or_ntokens); const dim3 block_dims(warp_size, nwarps, 1); return {block_nums, block_dims}; } +// Resolves the runtime SoA decision to the y_soa template parameter. The true instantiation exists +// only for PTQ1_0 at one column, so every other type/column count compiles exactly one kernel body. +template +static void mul_mat_vec_q_launch(const bool y_soa, const ggml_cuda_kernel_launch_params & launch_params, Args&&... args) { + if constexpr (type == GGML_TYPE_PTQ1_0 && c_ncols_dst == 1) { + if (y_soa) { + ggml_cuda_kernel_launch(mul_mat_vec_q, + launch_params, std::forward(args)...); + return; + } + } else { + GGML_ASSERT(!y_soa); + } + ggml_cuda_kernel_launch(mul_mat_vec_q, + launch_params, std::forward(args)...); +} + template static void mul_mat_vec_q_switch_fusion( const void * vx, const void * vy, const int32_t * ids, const ggml_cuda_mm_fusion_args_device fusion, float * dst, @@ -937,7 +1141,7 @@ static void mul_mat_vec_q_switch_fusion( const uint32_t stride_channel_y, const uint32_t stride_channel_dst, const uint3 sample_ratio, const uint32_t stride_sample_x, const uint32_t stride_sample_y, const uint32_t stride_sample_dst, const dim3 & block_nums, const dim3 & block_dims, const int nbytes_shared, - const uint32_t ids_stride, cudaStream_t stream) { + const uint32_t ids_stride, const bool y_soa, cudaStream_t stream) { const bool has_fusion = fusion.gate != nullptr || fusion.x_bias != nullptr || fusion.gate_bias != nullptr || fusion.x_scale != nullptr || fusion.gate_scale != nullptr; @@ -945,12 +1149,12 @@ static void mul_mat_vec_q_switch_fusion( if (has_fusion) { const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params(block_nums, block_dims, nbytes_shared, stream); if (fusion.gate != nullptr) { - ggml_cuda_kernel_launch(mul_mat_vec_q, launch_params, + mul_mat_vec_q_launch(y_soa, launch_params, vx, vy, ids, fusion, dst, ncols_x, nchannels_y, stride_row_x, stride_col_y, stride_col_dst, channel_ratio, stride_channel_x, stride_channel_y, stride_channel_dst, sample_ratio, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride); } else { - ggml_cuda_kernel_launch(mul_mat_vec_q, launch_params, + mul_mat_vec_q_launch(y_soa, launch_params, vx, vy, ids, fusion, dst, ncols_x, nchannels_y, stride_row_x, stride_col_y, stride_col_dst, channel_ratio, stride_channel_x, stride_channel_y, stride_channel_dst, sample_ratio, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride); @@ -962,7 +1166,7 @@ static void mul_mat_vec_q_switch_fusion( GGML_ASSERT(!has_fusion && "fusion only supported for ncols_dst=1"); const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params(block_nums, block_dims, nbytes_shared, stream); - ggml_cuda_kernel_launch(mul_mat_vec_q, launch_params, + mul_mat_vec_q_launch(y_soa, launch_params, vx, vy, ids, fusion, dst, ncols_x, nchannels_y, stride_row_x, stride_col_y, stride_col_dst, channel_ratio, stride_channel_x, stride_channel_y, stride_channel_dst, sample_ratio, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride); @@ -998,11 +1202,27 @@ static void mul_mat_vec_q_switch_ncols_dst( const int nchannels_x, const int nchannels_y, const int nchannels_dst, const int stride_channel_x, const int stride_channel_y, const int stride_channel_dst, const int nsamples_x, const int nsamples_dst, const int stride_sample_x, const int stride_sample_y, const int stride_sample_dst, - const int ids_stride, cudaStream_t stream) { + const int ids_stride, const ggml_cuda_q8_1_layout y_layout, cudaStream_t stream) { GGML_ASSERT(ncols_x % ggml_blck_size(type) == 0); GGML_ASSERT(ncols_dst <= MMVQ_MAX_BATCH_SIZE); + // Same layout the caller handed the quantizer. Do not recompute from ncols_dst here: + // ggml_cuda_mul_mat_vec_q derives the column count as (ids ? ne2 : ne11). + const bool y_soa = y_layout == GGML_CUDA_Q8_1_SOA_ISUM; + +#if !defined(GGML_USE_HIP) + if constexpr (type == GGML_TYPE_PTQ1_0) { + // plain 2D PTQ1_0 mat-vec with 2-8 columns: dedicated kernel with full lane utilization, + // see mmvq-ptq1_0.cuh. One column takes the SoA path in the generic kernel below. + if (!ids && !y_soa && mul_mat_vec_ptq1_0_pt_switch(vx, vy, fusion, dst, ncols_x, nrows_x, ncols_dst, + stride_row_x, stride_col_y, stride_col_dst, + nchannels_dst, nsamples_dst, stream)) { + return; + } + } +#endif + const uint3 nchannels_y_fd = ids ? init_fastdiv_values(nchannels_y) : make_uint3(0, 0, 0); const uint3 channel_ratio_fd = ids ? make_uint3(0, 0, 0) : init_fastdiv_values(nchannels_dst / nchannels_x); const uint3 sample_ratio_fd = init_fastdiv_values(nsamples_dst / nsamples_x); @@ -1118,7 +1338,7 @@ static void mul_mat_vec_q_switch_ncols_dst( mul_mat_vec_q_switch_fusion( vx, vy, ids, fusion, dst, ncols_x, nchannels_y_fd, stride_row_x, stride_col_y, stride_col_dst, channel_ratio_fd, stride_channel_x, stride_channel_y, stride_channel_dst, sample_ratio_fd, - stride_sample_x, stride_sample_y, stride_sample_dst, dims.first, dims.second, 0, ids_stride, + stride_sample_x, stride_sample_y, stride_sample_dst, dims.first, dims.second, 0, ids_stride, y_soa, stream); }; @@ -1136,7 +1356,7 @@ static void mul_mat_vec_q_switch_ncols_dst( mul_mat_vec_q_switch_fusion(vx, vy, ids, fusion, dst, ncols_x, nchannels_y_fd, stride_row_x, stride_col_y, stride_col_dst, channel_ratio_fd, stride_channel_x, stride_channel_y, stride_channel_dst, sample_ratio_fd, stride_sample_x, stride_sample_y, stride_sample_dst, - dims.first, dims.second, 0, ids_stride, stream); + dims.first, dims.second, 0, ids_stride, y_soa, stream); } break; case 3: { constexpr int c_ncols_dst = 3; @@ -1144,7 +1364,7 @@ static void mul_mat_vec_q_switch_ncols_dst( mul_mat_vec_q_switch_fusion(vx, vy, ids, fusion, dst, ncols_x, nchannels_y_fd, stride_row_x, stride_col_y, stride_col_dst, channel_ratio_fd, stride_channel_x, stride_channel_y, stride_channel_dst, sample_ratio_fd, stride_sample_x, stride_sample_y, stride_sample_dst, - dims.first, dims.second, 0, ids_stride, stream); + dims.first, dims.second, 0, ids_stride, y_soa, stream); } break; case 4: { constexpr int c_ncols_dst = 4; @@ -1152,7 +1372,7 @@ static void mul_mat_vec_q_switch_ncols_dst( mul_mat_vec_q_switch_fusion(vx, vy, ids, fusion, dst, ncols_x, nchannels_y_fd, stride_row_x, stride_col_y, stride_col_dst, channel_ratio_fd, stride_channel_x, stride_channel_y, stride_channel_dst, sample_ratio_fd, stride_sample_x, stride_sample_y, stride_sample_dst, - dims.first, dims.second, 0, ids_stride, stream); + dims.first, dims.second, 0, ids_stride, y_soa, stream); } break; case 5: { constexpr int c_ncols_dst = 5; @@ -1160,7 +1380,7 @@ static void mul_mat_vec_q_switch_ncols_dst( mul_mat_vec_q_switch_fusion(vx, vy, ids, fusion, dst, ncols_x, nchannels_y_fd, stride_row_x, stride_col_y, stride_col_dst, channel_ratio_fd, stride_channel_x, stride_channel_y, stride_channel_dst, sample_ratio_fd, stride_sample_x, stride_sample_y, stride_sample_dst, - dims.first, dims.second, 0, ids_stride, stream); + dims.first, dims.second, 0, ids_stride, y_soa, stream); } break; case 6: { constexpr int c_ncols_dst = 6; @@ -1168,7 +1388,7 @@ static void mul_mat_vec_q_switch_ncols_dst( mul_mat_vec_q_switch_fusion(vx, vy, ids, fusion, dst, ncols_x, nchannels_y_fd, stride_row_x, stride_col_y, stride_col_dst, channel_ratio_fd, stride_channel_x, stride_channel_y, stride_channel_dst, sample_ratio_fd, stride_sample_x, stride_sample_y, stride_sample_dst, - dims.first, dims.second, 0, ids_stride, stream); + dims.first, dims.second, 0, ids_stride, y_soa, stream); } break; case 7: { constexpr int c_ncols_dst = 7; @@ -1176,7 +1396,7 @@ static void mul_mat_vec_q_switch_ncols_dst( mul_mat_vec_q_switch_fusion(vx, vy, ids, fusion, dst, ncols_x, nchannels_y_fd, stride_row_x, stride_col_y, stride_col_dst, channel_ratio_fd, stride_channel_x, stride_channel_y, stride_channel_dst, sample_ratio_fd, stride_sample_x, stride_sample_y, stride_sample_dst, - dims.first, dims.second, 0, ids_stride, stream); + dims.first, dims.second, 0, ids_stride, y_soa, stream); } break; case 8: { constexpr int c_ncols_dst = 8; @@ -1184,7 +1404,7 @@ static void mul_mat_vec_q_switch_ncols_dst( mul_mat_vec_q_switch_fusion(vx, vy, ids, fusion, dst, ncols_x, nchannels_y_fd, stride_row_x, stride_col_y, stride_col_dst, channel_ratio_fd, stride_channel_x, stride_channel_y, stride_channel_dst, sample_ratio_fd, stride_sample_x, stride_sample_y, stride_sample_dst, - dims.first, dims.second, 0, ids_stride, stream); + dims.first, dims.second, 0, ids_stride, y_soa, stream); } break; default: GGML_ABORT("fatal error"); @@ -1198,157 +1418,157 @@ static void mul_mat_vec_q_switch_type( const int nchannels_x, const int nchannels_y, const int nchannels_dst, const int stride_channel_x, const int stride_channel_y, const int stride_channel_dst, const int nsamples_x, const int nsamples_dst, const int stride_sample_x, const int stride_sample_y, const int stride_sample_dst, - const int ids_stride, cudaStream_t stream) { + const int ids_stride, const ggml_cuda_q8_1_layout y_layout, cudaStream_t stream) { switch (type_x) { case GGML_TYPE_Q1_0: mul_mat_vec_q_switch_ncols_dst (vx, vy, ids, fusion, dst, ncols_x, nrows_x, ncols_dst, stride_row_x, stride_col_y, stride_col_dst, nchannels_x, nchannels_y, nchannels_dst, stride_channel_x, stride_channel_y, stride_channel_dst, - nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, stream); + nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, y_layout, stream); break; case GGML_TYPE_PQ2_0: mul_mat_vec_q_switch_ncols_dst (vx, vy, ids, fusion, dst, ncols_x, nrows_x, ncols_dst, stride_row_x, stride_col_y, stride_col_dst, nchannels_x, nchannels_y, nchannels_dst, stride_channel_x, stride_channel_y, stride_channel_dst, - nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, stream); + nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, y_layout, stream); break; case GGML_TYPE_PTQ1_0: mul_mat_vec_q_switch_ncols_dst (vx, vy, ids, fusion, dst, ncols_x, nrows_x, ncols_dst, stride_row_x, stride_col_y, stride_col_dst, nchannels_x, nchannels_y, nchannels_dst, stride_channel_x, stride_channel_y, stride_channel_dst, - nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, stream); + nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, y_layout, stream); break; case GGML_TYPE_Q2_0: mul_mat_vec_q_switch_ncols_dst (vx, vy, ids, fusion, dst, ncols_x, nrows_x, ncols_dst, stride_row_x, stride_col_y, stride_col_dst, nchannels_x, nchannels_y, nchannels_dst, stride_channel_x, stride_channel_y, stride_channel_dst, - nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, stream); + nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, y_layout, stream); break; case GGML_TYPE_Q4_0: mul_mat_vec_q_switch_ncols_dst (vx, vy, ids, fusion, dst, ncols_x, nrows_x, ncols_dst, stride_row_x, stride_col_y, stride_col_dst, nchannels_x, nchannels_y, nchannels_dst, stride_channel_x, stride_channel_y, stride_channel_dst, - nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, stream); + nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, y_layout, stream); break; case GGML_TYPE_Q4_1: mul_mat_vec_q_switch_ncols_dst (vx, vy, ids, fusion, dst, ncols_x, nrows_x, ncols_dst, stride_row_x, stride_col_y, stride_col_dst, nchannels_x, nchannels_y, nchannels_dst, stride_channel_x, stride_channel_y, stride_channel_dst, - nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, stream); + nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, y_layout, stream); break; case GGML_TYPE_Q5_0: mul_mat_vec_q_switch_ncols_dst (vx, vy, ids, fusion, dst, ncols_x, nrows_x, ncols_dst, stride_row_x, stride_col_y, stride_col_dst, nchannels_x, nchannels_y, nchannels_dst, stride_channel_x, stride_channel_y, stride_channel_dst, - nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, stream); + nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, y_layout, stream); break; case GGML_TYPE_Q5_1: mul_mat_vec_q_switch_ncols_dst (vx, vy, ids, fusion, dst, ncols_x, nrows_x, ncols_dst, stride_row_x, stride_col_y, stride_col_dst, nchannels_x, nchannels_y, nchannels_dst, stride_channel_x, stride_channel_y, stride_channel_dst, - nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, stream); + nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, y_layout, stream); break; case GGML_TYPE_Q8_0: mul_mat_vec_q_switch_ncols_dst (vx, vy, ids, fusion, dst, ncols_x, nrows_x, ncols_dst, stride_row_x, stride_col_y, stride_col_dst, nchannels_x, nchannels_y, nchannels_dst, stride_channel_x, stride_channel_y, stride_channel_dst, - nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, stream); + nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, y_layout, stream); break; case GGML_TYPE_MXFP4: mul_mat_vec_q_switch_ncols_dst (vx, vy, ids, fusion, dst, ncols_x, nrows_x, ncols_dst, stride_row_x, stride_col_y, stride_col_dst, nchannels_x, nchannels_y, nchannels_dst, stride_channel_x, stride_channel_y, stride_channel_dst, - nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, stream); + nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, y_layout, stream); break; case GGML_TYPE_NVFP4: mul_mat_vec_q_switch_ncols_dst (vx, vy, ids, fusion, dst, ncols_x, nrows_x, ncols_dst, stride_row_x, stride_col_y, stride_col_dst, nchannels_x, nchannels_y, nchannels_dst, stride_channel_x, stride_channel_y, stride_channel_dst, - nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, stream); + nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, y_layout, stream); break; case GGML_TYPE_Q2_K: mul_mat_vec_q_switch_ncols_dst (vx, vy, ids, fusion, dst, ncols_x, nrows_x, ncols_dst, stride_row_x, stride_col_y, stride_col_dst, nchannels_x, nchannels_y, nchannels_dst, stride_channel_x, stride_channel_y, stride_channel_dst, - nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, stream); + nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, y_layout, stream); break; case GGML_TYPE_Q3_K: mul_mat_vec_q_switch_ncols_dst (vx, vy, ids, fusion, dst, ncols_x, nrows_x, ncols_dst, stride_row_x, stride_col_y, stride_col_dst, nchannels_x, nchannels_y, nchannels_dst, stride_channel_x, stride_channel_y, stride_channel_dst, - nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, stream); + nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, y_layout, stream); break; case GGML_TYPE_Q4_K: mul_mat_vec_q_switch_ncols_dst (vx, vy, ids, fusion, dst, ncols_x, nrows_x, ncols_dst, stride_row_x, stride_col_y, stride_col_dst, nchannels_x, nchannels_y, nchannels_dst, stride_channel_x, stride_channel_y, stride_channel_dst, - nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, stream); + nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, y_layout, stream); break; case GGML_TYPE_Q5_K: mul_mat_vec_q_switch_ncols_dst (vx, vy, ids, fusion, dst, ncols_x, nrows_x, ncols_dst, stride_row_x, stride_col_y, stride_col_dst, nchannels_x, nchannels_y, nchannels_dst, stride_channel_x, stride_channel_y, stride_channel_dst, - nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, stream); + nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, y_layout, stream); break; case GGML_TYPE_Q6_K: mul_mat_vec_q_switch_ncols_dst (vx, vy, ids, fusion, dst, ncols_x, nrows_x, ncols_dst, stride_row_x, stride_col_y, stride_col_dst, nchannels_x, nchannels_y, nchannels_dst, stride_channel_x, stride_channel_y, stride_channel_dst, - nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, stream); + nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, y_layout, stream); break; case GGML_TYPE_IQ2_XXS: mul_mat_vec_q_switch_ncols_dst (vx, vy, ids, fusion, dst, ncols_x, nrows_x, ncols_dst, stride_row_x, stride_col_y, stride_col_dst, nchannels_x, nchannels_y, nchannels_dst, stride_channel_x, stride_channel_y, stride_channel_dst, - nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, stream); + nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, y_layout, stream); break; case GGML_TYPE_IQ2_XS: mul_mat_vec_q_switch_ncols_dst (vx, vy, ids, fusion, dst, ncols_x, nrows_x, ncols_dst, stride_row_x, stride_col_y, stride_col_dst, nchannels_x, nchannels_y, nchannels_dst, stride_channel_x, stride_channel_y, stride_channel_dst, - nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, stream); + nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, y_layout, stream); break; case GGML_TYPE_IQ2_S: mul_mat_vec_q_switch_ncols_dst (vx, vy, ids, fusion, dst, ncols_x, nrows_x, ncols_dst, stride_row_x, stride_col_y, stride_col_dst, nchannels_x, nchannels_y, nchannels_dst, stride_channel_x, stride_channel_y, stride_channel_dst, - nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, stream); + nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, y_layout, stream); break; case GGML_TYPE_IQ3_XXS: mul_mat_vec_q_switch_ncols_dst (vx, vy, ids, fusion, dst, ncols_x, nrows_x, ncols_dst, stride_row_x, stride_col_y, stride_col_dst, nchannels_x, nchannels_y, nchannels_dst, stride_channel_x, stride_channel_y, stride_channel_dst, - nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, stream); + nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, y_layout, stream); break; case GGML_TYPE_IQ1_S: mul_mat_vec_q_switch_ncols_dst (vx, vy, ids, fusion, dst, ncols_x, nrows_x, ncols_dst, stride_row_x, stride_col_y, stride_col_dst, nchannels_x, nchannels_y, nchannels_dst, stride_channel_x, stride_channel_y, stride_channel_dst, - nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, stream); + nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, y_layout, stream); break; case GGML_TYPE_IQ1_M: mul_mat_vec_q_switch_ncols_dst (vx, vy, ids, fusion, dst, ncols_x, nrows_x, ncols_dst, stride_row_x, stride_col_y, stride_col_dst, nchannels_x, nchannels_y, nchannels_dst, stride_channel_x, stride_channel_y, stride_channel_dst, - nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, stream); + nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, y_layout, stream); break; case GGML_TYPE_IQ4_NL: mul_mat_vec_q_switch_ncols_dst (vx, vy, ids, fusion, dst, ncols_x, nrows_x, ncols_dst, stride_row_x, stride_col_y, stride_col_dst, nchannels_x, nchannels_y, nchannels_dst, stride_channel_x, stride_channel_y, stride_channel_dst, - nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, stream); + nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, y_layout, stream); break; case GGML_TYPE_IQ4_XS: mul_mat_vec_q_switch_ncols_dst (vx, vy, ids, fusion, dst, ncols_x, nrows_x, ncols_dst, stride_row_x, stride_col_y, stride_col_dst, nchannels_x, nchannels_y, nchannels_dst, stride_channel_x, stride_channel_y, stride_channel_dst, - nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, stream); + nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, y_layout, stream); break; case GGML_TYPE_IQ3_S: mul_mat_vec_q_switch_ncols_dst (vx, vy, ids, fusion, dst, ncols_x, nrows_x, ncols_dst, stride_row_x, stride_col_y, stride_col_dst, nchannels_x, nchannels_y, nchannels_dst, stride_channel_x, stride_channel_y, stride_channel_dst, - nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, stream); + nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, y_layout, stream); break; default: GGML_ABORT("fatal error"); @@ -1433,13 +1653,31 @@ void ggml_cuda_mul_mat_vec_q( } } - const int64_t ne10_padded = GGML_PAD(ne10, MATRIX_ROW_PADDING); - ggml_cuda_pool_alloc src1_q8_1(ctx.pool(), ne13*ne12 * ne11*ne10_padded * sizeof(block_q8_1)/QK8_1); - { + // Same (type, ncols_dst, ids) triple the kernel switch evaluates, so quantizer and kernel agree. + const ggml_cuda_q8_1_layout y_layout = ggml_cuda_q8_1_layout_host(src0->type, (int) (ids ? ne2 : ne11), ids != nullptr); + + int64_t ne10_padded = GGML_PAD(ne10, MATRIX_ROW_PADDING); + if (y_layout == GGML_CUDA_Q8_1_SOA_ISUM) { + // Warp-transposed q8 layout needs whole 32-K-block groups per column. + ne10_padded = GGML_PAD(ne10_padded, GGML_CUDA_PTQ1_K_PAD); + } + // src1 may be a Hadamard transform output whose producer already wrote the q8_1 rows in this + // exact layout (ggml_cuda_fwht_q8, into src1's buffer or a held pool block); nothing to quantize. + const ggml_cuda_fwht_q8 * pre_q8 = ids ? nullptr : ctx.fwht_q8().find(src1); + if (pre_q8) { + GGML_ASSERT(pre_q8->layout == y_layout && pre_q8->ne0 == ne10_padded && pre_q8->ncols == ne11 && ne12 == 1 && ne13 == 1); + } + + ggml_cuda_pool_alloc src1_q8_1(ctx.pool()); + const char * src1_q8_1_d = nullptr; + if (pre_q8) { + src1_q8_1_d = (const char *) pre_q8->data; + } else { + src1_q8_1_d = src1_q8_1.alloc(ne13*ne12 * ne11*ne10_padded * sizeof(block_q8_1)/QK8_1); const int64_t s11 = src1->nb[1] / ts_src1; const int64_t s12 = src1->nb[2] / ts_src1; const int64_t s13 = src1->nb[3] / ts_src1; - quantize_row_q8_1_cuda(src1_d, nullptr, src1_q8_1.get(), src0->type, ne10, s11, s12, s13, ne10_padded, ne11, ne12, ne13, stream); + quantize_row_q8_1_cuda(src1_d, nullptr, src1_q8_1.get(), y_layout, ne10, s11, s12, s13, ne10_padded, ne11, ne12, ne13, stream); } const int64_t s01 = src0->nb[1] / ts_src0; @@ -1465,10 +1703,10 @@ void ggml_cuda_mul_mat_vec_q( const int64_t ids_stride = ids ? ids->nb[1] / ggml_type_size(ids->type) : 0; mul_mat_vec_q_switch_type( - src0->data, src0->type, src1_q8_1.get(), ids_d, fusion_local, dst_d, ne00, + src0->data, src0->type, src1_q8_1_d, ids_d, fusion_local, dst_d, ne00, ne01, ncols_dst, s01, stride_col_y, stride_col_dst, ne02, nchannels_y, nchannels_dst, s02, stride_channel_y, stride_channel_dst, - ne03, ne3, s03, s13, s3, ids_stride, stream); + ne03, ne3, s03, s13, s3, ids_stride, y_layout, stream); } void ggml_cuda_op_mul_mat_vec_q( @@ -1495,9 +1733,10 @@ void ggml_cuda_op_mul_mat_vec_q( const int stride_col_y = src1_padded_row_size / QK8_1; ggml_cuda_mm_fusion_args_device fusion_local{}; + const ggml_cuda_q8_1_layout y_layout = ggml_cuda_q8_1_layout_host(src0->type, (int) src1_ncols, false); mul_mat_vec_q_switch_type( src0_dd_i, src0->type, src1_ddq_i, nullptr, fusion_local, dst_dd_i, ne00, row_diff, src1_ncols, stride_row_x, stride_col_y, nrows_dst, - 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 0, stream); + 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 0, y_layout, stream); GGML_UNUSED_VARS(src1, dst, src1_ddf_i, src1_ncols, src1_padded_row_size); } diff --git a/ggml/src/ggml-cuda/quantize.cu b/ggml/src/ggml-cuda/quantize.cu index fbbc6314ab33..b2a435f09fea 100644 --- a/ggml/src/ggml-cuda/quantize.cu +++ b/ggml/src/ggml-cuda/quantize.cu @@ -1,4 +1,6 @@ #include "quantize.cuh" +#include "fwht.cuh" +#include "mmvq-ptq1_0.cuh" #include "unary.cuh" #include @@ -51,6 +53,16 @@ static __device__ __forceinline__ float nvfp4_native_scale_error( #endif // CUDART_VERSION >= 12080 #endif // defined(BLACKWELL_MMA_AVAILABLE) +// layout (ggml_cuda_q8_1_layout): +// AOS plain block_q8_1, float input sum in ds.y +// SOA_ISUM warp-transposed layout (ggml_cuda_ptq1_q8_word) with the exact integer sum of the +// quantized values bit-cast into ds.y; the 1-column PTQ1_0 vec-dot folds the digit bias +// {0,1,2} -> {-1,0,+1} into one subtraction per 32-block instead of a SIMD byte +// subtract per 4 weights (bit-identical to the biased path) +// PT planar-transposed layout consumed by the 2-8 column PTQ1_0 mat-vec path +// (see mmvq-ptq1_0.cuh); same quantization, same bytes per row as block_q8_1, and the +// same exact integer sum in ds.y as SOA_ISUM so that kernel can also use unbiased trits +template __launch_bounds__(CUDA_QUANTIZE_BLOCK_SIZE, 1) static __global__ void quantize_q8_1( const float * x_ptr, void * vy_ptr, @@ -92,6 +104,44 @@ static __global__ void quantize_q8_1( const float d = amax / 127.0f; const int8_t q = amax == 0.0f ? 0 : roundf(xi / d); + if constexpr (layout == GGML_CUDA_Q8_1_PT) { + const int64_t row_cont = (i3*ne2.z + i2) * ne1 + i1; + char * ycol = (char *) vy + row_cont * (ne0 * 9 / 8); // same row stride as block_q8_1 + const int64_t nblk = ne0 / QK_PTQ1_0; + const int64_t kb = i0 / QK_PTQ1_0; + const int e = i0 % QK_PTQ1_0; + ycol[((e / 16)*nblk + kb) * 16 + (e % 16)] = q; + + int isum = q; + isum = warp_reduce_sum(isum); + + if (iqs > 0) { + return; + } + + // |isum| <= 32*127 fits int16; keep the raw bits in the half slot (read back with __half_as_short). + half2 * ds = (half2 *) (ycol + 8*nblk*16) + kb*4 + e / QK8_1; + *ds = make_half2(__float2half(d), __short_as_half((short) isum)); + return; + } + if constexpr (layout == GGML_CUDA_Q8_1_SOA_ISUM) { + // Warp-transposed layout (see ggml_cuda_ptq1_q8_word); ne0 is padded to GGML_CUDA_PTQ1_K_PAD. + int isum = q; + isum = warp_reduce_sum(isum); + int32_t * yw = (int32_t *) vy; + const int64_t col = i_cont / ne0; + const int64_t col_base = col * (ne0 / QK8_1) * (int64_t) (sizeof(block_q8_1) / 4); + const int ib_col = (int) (i0 / QK8_1); + ((int8_t *) (yw + col_base + ggml_cuda_ptq1_q8_word(ib_col, iqs / 4)))[iqs % 4] = q; + if (iqs > 0) { + return; + } + // |isum| <= 32*127 fits int16; keep the raw bits in the half slot. + const half2 ds = make_half2(__float2half(d), __short_as_half((short) isum)); + yw[col_base + ggml_cuda_ptq1_q8_word(ib_col, 8)] = *reinterpret_cast(&ds); + return; + } + y[ib].qs[iqs] = q; if (iqs > 0) { @@ -101,6 +151,176 @@ static __global__ void quantize_q8_1( y[ib].ds = make_half2(d, sum); } +// Hadamard transform folded into the activation quantizer (see ggml_cuda_fwht_skip): one block per +// (transform width N, row), the same butterfly as fwht_cuda_block, then the same per-32 quantization +// and stores as quantize_q8_1 above. Element i*NT + tid of the transform sits in reg[i], so +// warp w holds the complete 32-blocks kb = i*(NT/32) + w and reduces them with the usual warp reductions. +// Rows of x are contiguous (ne00 elements); row index = ((i3*ne2 + i2)*ne1 + i1) like a contiguous src1. +template +__launch_bounds__(NT, 1) +static __global__ void fwht_quantize_q8_1( + const T * x_ptr, const float * signs, void * vy_ptr, + const int64_t ne00, const int64_t ne0, const uint32_t ne1, const uint3 ne2, const float scale) { + constexpr int warp_size = ggml_cuda_get_physical_warp_size(); + constexpr int NE = N / NT; + static_assert(NE >= 1 && N % NT == 0 && NT % warp_size == 0 && NT % QK8_1 == 0, "bad fused FWHT shape"); + static_assert(QK8_1 == warp_size, "fused quantizer reduces one 32-block per warp"); + + __shared__ float s[N]; + + ggml_cuda_pdl_lc(); + + const int64_t blk = blockIdx.x; // transform block within the row + const int64_t i3 = fastdiv(blockIdx.z, ne2); + const int64_t i2 = blockIdx.z - i3*ne2.z; + const int64_t i1 = blockIdx.y; + const int64_t row = (i3*ne2.z + i2) * ne1 + i1; + + const int tid = threadIdx.x; + const int lane = tid % warp_size; + + const int64_t base = blk * N; // first element of this transform in the row + + float reg[NE]; + ggml_cuda_pdl_sync(); + if (base < ne00) { + const T * src = x_ptr + row * ne00 + base; + const float * signs_row = has_signs ? signs + base : nullptr; // (r % n_blk) * N with r = row*n_blk + blk +#pragma unroll + for (int i = 0; i < NE; ++i) { + reg[i] = (float) src[i * NT + tid] * scale; + if (has_signs) { + reg[i] *= signs_row[i * NT + tid]; + } + } + ggml_cuda_fwht_block_butterfly(reg, s, tid, lane); + } else { + // padding beyond the real row: zero blocks, same as the plain quantizer's x = 0 path +#pragma unroll + for (int i = 0; i < NE; ++i) { + reg[i] = 0.0f; + } + } + + void * GGML_CUDA_RESTRICT vy = vy_ptr; + +#pragma unroll + for (int i = 0; i < NE; ++i) { + const float xi = reg[i]; + const int64_t i0 = base + i * NT + tid; // element within the row (i0 % 32 == lane) + + float amax = fabsf(xi); + float sum = xi; + amax = warp_reduce_max(amax); + sum = warp_reduce_sum(sum); + + const float d = amax / 127.0f; + const int8_t q = amax == 0.0f ? 0 : roundf(xi / d); + + const int iqs = lane; + + if constexpr (layout == GGML_CUDA_Q8_1_PT) { + char * ycol = (char *) vy + row * (ne0 * 9 / 8); + const int64_t nblk = ne0 / QK_PTQ1_0; + const int64_t kb = i0 / QK_PTQ1_0; + const int e = i0 % QK_PTQ1_0; + ycol[((e / 16)*nblk + kb) * 16 + (e % 16)] = q; + + int isum = q; + isum = warp_reduce_sum(isum); + if (iqs == 0) { + half2 * ds = (half2 *) (ycol + 8*nblk*16) + kb*4 + e / QK8_1; + *ds = make_half2(__float2half(d), __short_as_half((short) isum)); + } + } else if constexpr (layout == GGML_CUDA_Q8_1_SOA_ISUM) { + int isum = q; + isum = warp_reduce_sum(isum); + int32_t * yw = (int32_t *) vy; + const int64_t col_base = row * (ne0 / QK8_1) * (int64_t) (sizeof(block_q8_1) / 4); + const int ib_col = (int) (i0 / QK8_1); + ((int8_t *) (yw + col_base + ggml_cuda_ptq1_q8_word(ib_col, iqs / 4)))[iqs % 4] = q; + if (iqs == 0) { + const half2 ds = make_half2(__float2half(d), __short_as_half((short) isum)); + yw[col_base + ggml_cuda_ptq1_q8_word(ib_col, 8)] = *reinterpret_cast(&ds); + } + } else { + block_q8_1 * y = (block_q8_1 *) vy; + const int64_t ib = (row * ne0 + i0) / QK8_1; + y[ib].qs[iqs] = q; + if (iqs == 0) { + y[ib].ds = make_half2(d, sum); + } + } + } +} + +template +static void fwht_quantize_launch_layout( + const T * x, const float * signs, void * vy, const ggml_cuda_q8_1_layout layout, + const int64_t ne00, const int64_t ne0, const int64_t ne1, const int64_t ne2, const int64_t ne3, + const float scale, cudaStream_t stream) { + constexpr int NT = N < 256 ? N : 256; + const uint3 ne2_fastdiv = init_fastdiv_values(ne2); + const dim3 num_blocks((unsigned) (ne0 / N), (unsigned) ne1, (unsigned) (ne2*ne3)); + const dim3 block_size(NT, 1, 1); + const ggml_cuda_kernel_launch_params lp = ggml_cuda_kernel_launch_params(num_blocks, block_size, 0, stream); +#define FWHT_Q_LAUNCH(LAYOUT) \ + if (signs) { \ + ggml_cuda_kernel_launch(fwht_quantize_q8_1, lp, x, signs, vy, ne00, ne0, (uint32_t) ne1, ne2_fastdiv, scale); \ + } else { \ + ggml_cuda_kernel_launch(fwht_quantize_q8_1, lp, x, signs, vy, ne00, ne0, (uint32_t) ne1, ne2_fastdiv, scale); \ + } + switch (layout) { + case GGML_CUDA_Q8_1_PT: FWHT_Q_LAUNCH(GGML_CUDA_Q8_1_PT); break; + case GGML_CUDA_Q8_1_SOA_ISUM: FWHT_Q_LAUNCH(GGML_CUDA_Q8_1_SOA_ISUM); break; + default: FWHT_Q_LAUNCH(GGML_CUDA_Q8_1_AOS); break; + } +#undef FWHT_Q_LAUNCH +} + +template +static void fwht_quantize_launch( + const T * x, const float * signs, void * vy, const ggml_cuda_q8_1_layout layout, const int n, + const int64_t ne00, const int64_t ne0, const int64_t ne1, const int64_t ne2, const int64_t ne3, + const float scale, cudaStream_t stream) { + switch (n) { + case 64: fwht_quantize_launch_layout< 64, T>(x, signs, vy, layout, ne00, ne0, ne1, ne2, ne3, scale, stream); break; + case 128: fwht_quantize_launch_layout< 128, T>(x, signs, vy, layout, ne00, ne0, ne1, ne2, ne3, scale, stream); break; + case 256: fwht_quantize_launch_layout< 256, T>(x, signs, vy, layout, ne00, ne0, ne1, ne2, ne3, scale, stream); break; + case 512: fwht_quantize_launch_layout< 512, T>(x, signs, vy, layout, ne00, ne0, ne1, ne2, ne3, scale, stream); break; + case 1024: fwht_quantize_launch_layout<1024, T>(x, signs, vy, layout, ne00, ne0, ne1, ne2, ne3, scale, stream); break; + case 2048: fwht_quantize_launch_layout<2048, T>(x, signs, vy, layout, ne00, ne0, ne1, ne2, ne3, scale, stream); break; + default: GGML_ABORT("fused FWHT quantizer: unsupported transform width %d", n); + } +} + +bool ggml_cuda_fwht_quantize_supported(const int n, const int64_t ne00) { + switch (n) { + case 64: case 128: case 256: case 512: case 1024: case 2048: + break; + default: + return false; + } + return ne00 % n == 0 && ne00 % QK8_1 == 0; +} + +void fwht_quantize_row_q8_1_cuda( + const void * x, const ggml_type x_type, const float * signs, const int n, void * vy, + const ggml_cuda_q8_1_layout layout, const int64_t ne00, const int64_t ne0, const int64_t ncols, cudaStream_t stream) { + GGML_ASSERT(ggml_cuda_fwht_quantize_supported(n, ne00)); + GGML_ASSERT(ne0 % n == 0 && ne0 >= ne00); + if (layout == GGML_CUDA_Q8_1_PT) { + GGML_ASSERT(ne0 % QK_PTQ1_0 == 0); + } + const float scale = 1.0f / sqrtf((float) n); + if (x_type == GGML_TYPE_F16) { + fwht_quantize_launch((const half *) x, signs, vy, layout, n, ne00, ne0, ncols, 1, 1, scale, stream); + } else { + GGML_ASSERT(x_type == GGML_TYPE_F32); + fwht_quantize_launch((const float *) x, signs, vy, layout, n, ne00, ne0, ncols, 1, 1, scale, stream); + } +} + __device__ __forceinline__ uint8_t compute_e8m0_scale(float amax) { if (!(amax > 0.0f)) { return 0; @@ -636,7 +856,7 @@ void quantize_mmq_q8_1_rms_cuda( } void quantize_row_q8_1_cuda( - const float * x, const int32_t * ids, void * vy, const ggml_type type_src0, + const float * x, const int32_t * ids, void * vy, const ggml_cuda_q8_1_layout layout, const int64_t ne00, const int64_t s01, const int64_t s02, const int64_t s03, const int64_t ne0, const int64_t ne1, const int64_t ne2, const int64_t ne3, cudaStream_t stream) { GGML_ASSERT(!ids); @@ -648,8 +868,19 @@ void quantize_row_q8_1_cuda( const dim3 num_blocks(block_num_x, ne1, ne2*ne3); const dim3 block_size(CUDA_QUANTIZE_BLOCK_SIZE, 1, 1); const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params(num_blocks, block_size, 0, stream); - ggml_cuda_kernel_launch(quantize_q8_1, launch_params, x, vy, ne00, s01, s02, s03, ne0, ne1, ne2_fastdiv); - GGML_UNUSED(type_src0); + switch (layout) { + case GGML_CUDA_Q8_1_PT: + GGML_ASSERT(ne0 % QK_PTQ1_0 == 0); + ggml_cuda_kernel_launch(quantize_q8_1, launch_params, x, vy, ne00, s01, s02, s03, ne0, ne1, ne2_fastdiv); + break; + case GGML_CUDA_Q8_1_SOA_ISUM: + GGML_ASSERT(ne0 % GGML_CUDA_PTQ1_K_PAD == 0); + ggml_cuda_kernel_launch(quantize_q8_1, launch_params, x, vy, ne00, s01, s02, s03, ne0, ne1, ne2_fastdiv); + break; + default: + ggml_cuda_kernel_launch(quantize_q8_1, launch_params, x, vy, ne00, s01, s02, s03, ne0, ne1, ne2_fastdiv); + break; + } } void quantize_mmq_q8_1_cuda( diff --git a/ggml/src/ggml-cuda/quantize.cuh b/ggml/src/ggml-cuda/quantize.cuh index 335ff587adb5..f9a80f2f03d4 100644 --- a/ggml/src/ggml-cuda/quantize.cuh +++ b/ggml/src/ggml-cuda/quantize.cuh @@ -16,11 +16,21 @@ typedef void (*quantize_cuda_t)( ggml_type type_src0, int64_t ne00, int64_t s01, int64_t s02, int64_t s03, int64_t ne0, int64_t ne1, int64_t ne2, int64_t ne3, cudaStream_t stream); +// layout: see ggml_cuda_q8_1_layout_for() in common.cuh. The caller must pass the same layout the +// consuming mat-vec kernel expects for (type_src0, ncols_dst, ids). void quantize_row_q8_1_cuda( const float * x, const int32_t * ids, void * vy, - ggml_type type_src0, int64_t ne00, int64_t s01, int64_t s02, int64_t s03, + ggml_cuda_q8_1_layout layout, int64_t ne00, int64_t s01, int64_t s02, int64_t s03, int64_t ne0, int64_t ne1, int64_t ne2, int64_t ne3, cudaStream_t stream); +// Hadamard transform (width n, optional per-element signs, scale 1/sqrt(n)) + q8_1 quantization in one +// kernel (ggml_cuda_fwht_q8). x holds ncols contiguous rows of ne00 elements; ne0 is the padded row +// width the consuming mat-vec expects (pad blocks are written as zeros). +bool ggml_cuda_fwht_quantize_supported(int n, int64_t ne00); +void fwht_quantize_row_q8_1_cuda( + const void * x, ggml_type x_type, const float * signs, int n, void * vy, + ggml_cuda_q8_1_layout layout, int64_t ne00, int64_t ne0, int64_t ncols, cudaStream_t stream); + void quantize_mmq_q8_1_cuda( const float * x, const int32_t * ids, void * vy, ggml_type type_src0, int64_t ne00, int64_t s01, int64_t s02, int64_t s03, diff --git a/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_1-ncols2_8.cu b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_1-ncols2_8.cu new file mode 100644 index 000000000000..355554a1689c --- /dev/null +++ b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_1-ncols2_8.cu @@ -0,0 +1,6 @@ +// This file has been autogenerated by generate_cu_files.py, do not edit manually. + +#include "../fattn-mma-f16.cuh" + +DECL_FATTN_MMA_F16_CASE_KV(128, 128, 1, 8, GGML_TYPE_Q4_0); +DECL_FATTN_MMA_F16_CASE_KV(256, 256, 1, 8, GGML_TYPE_Q4_0); diff --git a/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_16-ncols2_1.cu b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_16-ncols2_1.cu new file mode 100644 index 000000000000..a79fa0782a83 --- /dev/null +++ b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_16-ncols2_1.cu @@ -0,0 +1,6 @@ +// This file has been autogenerated by generate_cu_files.py, do not edit manually. + +#include "../fattn-mma-f16.cuh" + +DECL_FATTN_MMA_F16_CASE_KV(128, 128, 16, 1, GGML_TYPE_Q4_0); +DECL_FATTN_MMA_F16_CASE_KV(256, 256, 16, 1, GGML_TYPE_Q4_0); diff --git a/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_16-ncols2_2.cu b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_16-ncols2_2.cu new file mode 100644 index 000000000000..709e5de20542 --- /dev/null +++ b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_16-ncols2_2.cu @@ -0,0 +1,6 @@ +// This file has been autogenerated by generate_cu_files.py, do not edit manually. + +#include "../fattn-mma-f16.cuh" + +DECL_FATTN_MMA_F16_CASE_KV(128, 128, 16, 2, GGML_TYPE_Q4_0); +DECL_FATTN_MMA_F16_CASE_KV(256, 256, 16, 2, GGML_TYPE_Q4_0); diff --git a/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_16-ncols2_4.cu b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_16-ncols2_4.cu new file mode 100644 index 000000000000..6f475ec0d22b --- /dev/null +++ b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_16-ncols2_4.cu @@ -0,0 +1,6 @@ +// This file has been autogenerated by generate_cu_files.py, do not edit manually. + +#include "../fattn-mma-f16.cuh" + +DECL_FATTN_MMA_F16_CASE_KV(128, 128, 16, 4, GGML_TYPE_Q4_0); +DECL_FATTN_MMA_F16_CASE_KV(256, 256, 16, 4, GGML_TYPE_Q4_0); diff --git a/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_2-ncols2_4.cu b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_2-ncols2_4.cu new file mode 100644 index 000000000000..4144c8e3ab40 --- /dev/null +++ b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_2-ncols2_4.cu @@ -0,0 +1,6 @@ +// This file has been autogenerated by generate_cu_files.py, do not edit manually. + +#include "../fattn-mma-f16.cuh" + +DECL_FATTN_MMA_F16_CASE_KV(128, 128, 2, 4, GGML_TYPE_Q4_0); +DECL_FATTN_MMA_F16_CASE_KV(256, 256, 2, 4, GGML_TYPE_Q4_0); diff --git a/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_2-ncols2_8.cu b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_2-ncols2_8.cu new file mode 100644 index 000000000000..859de760e9aa --- /dev/null +++ b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_2-ncols2_8.cu @@ -0,0 +1,6 @@ +// This file has been autogenerated by generate_cu_files.py, do not edit manually. + +#include "../fattn-mma-f16.cuh" + +DECL_FATTN_MMA_F16_CASE_KV(128, 128, 2, 8, GGML_TYPE_Q4_0); +DECL_FATTN_MMA_F16_CASE_KV(256, 256, 2, 8, GGML_TYPE_Q4_0); diff --git a/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_32-ncols2_1.cu b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_32-ncols2_1.cu new file mode 100644 index 000000000000..5b116735761e --- /dev/null +++ b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_32-ncols2_1.cu @@ -0,0 +1,6 @@ +// This file has been autogenerated by generate_cu_files.py, do not edit manually. + +#include "../fattn-mma-f16.cuh" + +DECL_FATTN_MMA_F16_CASE_KV(128, 128, 32, 1, GGML_TYPE_Q4_0); +DECL_FATTN_MMA_F16_CASE_KV(256, 256, 32, 1, GGML_TYPE_Q4_0); diff --git a/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_32-ncols2_2.cu b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_32-ncols2_2.cu new file mode 100644 index 000000000000..861a34c1444c --- /dev/null +++ b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_32-ncols2_2.cu @@ -0,0 +1,6 @@ +// This file has been autogenerated by generate_cu_files.py, do not edit manually. + +#include "../fattn-mma-f16.cuh" + +DECL_FATTN_MMA_F16_CASE_KV(128, 128, 32, 2, GGML_TYPE_Q4_0); +DECL_FATTN_MMA_F16_CASE_KV(256, 256, 32, 2, GGML_TYPE_Q4_0); diff --git a/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_4-ncols2_2.cu b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_4-ncols2_2.cu new file mode 100644 index 000000000000..99c3234ee476 --- /dev/null +++ b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_4-ncols2_2.cu @@ -0,0 +1,6 @@ +// This file has been autogenerated by generate_cu_files.py, do not edit manually. + +#include "../fattn-mma-f16.cuh" + +DECL_FATTN_MMA_F16_CASE_KV(128, 128, 4, 2, GGML_TYPE_Q4_0); +DECL_FATTN_MMA_F16_CASE_KV(256, 256, 4, 2, GGML_TYPE_Q4_0); diff --git a/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_4-ncols2_4.cu b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_4-ncols2_4.cu new file mode 100644 index 000000000000..d909b720b7b7 --- /dev/null +++ b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_4-ncols2_4.cu @@ -0,0 +1,6 @@ +// This file has been autogenerated by generate_cu_files.py, do not edit manually. + +#include "../fattn-mma-f16.cuh" + +DECL_FATTN_MMA_F16_CASE_KV(128, 128, 4, 4, GGML_TYPE_Q4_0); +DECL_FATTN_MMA_F16_CASE_KV(256, 256, 4, 4, GGML_TYPE_Q4_0); diff --git a/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_4-ncols2_8.cu b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_4-ncols2_8.cu new file mode 100644 index 000000000000..cdee7c6fbbdd --- /dev/null +++ b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_4-ncols2_8.cu @@ -0,0 +1,6 @@ +// This file has been autogenerated by generate_cu_files.py, do not edit manually. + +#include "../fattn-mma-f16.cuh" + +DECL_FATTN_MMA_F16_CASE_KV(128, 128, 4, 8, GGML_TYPE_Q4_0); +DECL_FATTN_MMA_F16_CASE_KV(256, 256, 4, 8, GGML_TYPE_Q4_0); diff --git a/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_64-ncols2_1.cu b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_64-ncols2_1.cu new file mode 100644 index 000000000000..ab677bed38bb --- /dev/null +++ b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_64-ncols2_1.cu @@ -0,0 +1,6 @@ +// This file has been autogenerated by generate_cu_files.py, do not edit manually. + +#include "../fattn-mma-f16.cuh" + +DECL_FATTN_MMA_F16_CASE_KV(128, 128, 64, 1, GGML_TYPE_Q4_0); +DECL_FATTN_MMA_F16_CASE_KV(256, 256, 64, 1, GGML_TYPE_Q4_0); diff --git a/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_8-ncols2_1.cu b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_8-ncols2_1.cu new file mode 100644 index 000000000000..d9de18d31435 --- /dev/null +++ b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_8-ncols2_1.cu @@ -0,0 +1,6 @@ +// This file has been autogenerated by generate_cu_files.py, do not edit manually. + +#include "../fattn-mma-f16.cuh" + +DECL_FATTN_MMA_F16_CASE_KV(128, 128, 8, 1, GGML_TYPE_Q4_0); +DECL_FATTN_MMA_F16_CASE_KV(256, 256, 8, 1, GGML_TYPE_Q4_0); diff --git a/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_8-ncols2_2.cu b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_8-ncols2_2.cu new file mode 100644 index 000000000000..68f48ae6a4f4 --- /dev/null +++ b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_8-ncols2_2.cu @@ -0,0 +1,6 @@ +// This file has been autogenerated by generate_cu_files.py, do not edit manually. + +#include "../fattn-mma-f16.cuh" + +DECL_FATTN_MMA_F16_CASE_KV(128, 128, 8, 2, GGML_TYPE_Q4_0); +DECL_FATTN_MMA_F16_CASE_KV(256, 256, 8, 2, GGML_TYPE_Q4_0); diff --git a/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_8-ncols2_4.cu b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_8-ncols2_4.cu new file mode 100644 index 000000000000..c7655a7b1325 --- /dev/null +++ b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_8-ncols2_4.cu @@ -0,0 +1,6 @@ +// This file has been autogenerated by generate_cu_files.py, do not edit manually. + +#include "../fattn-mma-f16.cuh" + +DECL_FATTN_MMA_F16_CASE_KV(128, 128, 8, 4, GGML_TYPE_Q4_0); +DECL_FATTN_MMA_F16_CASE_KV(256, 256, 8, 4, GGML_TYPE_Q4_0); diff --git a/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_8-ncols2_8.cu b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_8-ncols2_8.cu new file mode 100644 index 000000000000..e9cb67038d16 --- /dev/null +++ b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q4_0-ncols1_8-ncols2_8.cu @@ -0,0 +1,6 @@ +// This file has been autogenerated by generate_cu_files.py, do not edit manually. + +#include "../fattn-mma-f16.cuh" + +DECL_FATTN_MMA_F16_CASE_KV(128, 128, 8, 8, GGML_TYPE_Q4_0); +DECL_FATTN_MMA_F16_CASE_KV(256, 256, 8, 8, GGML_TYPE_Q4_0); diff --git a/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_1-ncols2_8.cu b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_1-ncols2_8.cu new file mode 100644 index 000000000000..ef09f7b7bcbd --- /dev/null +++ b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_1-ncols2_8.cu @@ -0,0 +1,6 @@ +// This file has been autogenerated by generate_cu_files.py, do not edit manually. + +#include "../fattn-mma-f16.cuh" + +DECL_FATTN_MMA_F16_CASE_KV(128, 128, 1, 8, GGML_TYPE_Q8_0); +DECL_FATTN_MMA_F16_CASE_KV(256, 256, 1, 8, GGML_TYPE_Q8_0); diff --git a/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_16-ncols2_1.cu b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_16-ncols2_1.cu new file mode 100644 index 000000000000..c7a048adef7e --- /dev/null +++ b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_16-ncols2_1.cu @@ -0,0 +1,6 @@ +// This file has been autogenerated by generate_cu_files.py, do not edit manually. + +#include "../fattn-mma-f16.cuh" + +DECL_FATTN_MMA_F16_CASE_KV(128, 128, 16, 1, GGML_TYPE_Q8_0); +DECL_FATTN_MMA_F16_CASE_KV(256, 256, 16, 1, GGML_TYPE_Q8_0); diff --git a/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_16-ncols2_2.cu b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_16-ncols2_2.cu new file mode 100644 index 000000000000..6dd74f06b9be --- /dev/null +++ b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_16-ncols2_2.cu @@ -0,0 +1,6 @@ +// This file has been autogenerated by generate_cu_files.py, do not edit manually. + +#include "../fattn-mma-f16.cuh" + +DECL_FATTN_MMA_F16_CASE_KV(128, 128, 16, 2, GGML_TYPE_Q8_0); +DECL_FATTN_MMA_F16_CASE_KV(256, 256, 16, 2, GGML_TYPE_Q8_0); diff --git a/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_16-ncols2_4.cu b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_16-ncols2_4.cu new file mode 100644 index 000000000000..8341b5f9276b --- /dev/null +++ b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_16-ncols2_4.cu @@ -0,0 +1,6 @@ +// This file has been autogenerated by generate_cu_files.py, do not edit manually. + +#include "../fattn-mma-f16.cuh" + +DECL_FATTN_MMA_F16_CASE_KV(128, 128, 16, 4, GGML_TYPE_Q8_0); +DECL_FATTN_MMA_F16_CASE_KV(256, 256, 16, 4, GGML_TYPE_Q8_0); diff --git a/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_2-ncols2_4.cu b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_2-ncols2_4.cu new file mode 100644 index 000000000000..991759e39e81 --- /dev/null +++ b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_2-ncols2_4.cu @@ -0,0 +1,6 @@ +// This file has been autogenerated by generate_cu_files.py, do not edit manually. + +#include "../fattn-mma-f16.cuh" + +DECL_FATTN_MMA_F16_CASE_KV(128, 128, 2, 4, GGML_TYPE_Q8_0); +DECL_FATTN_MMA_F16_CASE_KV(256, 256, 2, 4, GGML_TYPE_Q8_0); diff --git a/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_2-ncols2_8.cu b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_2-ncols2_8.cu new file mode 100644 index 000000000000..0cdd4b1e7c21 --- /dev/null +++ b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_2-ncols2_8.cu @@ -0,0 +1,6 @@ +// This file has been autogenerated by generate_cu_files.py, do not edit manually. + +#include "../fattn-mma-f16.cuh" + +DECL_FATTN_MMA_F16_CASE_KV(128, 128, 2, 8, GGML_TYPE_Q8_0); +DECL_FATTN_MMA_F16_CASE_KV(256, 256, 2, 8, GGML_TYPE_Q8_0); diff --git a/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_32-ncols2_1.cu b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_32-ncols2_1.cu new file mode 100644 index 000000000000..b0d5f0402748 --- /dev/null +++ b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_32-ncols2_1.cu @@ -0,0 +1,6 @@ +// This file has been autogenerated by generate_cu_files.py, do not edit manually. + +#include "../fattn-mma-f16.cuh" + +DECL_FATTN_MMA_F16_CASE_KV(128, 128, 32, 1, GGML_TYPE_Q8_0); +DECL_FATTN_MMA_F16_CASE_KV(256, 256, 32, 1, GGML_TYPE_Q8_0); diff --git a/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_32-ncols2_2.cu b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_32-ncols2_2.cu new file mode 100644 index 000000000000..e4086fba93bd --- /dev/null +++ b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_32-ncols2_2.cu @@ -0,0 +1,6 @@ +// This file has been autogenerated by generate_cu_files.py, do not edit manually. + +#include "../fattn-mma-f16.cuh" + +DECL_FATTN_MMA_F16_CASE_KV(128, 128, 32, 2, GGML_TYPE_Q8_0); +DECL_FATTN_MMA_F16_CASE_KV(256, 256, 32, 2, GGML_TYPE_Q8_0); diff --git a/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_4-ncols2_2.cu b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_4-ncols2_2.cu new file mode 100644 index 000000000000..0388387f6c65 --- /dev/null +++ b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_4-ncols2_2.cu @@ -0,0 +1,6 @@ +// This file has been autogenerated by generate_cu_files.py, do not edit manually. + +#include "../fattn-mma-f16.cuh" + +DECL_FATTN_MMA_F16_CASE_KV(128, 128, 4, 2, GGML_TYPE_Q8_0); +DECL_FATTN_MMA_F16_CASE_KV(256, 256, 4, 2, GGML_TYPE_Q8_0); diff --git a/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_4-ncols2_4.cu b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_4-ncols2_4.cu new file mode 100644 index 000000000000..29886f86766a --- /dev/null +++ b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_4-ncols2_4.cu @@ -0,0 +1,6 @@ +// This file has been autogenerated by generate_cu_files.py, do not edit manually. + +#include "../fattn-mma-f16.cuh" + +DECL_FATTN_MMA_F16_CASE_KV(128, 128, 4, 4, GGML_TYPE_Q8_0); +DECL_FATTN_MMA_F16_CASE_KV(256, 256, 4, 4, GGML_TYPE_Q8_0); diff --git a/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_4-ncols2_8.cu b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_4-ncols2_8.cu new file mode 100644 index 000000000000..b0d857926025 --- /dev/null +++ b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_4-ncols2_8.cu @@ -0,0 +1,6 @@ +// This file has been autogenerated by generate_cu_files.py, do not edit manually. + +#include "../fattn-mma-f16.cuh" + +DECL_FATTN_MMA_F16_CASE_KV(128, 128, 4, 8, GGML_TYPE_Q8_0); +DECL_FATTN_MMA_F16_CASE_KV(256, 256, 4, 8, GGML_TYPE_Q8_0); diff --git a/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_64-ncols2_1.cu b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_64-ncols2_1.cu new file mode 100644 index 000000000000..c5c46959f2d2 --- /dev/null +++ b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_64-ncols2_1.cu @@ -0,0 +1,6 @@ +// This file has been autogenerated by generate_cu_files.py, do not edit manually. + +#include "../fattn-mma-f16.cuh" + +DECL_FATTN_MMA_F16_CASE_KV(128, 128, 64, 1, GGML_TYPE_Q8_0); +DECL_FATTN_MMA_F16_CASE_KV(256, 256, 64, 1, GGML_TYPE_Q8_0); diff --git a/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_8-ncols2_1.cu b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_8-ncols2_1.cu new file mode 100644 index 000000000000..ba806751602d --- /dev/null +++ b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_8-ncols2_1.cu @@ -0,0 +1,6 @@ +// This file has been autogenerated by generate_cu_files.py, do not edit manually. + +#include "../fattn-mma-f16.cuh" + +DECL_FATTN_MMA_F16_CASE_KV(128, 128, 8, 1, GGML_TYPE_Q8_0); +DECL_FATTN_MMA_F16_CASE_KV(256, 256, 8, 1, GGML_TYPE_Q8_0); diff --git a/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_8-ncols2_2.cu b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_8-ncols2_2.cu new file mode 100644 index 000000000000..b2fdf9639c55 --- /dev/null +++ b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_8-ncols2_2.cu @@ -0,0 +1,6 @@ +// This file has been autogenerated by generate_cu_files.py, do not edit manually. + +#include "../fattn-mma-f16.cuh" + +DECL_FATTN_MMA_F16_CASE_KV(128, 128, 8, 2, GGML_TYPE_Q8_0); +DECL_FATTN_MMA_F16_CASE_KV(256, 256, 8, 2, GGML_TYPE_Q8_0); diff --git a/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_8-ncols2_4.cu b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_8-ncols2_4.cu new file mode 100644 index 000000000000..cfe288e52d4a --- /dev/null +++ b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_8-ncols2_4.cu @@ -0,0 +1,6 @@ +// This file has been autogenerated by generate_cu_files.py, do not edit manually. + +#include "../fattn-mma-f16.cuh" + +DECL_FATTN_MMA_F16_CASE_KV(128, 128, 8, 4, GGML_TYPE_Q8_0); +DECL_FATTN_MMA_F16_CASE_KV(256, 256, 8, 4, GGML_TYPE_Q8_0); diff --git a/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_8-ncols2_8.cu b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_8-ncols2_8.cu new file mode 100644 index 000000000000..1c35d44c201b --- /dev/null +++ b/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-q8_0-ncols1_8-ncols2_8.cu @@ -0,0 +1,6 @@ +// This file has been autogenerated by generate_cu_files.py, do not edit manually. + +#include "../fattn-mma-f16.cuh" + +DECL_FATTN_MMA_F16_CASE_KV(128, 128, 8, 8, GGML_TYPE_Q8_0); +DECL_FATTN_MMA_F16_CASE_KV(256, 256, 8, 8, GGML_TYPE_Q8_0); diff --git a/ggml/src/ggml-cuda/template-instances/generate_cu_files.py b/ggml/src/ggml-cuda/template-instances/generate_cu_files.py index c54928a50ae6..073e41546449 100755 --- a/ggml/src/ggml-cuda/template-instances/generate_cu_files.py +++ b/ggml/src/ggml-cuda/template-instances/generate_cu_files.py @@ -34,6 +34,11 @@ SOURCE_FATTN_MMA_CASE = "DECL_FATTN_MMA_F16_CASE({head_size_kq}, {head_size_v}, {ncols1}, {ncols2});\n" +SOURCE_FATTN_MMA_CASE_KV = "DECL_FATTN_MMA_F16_CASE_KV({head_size_kq}, {head_size_v}, {ncols1}, {ncols2}, {type_kv});\n" + +TYPES_KV_MMA_NATIVE = ["GGML_TYPE_Q4_0", "GGML_TYPE_Q8_0"] +HEAD_SIZES_MMA_NATIVE = [128, 256] + TYPES_MMQ = [ "GGML_TYPE_Q1_0", "GGML_TYPE_Q2_0", @@ -113,6 +118,17 @@ def get_short_name(long_quant_name): head_size_v = HEAD_SIZES_V_OVERRIDE.get(head_size_kq, head_size_kq) f.write(SOURCE_FATTN_MMA_CASE.format(ncols1=ncols1, ncols2=ncols2, head_size_kq=head_size_kq, head_size_v=head_size_v)) +# In-place quantized K/V variants of the MMA kernel (no F16 copy of the cache): the head sizes and GQA tilings +# a q4_0 / q8_0 cache is served with in practice. Kept to one file per (type, ncols1, ncols2) for parallel builds. +for type_kv in TYPES_KV_MMA_NATIVE: + for ncols in [8, 16, 32, 64]: + for ncols2 in [1, 2, 4, 8]: + ncols1 = ncols // ncols2 + with open(f"fattn-mma-f16-instance-{get_short_name(type_kv)}-ncols1_{ncols1}-ncols2_{ncols2}.cu", "w") as f: + f.write(SOURCE_FATTN_MMA_START) + for head_size in HEAD_SIZES_MMA_NATIVE: + f.write(SOURCE_FATTN_MMA_CASE_KV.format(ncols1=ncols1, ncols2=ncols2, head_size_kq=head_size, head_size_v=head_size, type_kv=type_kv)) + for type in TYPES_MMQ: with open(f"mmq-instance-{get_short_name(type)}.cu", "w") as f: source = SOURCE_MMQ_CUDA_ONLY if type == "GGML_TYPE_PTQ1_0" else SOURCE_MMQ diff --git a/ggml/src/ggml-cuda/vecdotq.cuh b/ggml/src/ggml-cuda/vecdotq.cuh index 65dbbb9573be..ce5f47ee3292 100644 --- a/ggml/src/ggml-cuda/vecdotq.cuh +++ b/ggml/src/ggml-cuda/vecdotq.cuh @@ -805,16 +805,26 @@ static __device__ __forceinline__ float vec_dot_q2_0_q8_1( } #if !defined(GGML_USE_HIP) +// y_col0: activation column 0 base (warp-transposed layout, see ggml_cuda_ptq1_q8_word); +// kbx_x: absolute PTQ1 block index into vbq; kb: K-block index within the row; +// stride_col_y: column stride in block_q8_1 units. template static __device__ __forceinline__ void vec_dot_ptq1_0_q8_1_multi(const void * __restrict__ vbq, - const block_q8_1 * __restrict__ bq8_1, - const int & kbx, - const int & iqs, + const block_q8_1 * __restrict__ y_col0, + const int & kbx_x, + const int & kb, const uint32_t stride_col_y, float * result) { - const block_ptq1_0 * bq = (const block_ptq1_0 *) vbq + kbx; + const block_ptq1_0 * bq = (const block_ptq1_0 *) vbq + kbx_x; int sumi[ncols_dst][4] = {}; + // Lane-coalesced activation words: word ww of this K-block is at yw[ww*32]; 32 lanes of a + // warp own 32 consecutive K-blocks, so every load below is 32 consecutive words. + const int * __restrict__ yw = (const int *) y_col0 + (kb >> 5) * GGML_CUDA_PTQ1_Q8_GROUP_WORDS + (kb & 31); + const uint32_t sy = stride_col_y * (sizeof(block_q8_1) / 4); +# define PTQ1_U(j, sub, m) yw[(j) * sy + ((sub) * 8 + (m)) * GGML_CUDA_PTQ1_Q8_GROUP_KB] +# define PTQ1_DS(j, sub) yw[(j) * sy + (32 + (sub)) * GGML_CUDA_PTQ1_Q8_GROUP_KB] + // Widen four bytes to 16-bit lanes so multiply-by-three cannot carry between bytes. # pragma unroll for (int g = 0; g < 4; ++g) { @@ -829,11 +839,12 @@ static __device__ __forceinline__ void vec_dot_ptq1_0_q8_1_multi(const void * __ v_lo = w_lo & 0x00FF00FF; v_hi = w_hi & 0x00FF00FF; - const int q = __vsub4(__byte_perm(w_lo, w_hi, 0x7531), 0x01010101); + // Raw digits {0,1,2}; the -1 bias is folded into the per-32-block exact q8 sum below. + const int q = __byte_perm(w_lo, w_hi, 0x7531); const int e = t * 16 + 4 * g; # pragma unroll for (int j = 0; j < ncols_dst; ++j) { - const int u = get_int_b4(bq8_1[j * stride_col_y + iqs + (e >> 5)].qs, (e & 31) >> 2); + const int u = PTQ1_U(j, e >> 5, (e & 31) >> 2); sumi[j][e >> 5] = ggml_cuda_dp4a(q, u, sumi[j][e >> 5]); } } @@ -852,11 +863,11 @@ static __device__ __forceinline__ void vec_dot_ptq1_0_q8_1_multi(const void * __ v_lo = w_lo & 0x00FF00FF; v_hi = w_hi & 0x00FF00FF; - const int q = __vsub4(__byte_perm(w_lo, w_hi, 0x7531), 0x01010101); + const int q = __byte_perm(w_lo, w_hi, 0x7531); const int e = 80 + t * 8 + 4 * g; # pragma unroll for (int j = 0; j < ncols_dst; ++j) { - const int u = get_int_b4(bq8_1[j * stride_col_y + iqs + (e >> 5)].qs, (e & 31) >> 2); + const int u = PTQ1_U(j, e >> 5, (e & 31) >> 2); sumi[j][e >> 5] = ggml_cuda_dp4a(q, u, sumi[j][e >> 5]); } } @@ -870,10 +881,10 @@ static __device__ __forceinline__ void vec_dot_ptq1_0_q8_1_multi(const void * __ const uint32_t w1 = v * 3; v = w1 & 0x00FF00FF; - const int q = __vsub4(__byte_perm(w0, w1, 0x7531), 0x01010101); + const int q = __byte_perm(w0, w1, 0x7531); # pragma unroll for (int j = 0; j < ncols_dst; ++j) { - const int u = get_int_b4(bq8_1[j * stride_col_y + iqs + 3].qs, 6 + t / 2); + const int u = PTQ1_U(j, 3, 6 + t / 2); sumi[j][3] = ggml_cuda_dp4a(q, u, sumi[j][3]); } } @@ -884,10 +895,16 @@ static __device__ __forceinline__ void vec_dot_ptq1_0_q8_1_multi(const void * __ float acc = 0.0f; # pragma unroll for (int k = 0; k < 4; ++k) { - acc += __low2float(bq8_1[j * stride_col_y + iqs + k].ds) * (float) sumi[j][k]; + const int ds_bits = PTQ1_DS(j, k); + const half2 ds = *reinterpret_cast(&ds_bits); + // ds.y carries the exact int16 sum of the 32 q8 values (see quantize_q8_1). + const int isum = (int) __half_as_short(__high2half(ds)); + acc += __low2float(ds) * (float) (sumi[j][k] - isum); } result[j] = d * acc; } +# undef PTQ1_U +# undef PTQ1_DS } #endif @@ -968,10 +985,64 @@ static __device__ __forceinline__ float vec_dot_ptq1_0_q8_1(const void * __restr acc += __low2float(bq8_1[iqs + k].ds) * (float) (sumi[k] - sumu[k]); } return (float) bq->d * acc; +#elif defined(GGML_USE_MUSA) + // MUSA cannot compile the HIP amdgcn path. Scalar decode, same as the pre-#211 generic entry. + const block_ptq1_0 * bq = (const block_ptq1_0 *) vbq + kbx; + int sumi[4] = { 0, 0, 0, 0 }; + +# pragma unroll + for (int m = 0; m < 16; ++m) { + uint32_t v = bq->qs[m]; +# pragma unroll + for (int t = 0; t < 5; ++t) { + const uint32_t w = v * 3; + const int q = (int) (w >> 8) - 1; + v = w & 0xFF; + const int e = t * 16 + m; + sumi[e >> 5] += q * (int) bq8_1[iqs + (e >> 5)].qs[e & 31]; + } + } + +# pragma unroll + for (int m = 0; m < 8; ++m) { + uint32_t v = bq->qs[16 + m]; +# pragma unroll + for (int t = 0; t < 5; ++t) { + const uint32_t w = v * 3; + const int q = (int) (w >> 8) - 1; + v = w & 0xFF; + const int e = 80 + t * 8 + m; + sumi[e >> 5] += q * (int) bq8_1[iqs + (e >> 5)].qs[e & 31]; + } + } + +# pragma unroll + for (int h = 0; h < 2; ++h) { + uint32_t v = bq->qh[h]; +# pragma unroll + for (int t = 0; t < 4; ++t) { + const uint32_t w = v * 3; + const int q = (int) (w >> 8) - 1; + v = w & 0xFF; + const int e = 120 + t * 2 + h; + sumi[e >> 5] += q * (int) bq8_1[iqs + (e >> 5)].qs[e & 31]; + } + } + + float acc = 0.0f; +# pragma unroll + for (int k = 0; k < 4; ++k) { + acc += __low2float(bq8_1[iqs + k].ds) * (float) sumi[k]; + } + return (float) bq->d * acc; #else - float result; - vec_dot_ptq1_0_q8_1_multi<1>(vbq, bq8_1, kbx, iqs, 0, &result); - return result; + // NVIDIA PTQ1_0 goes through vec_dot_ptq1_0_q8_1_multi; this generic entry is unused there. + // The SoA _multi signature does not match a plain AoS call, so do not forward. + GGML_UNUSED(vbq); + GGML_UNUSED(bq8_1); + GGML_UNUSED(kbx); + GGML_UNUSED(iqs); + return 0.0f; #endif } diff --git a/ggml/src/ggml-hip/CMakeLists.txt b/ggml/src/ggml-hip/CMakeLists.txt index 47f16f56c470..8ce638f16616 100644 --- a/ggml/src/ggml-hip/CMakeLists.txt +++ b/ggml/src/ggml-hip/CMakeLists.txt @@ -64,6 +64,8 @@ file(GLOB GGML_SOURCES_ROCM "../ggml-cuda/*.cu") file(GLOB SRCS "../ggml-cuda/template-instances/fattn-tile*.cu") list(APPEND GGML_SOURCES_ROCM ${SRCS}) file(GLOB SRCS "../ggml-cuda/template-instances/fattn-mma*.cu") +# Native q4_0/q8_0 K/V MMA instances are NVIDIA-only (ggml_cuda_fattn_mma_kv_native_supported). +list(FILTER SRCS EXCLUDE REGEX "fattn-mma-f16-instance-q[48]_0-") list(APPEND GGML_SOURCES_ROCM ${SRCS}) file(GLOB SRCS "../ggml-cuda/template-instances/mmq*.cu") list(APPEND GGML_SOURCES_ROCM ${SRCS}) diff --git a/ggml/src/ggml-musa/CMakeLists.txt b/ggml/src/ggml-musa/CMakeLists.txt index cc53c812ce5f..af2334e9f18d 100644 --- a/ggml/src/ggml-musa/CMakeLists.txt +++ b/ggml/src/ggml-musa/CMakeLists.txt @@ -32,8 +32,10 @@ if (MUSAToolkit_FOUND) file(GLOB GGML_SOURCES_MUSA "../ggml-cuda/*.cu") file(GLOB SRCS "../ggml-cuda/template-instances/fattn-tile*.cu") list(APPEND GGML_SOURCES_MUSA ${SRCS}) - file(GLOB SRCS "../ggml-cuda/template-instances/fattn-mma*.cu") - list(APPEND GGML_SOURCES_MUSA ${SRCS}) +file(GLOB SRCS "../ggml-cuda/template-instances/fattn-mma*.cu") +# Native q4_0/q8_0 K/V MMA instances are NVIDIA-only (ggml_cuda_fattn_mma_kv_native_supported). +list(FILTER SRCS EXCLUDE REGEX "fattn-mma-f16-instance-q[48]_0-") +list(APPEND GGML_SOURCES_MUSA ${SRCS}) file(GLOB SRCS "../ggml-cuda/template-instances/mmq*.cu") list(APPEND GGML_SOURCES_MUSA ${SRCS}) diff --git a/src/models/qwen35.cpp b/src/models/qwen35.cpp index 2f143d367a4d..1fb752e86448 100644 --- a/src/models/qwen35.cpp +++ b/src/models/qwen35.cpp @@ -755,9 +755,18 @@ llama_model_qwen35::graph_mtp::graph_mtp(const llama_model & model, const llm_gr cur = build_norm(cur, head_norm_w, nullptr, LLM_NORM_RMS, -1); cb(cur, "h_nextn", -1); - res->t_h_nextn = cur; - cur = ggml_get_rows(ctx0, cur, inp_out_ids); + // llama_context copies the first n_outputs rows of t_h_nextn when embeddings_nextn_masked is set, + // so publish the output rows only (as the trunk graph does); unmasked readers index by token. + // Matters as soon as an MTP batch has more tokens than outputs, e.g. catch-up rows decoded + // together with the first draft row. + if (cparams.embeddings_nextn_masked) { + cur = ggml_get_rows(ctx0, cur, inp_out_ids); + res->t_h_nextn = cur; + } else { + res->t_h_nextn = cur; + cur = ggml_get_rows(ctx0, cur, inp_out_ids); + } cb(cur, "mtp_shared_head_norm", -1); ggml_tensor * head_w = layer.nextn.shared_head_head ? layer.nextn.shared_head_head : model.output; diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index 3e5874ce61c4..f0bbee4ce969 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -301,6 +301,7 @@ llama_build_and_test(test-thread-safety.cpp ARGS -m "${MODEL_DEST}" -ngl 99 -p " set_tests_properties(test-thread-safety PROPERTIES FIXTURES_REQUIRED test-download-model) llama_build_and_test(test-arg-parser.cpp) +llama_build_and_test(test-mtp-catchup-batch.cpp) llama_build_and_test(test-model-resolution.cpp) # the test serves its repos from an httplib server, and the library links it privately target_link_libraries(test-model-resolution PRIVATE cpp-httplib) diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 6510e15f5446..b2239f4c4ebe 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -10138,6 +10138,32 @@ static std::vector> make_test_cases_eval() { test_cases.emplace_back(new test_flash_attn_ext(64, 64, 4, {1, 1}, 1024, 75, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q8_0, {0, 2, 1, 3}, false)); test_cases.emplace_back(new test_flash_attn_ext(64, 64, 4, {1, 1}, 512, 75, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q8_0, {0, 1, 2, 3}, false)); + // in-place quantized K/V for the tensor-core prompt kernels (CUDA reads q4_0 / q8_0 straight from the cache at + // hs 128 / 256 instead of converting the whole cache to F16): every GQA tiling, prompt-sized batches, a KV + // length that is not a multiple of the tile (out-of-bounds rows), and the permuted cache layout. + for (ggml_type type_KV : { GGML_TYPE_Q4_0, GGML_TYPE_Q8_0 }) { + for (int64_t kv : { 1024, 4096, 16384 }) { + for (int64_t nb : { 3, 8, 35, 512 }) { + test_cases.emplace_back(new test_flash_attn_ext(256, 256, 4, {6, 1}, kv, nb, true, false, 0, 0, GGML_PREC_F32, type_KV, type_KV)); + } + } + for (int64_t nr : { 1, 2, 4, 8 }) { + test_cases.emplace_back(new test_flash_attn_ext(128, 128, 8, {nr, 1}, 4096, 512, true, false, 0, 0, GGML_PREC_F32, type_KV, type_KV)); + test_cases.emplace_back(new test_flash_attn_ext(256, 256, 4, {nr, 1}, 4096, 35, true, false, 0, 0, GGML_PREC_F32, type_KV, type_KV)); + } + test_cases.emplace_back(new test_flash_attn_ext(256, 256, 4, {6, 1}, 1025, 512, true, false, 0, 0, GGML_PREC_F32, type_KV, type_KV)); + test_cases.emplace_back(new test_flash_attn_ext(128, 128, 8, {4, 1}, 1025, 64, true, false, 0, 0, GGML_PREC_F32, type_KV, type_KV)); + test_cases.emplace_back(new test_flash_attn_ext(256, 256, 4, {6, 1}, 4096, 512, true, false, 0, 0, GGML_PREC_F32, type_KV, type_KV, {0, 2, 1, 3})); + test_cases.emplace_back(new test_flash_attn_ext(256, 256, 4, {6, 1}, 4096, 35, true, false, 0, 0, GGML_PREC_F32, type_KV, type_KV, {0, 1, 2, 3}, false)); + // Isolated ALiBi / Gemma softcap on the native q4_0/q8_0 MMA path. The loop above is + // max_bias=0, logit_softcap=0; the one older q8_0 case at D=256 sets both at once with sinks. + test_cases.emplace_back(new test_flash_attn_ext(128, 128, 8, {4, 1}, 1024, 8, true, false, 8.0f, 0, GGML_PREC_F32, type_KV, type_KV)); + test_cases.emplace_back(new test_flash_attn_ext(128, 128, 8, {4, 1}, 1024, 8, true, false, 0, 10.0f, GGML_PREC_F32, type_KV, type_KV)); + test_cases.emplace_back(new test_flash_attn_ext(256, 256, 4, {6, 1}, 1024, 8, true, false, 8.0f, 0, GGML_PREC_F32, type_KV, type_KV)); + test_cases.emplace_back(new test_flash_attn_ext(128, 128, 8, {4, 1}, 1024, 8, true, false, 8.0f, 0, GGML_PREC_F32, type_KV, type_KV, {0, 2, 1, 3})); + test_cases.emplace_back(new test_flash_attn_ext(128, 128, 8, {4, 1}, 1024, 8, true, false, 0, 10.0f, GGML_PREC_F32, type_KV, type_KV, {0, 2, 1, 3})); + } + test_cases.emplace_back(new test_cross_entropy_loss (GGML_TYPE_F32, { 10, 5, 4, 3})); test_cases.emplace_back(new test_cross_entropy_loss (GGML_TYPE_F32, {30000, 1, 1, 1})); test_cases.emplace_back(new test_cross_entropy_loss_back(GGML_TYPE_F32, { 10, 5, 4, 3})); @@ -10187,6 +10213,18 @@ static std::vector> make_test_cases_eval() { } } + // PTQ1_0 dedicated mat-vec, shared-memory boundary on this branch. The launch asks for + // ncols * rows_per_cta * (K/128 + 1) * 4 bytes (the +1 is the odd epilogue stride), twice + // that with a gate, against the 48 KiB default. One column and 4 rows per CTA: last K that + // fits is 393088 without a gate and 196480 with one; the next block over each falls back. + for (int64_t k : {393088, 393216}) { + test_cases.emplace_back(new test_mul_mat(GGML_TYPE_PTQ1_0, GGML_TYPE_F32, 64, 1, k, {1, 1}, {1, 1})); + } + for (int64_t k : {196480, 196608}) { + test_cases.emplace_back(new test_mul_mat_vec_fusion(GGML_TYPE_PTQ1_0, GGML_GLU_OP_SWIGLU, 1, 64, k, + false, 1, 1, false, false, true, false, {1, 1})); + } + for (auto gate : {GATING_FUNC_SOFTMAX, GATING_FUNC_SIGMOID, GATING_FUNC_SOFTMAX_WEIGHT, GATING_FUNC_SQRT_SOFTPLUS}) { for (bool with_norm : {false, true}) { for (bool bias_probs : {false, true}) { @@ -10489,6 +10527,23 @@ static std::vector> make_test_cases_perf() { } } + // Ternary Bonsai 2 27B (qwen35): the bf16 gated-delta-net gate projections + for (int bs : {1, 2, 3, 4, 8}) { + test_cases.emplace_back(new test_mul_mat(GGML_TYPE_BF16, GGML_TYPE_F32, 48, bs, 5120, {1, 1}, {1, 1})); // ssm_alpha, ssm_beta + } + + // Ternary Bonsai 2 27B (qwen35, PTQ1_0) projections at speculative-decoding batch sizes + for (int bs : {1, 2, 3, 4, 8}) { + for (ggml_type type_a : {GGML_TYPE_PTQ1_0, GGML_TYPE_Q4_0}) { + test_cases.emplace_back(new test_mul_mat(type_a, GGML_TYPE_F32, 10240, bs, 5120, {1, 1}, {1, 1})); // attn_qkv + test_cases.emplace_back(new test_mul_mat(type_a, GGML_TYPE_F32, 6144, bs, 5120, {1, 1}, {1, 1})); // attn_gate + test_cases.emplace_back(new test_mul_mat(type_a, GGML_TYPE_F32, 5120, bs, 6144, {1, 1}, {1, 1})); // ssm_out + test_cases.emplace_back(new test_mul_mat(type_a, GGML_TYPE_F32, 17408, bs, 5120, {1, 1}, {1, 1})); // ffn_up, ffn_gate + test_cases.emplace_back(new test_mul_mat(type_a, GGML_TYPE_F32, 5120, bs, 17408, {1, 1}, {1, 1})); // ffn_down + test_cases.emplace_back(new test_mul_mat(type_a, GGML_TYPE_F32, 12288, bs, 5120, {1, 1}, {1, 1})); // attn_q + } + } + // qwen3-30b-a3b for (int bs : {1, 4, 8, 32, 64, 128, 256, 512}) { for (ggml_type type_a : {GGML_TYPE_F32, GGML_TYPE_F16, GGML_TYPE_Q4_0, GGML_TYPE_Q8_0, GGML_TYPE_Q4_K, GGML_TYPE_Q6_K, GGML_TYPE_IQ2_XS}) { diff --git a/tests/test-mtp-catchup-batch.cpp b/tests/test-mtp-catchup-batch.cpp new file mode 100644 index 000000000000..5de04e8288c8 --- /dev/null +++ b/tests/test-mtp-catchup-batch.cpp @@ -0,0 +1,23 @@ +// Host-only bounds for draft-mtp's first decode: deferred catch-up plus one +// anchor per drafting sequence must fit in llama_n_batch(ctx_dft). +#include "speculative.h" + +#undef NDEBUG +#include +#include + +int main() { + // full-prefill boundary: n_max=31, n_batch=32, one 32-token prompt stash + assert(common_speculative_mtp_first_decode_fits(32, 31, 1)); + assert(!common_speculative_mtp_first_decode_fits(32, 32, 1)); + + // two sequences, combined catch-up plus two anchors + assert(common_speculative_mtp_first_decode_fits(32, 15 + 15, 2)); + assert(!common_speculative_mtp_first_decode_fits(32, 16 + 16, 2)); + + assert(!common_speculative_mtp_first_decode_fits(0, 0, 0)); + assert(!common_speculative_mtp_first_decode_fits(32, -1, 1)); + + printf("ok\n"); + return 0; +} diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp index a9edbd7be8b4..b063810acb18 100644 --- a/tools/server/server-context.cpp +++ b/tools/server/server-context.cpp @@ -205,6 +205,7 @@ struct server_slot { // speculative decoding common_speculative * spec; + int32_t spec_depth_max = 0; // no drafting once the sequence is longer than this (0 = always draft) llama_tokens spec_draft; llama_tokens spec_prompt; @@ -441,6 +442,12 @@ struct server_slot { return 0; } + // deep in the context the draft passes and the multi-column verify cost more than the + // accepted tokens save (see --spec-draft-depth-max); decode one token per step from here + if (spec_depth_max > 0 && prompt.n_tokens() > spec_depth_max) { + return 0; + } + // determine the max draft that fits the current slot state // note: slot.prompt is not yet expanded with the `id` token sampled above // also, need to leave space for 1 extra token to allow context shifts @@ -1210,6 +1217,7 @@ struct server_context_impl { slot.ctx_dft = ctx_dft; slot.mem.init(ctx_tgt, ctx_dft); slot.spec = spec.get(); + slot.spec_depth_max = params_base.speculative.draft.n_depth_max; slot.n_ctx = n_ctx_slot; slot.mctx = mctx; @@ -3625,7 +3633,19 @@ struct server_context_impl { // TODO: avoid restoring the draft context and re-evaluating the drafted tokens when not needed [TAG_SPEC_AVOID_DRAFT_REEVAL] // for now, always re-evaluate for simplicity // ref: https://github.com/ggml-org/llama.cpp/pull/22728#issuecomment-4400925384 - if (spec) { + // past --spec-draft-depth-max a slot does not draft, so its draft context has no use for + // these rows: skip the hook (and the draft-model decode it implies) when every token of this + // view sits beyond the cutoff. Positions, not the slot's token count, decide it, so the early + // ubatches of a long prompt still reach the draft context and stay reusable as a prefix. + bool spec_process = spec != nullptr; + if (spec_process && params_base.speculative.draft.n_depth_max > 0) { + spec_process = false; + for (int i = 0; i < batch_view.n_tokens && !spec_process; ++i) { + spec_process = batch_view.pos[i] <= params_base.speculative.draft.n_depth_max; + } + } + + if (spec_process) { bool ok = true; queue_tasks.yield_to_queue([&]() { ok = common_speculative_process(spec.get(), batch_view);