diff --git a/include/s2_model.h b/include/s2_model.h index eb36a6e..29bf056 100644 --- a/include/s2_model.h +++ b/include/s2_model.h @@ -6,6 +6,7 @@ #include "ggml-backend.h" #include "ggml-cpu.h" #include "gguf.h" +#include "s2_sampler.h" #ifdef GGML_USE_VULKAN #include "ggml-vulkan.h" #endif @@ -123,11 +124,18 @@ class SlowARModel { bool step(const std::vector & flat_tokens, int32_t n_threads, StepResult & result); - bool fast_decode(const std::vector & hidden, - const std::vector & prefix_codes, + bool fast_decode(const std::vector & hidden_in, + const std::vector & prefix_tokens, int32_t n_threads, std::vector & logits_out); + bool fast_decode_batch( + const std::vector & hidden_in, + int32_t semantic_code, + int32_t n_threads, + const SamplerParams & sparams, + std::vector & codebooks_out); + const ModelHParams & hparams() const { return hparams_; } ModelWeights weights_; @@ -138,6 +146,7 @@ class SlowARModel { ggml_backend_t backend_gpu_ = nullptr; ggml_backend_sched_t sched_ = nullptr; ggml_backend_sched_t fast_sched_ = nullptr; + ggml_gallocr_t fast_gallocr_ = nullptr; ggml_context * ctx_kv_ = nullptr; ggml_backend_buffer_t kv_buf_ = nullptr; ggml_tensor * memory_k_ = nullptr; @@ -152,6 +161,27 @@ class SlowARModel { std::unordered_set weight_tensor_set_; + bool fast_decoder_cpu_ = false; + + struct FastGraphSlot { + ggml_context * ctx = nullptr; + ggml_cgraph * gf = nullptr; + ggml_tensor * hidden0 = nullptr; + ggml_tensor * prefix_ids = nullptr; + ggml_tensor * positions = nullptr; + ggml_tensor * logits = nullptr; + ggml_tensor * logits_all = nullptr; + ggml_gallocr_t allocr = nullptr; + int32_t n_tokens = 0; + bool valid = false; + }; + + FastGraphSlot fast_slot_; + std::vector fast_slot_buf_; + size_t fast_slot_buf_size_ = 0; + + std::vector fast_slots_; + }; } diff --git a/src/s2_generate.cpp b/src/s2_generate.cpp index 6ee4c83..cf04681 100755 --- a/src/s2_generate.cpp +++ b/src/s2_generate.cpp @@ -19,34 +19,31 @@ GenerateResult generate( out.num_codebooks = model.hparams().num_codebooks; if (out.num_codebooks <= 0) out.num_codebooks = 1; - const int32_t vocab_size = model.hparams().vocab_size; - const int32_t sem_begin = model.hparams().semantic_begin_id; - const int32_t sem_end = model.hparams().semantic_end_id; + const int32_t vocab_size = model.hparams().vocab_size; + const int32_t sem_begin = model.hparams().semantic_begin_id; + const int32_t sem_end = model.hparams().semantic_end_id; const int32_t codebook_size = model.hparams().codebook_size; - const int32_t im_end_id = config.im_end_id; - const int32_t num_cb = out.num_codebooks; + const int32_t im_end_id = config.im_end_id; + const int32_t num_cb = out.num_codebooks; std::vector sem_mask(vocab_size, -std::numeric_limits::infinity()); - for (int32_t i = sem_begin; i <= sem_end && i < vocab_size; ++i) { + for (int32_t i = sem_begin; i <= sem_end && i < vocab_size; ++i) sem_mask[i] = 0.0f; - } - if (im_end_id >= 0 && im_end_id < vocab_size) { + if (im_end_id >= 0 && im_end_id < vocab_size) sem_mask[im_end_id] = 0.0f; - } const int32_t rows = prompt.rows; const int32_t cols = prompt.cols; std::vector prompt_tm(static_cast(rows) * cols); - for (int32_t r = 0; r < rows; ++r) { - for (int32_t c = 0; c < cols; ++c) { - prompt_tm[static_cast(c) * rows + r] = prompt.data[static_cast(r) * cols + c]; - } - } + for (int32_t r = 0; r < rows; ++r) + for (int32_t c = 0; c < cols; ++c) + prompt_tm[static_cast(c) * rows + r] = + prompt.data[static_cast(r) * cols + c]; StepResult state; - if (params.verbose && log_enabled(LogLevel::Info)) { + if (params.verbose && log_enabled(LogLevel::Info)) std::cout << "[Generate] Prefilling " << prompt.cols << " tokens..." << std::endl; - } + const auto prefill_t0 = std::chrono::steady_clock::now(); if (!model.prefill_fast(prompt_tm, prompt.cols, params.n_threads, state)) { std::cerr << "[Generate] Prefill failed." << std::endl; @@ -54,19 +51,18 @@ GenerateResult generate( } const auto prefill_t1 = std::chrono::steady_clock::now(); + SamplerParams sparams; + sparams.temperature = params.temperature; + sparams.top_p = params.top_p; + sparams.top_k = params.top_k; + auto apply_mask_and_sample = [&](const std::vector & logits, bool block_im_end) -> int32_t { std::vector biased(vocab_size); - for (int32_t i = 0; i < vocab_size; ++i) { + for (int32_t i = 0; i < vocab_size; ++i) biased[i] = logits[i] + sem_mask[i]; - } - if (block_im_end && im_end_id >= 0 && im_end_id < vocab_size) { + if (block_im_end && im_end_id >= 0 && im_end_id < vocab_size) biased[im_end_id] = -std::numeric_limits::infinity(); - } - SamplerParams sparams; - sparams.temperature = params.temperature; - sparams.top_p = params.top_p; - sparams.top_k = params.top_k; return sample_token(biased.data(), vocab_size, sparams); }; @@ -76,97 +72,89 @@ GenerateResult generate( out.codes.resize(static_cast(num_cb) * params.max_new_tokens, 0); out.n_frames = 0; - std::vector fast_logits; - - SamplerParams sparams; - sparams.temperature = params.temperature; - sparams.top_p = params.top_p; - sparams.top_k = params.top_k; - - std::vector ras_window; + int32_t ras_buf[10]; + int32_t ras_head = 0; + int32_t ras_count = 0; const int32_t ras_window_size = 10; - const float ras_high_temp = 1.0f; - const float ras_high_top_p = 0.9f; + const float ras_high_temp = 1.0f; + const float ras_high_top_p = 0.9f; - if (params.verbose && log_enabled(LogLevel::Info)) { + if (params.verbose && log_enabled(LogLevel::Info)) std::cout << "[Generate] Generating (max " << params.max_new_tokens << " tokens)..." << std::endl; - } int32_t step = 0; const auto loop_t0 = std::chrono::steady_clock::now(); + while (main_token != im_end_id && step < params.max_new_tokens) { - if (!ras_window.empty() && - std::find(ras_window.begin(), ras_window.end(), main_token) != ras_window.end() && - main_token >= sem_begin && main_token <= sem_end) + bool ras_hit = false; + if (main_token >= sem_begin && main_token <= sem_end) { + for (int32_t i = 0; i < ras_count; ++i) { + if (ras_buf[i] == main_token) { ras_hit = true; break; } + } + } + if (ras_hit) { - std::vector biased(vocab_size); - for (int32_t i = 0; i < vocab_size; ++i) { + for (int32_t i = 0; i < vocab_size; ++i) biased[i] = state.logits[i] + sem_mask[i]; - } - if (step < params.min_tokens_before_end && im_end_id >= 0 && im_end_id < vocab_size) { + if (step < params.min_tokens_before_end && + im_end_id >= 0 && im_end_id < vocab_size) biased[im_end_id] = -std::numeric_limits::infinity(); - } - SamplerParams ras_sparams; - ras_sparams.temperature = ras_high_temp; - ras_sparams.top_p = ras_high_top_p; - ras_sparams.top_k = params.top_k; - main_token = sample_token(biased.data(), vocab_size, ras_sparams); + SamplerParams ras_sp; + ras_sp.temperature = ras_high_temp; + ras_sp.top_p = ras_high_top_p; + ras_sp.top_k = params.top_k; + main_token = sample_token(biased.data(), vocab_size, ras_sp); } - ras_window.push_back(main_token); - if ((int32_t)ras_window.size() > ras_window_size) { - ras_window.erase(ras_window.begin()); - } + ras_buf[ras_head] = main_token; + ras_head = (ras_head + 1) % ras_window_size; + if (ras_count < ras_window_size) ras_count++; int32_t sem_code = main_token - sem_begin; - if (sem_code < 0) sem_code = 0; + if (sem_code < 0) sem_code = 0; if (sem_code >= codebook_size) sem_code = codebook_size - 1; + std::vector residual_cbs; + if (!model.fast_decode_batch(state.hidden, sem_code, + params.n_threads, sparams, + residual_cbs)) + { + std::cerr << "[Generate] fast_decode_batch failed at step " + << step << std::endl; + residual_cbs.assign(num_cb - 1, 0); + } + std::vector codebooks_cb; codebooks_cb.reserve(num_cb); codebooks_cb.push_back(sem_code); + for (int32_t r : residual_cbs) + codebooks_cb.push_back(r); - for (int32_t cb_idx = 1; cb_idx < num_cb; ++cb_idx) { - - if (!model.fast_decode(state.hidden, codebooks_cb, params.n_threads, fast_logits)) { - std::cerr << "[Generate] fast_decode failed at cb " << cb_idx << std::endl; - - for (int32_t r = cb_idx; r < num_cb; ++r) { - codebooks_cb.push_back(0); - } - break; - } - int32_t cb_token = sample_token(fast_logits.data(), (int32_t)fast_logits.size(), sparams); - codebooks_cb.push_back(cb_token); - } - - for (int32_t cb = 0; cb < num_cb; ++cb) { - out.codes[static_cast(cb) * params.max_new_tokens + step] = codebooks_cb[cb]; - } + for (int32_t cb = 0; cb < num_cb; ++cb) + out.codes[static_cast(cb) * params.max_new_tokens + step] = + codebooks_cb[cb]; out.n_frames++; if (params.on_frame) { FrameCallbackData fcd; - fcd.codes = codebooks_cb.data(); - fcd.frame_index = step; - fcd.total_frames = out.n_frames; + fcd.codes = codebooks_cb.data(); + fcd.frame_index = step; + fcd.total_frames = out.n_frames; fcd.num_codebooks = num_cb; if (!params.on_frame(fcd)) { - - if (params.verbose && log_enabled(LogLevel::Info)) { - std::cout << "\n[Generate] Aborted by callback at frame " << step << std::endl; - } + if (params.verbose && log_enabled(LogLevel::Info)) + std::cout << "\n[Generate] Aborted by callback at frame " + << step << std::endl; break; } } std::vector step_input(num_cb + 1); step_input[0] = main_token; - for (int32_t cb = 0; cb < num_cb; ++cb) { + for (int32_t cb = 0; cb < num_cb; ++cb) step_input[cb + 1] = codebooks_cb[cb]; - } if (!model.step(step_input, params.n_threads, state)) { std::cerr << "[Generate] step() failed at step " << step << std::endl; @@ -174,23 +162,28 @@ GenerateResult generate( } step++; - if (params.verbose && log_enabled(LogLevel::Info) && step % 50 == 0) { - std::cout << "\r[Generate] " << step << " / " << params.max_new_tokens << " tokens..." << std::flush; - } + + if (params.verbose && log_enabled(LogLevel::Info) && step % 50 == 0) + std::cout << "\r[Generate] " << step << " / " + << params.max_new_tokens << " tokens..." << std::flush; bool block_next_end = (step < params.min_tokens_before_end); main_token = apply_mask_and_sample(state.logits, block_next_end); } if (params.verbose && log_enabled(LogLevel::Info)) { - const auto loop_t1 = std::chrono::steady_clock::now(); + const auto loop_t1 = std::chrono::steady_clock::now(); const auto generate_t1 = std::chrono::steady_clock::now(); - const double prefill_ms = std::chrono::duration(prefill_t1 - prefill_t0).count(); - const double loop_ms = std::chrono::duration(loop_t1 - loop_t0).count(); - const double total_ms = std::chrono::duration(generate_t1 - generate_t0).count(); + const double prefill_ms = std::chrono::duration( + prefill_t1 - prefill_t0).count(); + const double loop_ms = std::chrono::duration( + loop_t1 - loop_t0).count(); + const double total_ms = std::chrono::duration( + generate_t1 - generate_t0).count(); const double ms_per_frame = out.n_frames > 0 ? (loop_ms / out.n_frames) : 0.0; std::cout << std::endl; - std::cout << "[Generate] Done: " << out.n_frames << " frames generated." << std::endl; + std::cout << "[Generate] Done: " << out.n_frames + << " frames generated." << std::endl; std::cout << "[Metrics] Generate: prefill=" << prefill_ms << " ms, loop=" << loop_ms << " ms, total=" << total_ms @@ -201,12 +194,10 @@ GenerateResult generate( const int32_t n_frames = out.n_frames; if (n_frames < params.max_new_tokens) { std::vector compacted(static_cast(num_cb) * n_frames); - for (int32_t cb = 0; cb < num_cb; ++cb) { - for (int32_t t = 0; t < n_frames; ++t) { + for (int32_t cb = 0; cb < num_cb; ++cb) + for (int32_t t = 0; t < n_frames; ++t) compacted[static_cast(cb) * n_frames + t] = out.codes[static_cast(cb) * params.max_new_tokens + t]; - } - } out.codes = std::move(compacted); } else { out.codes.resize(static_cast(num_cb) * n_frames); diff --git a/src/s2_model.cpp b/src/s2_model.cpp index 9ced623..fa04e2b 100755 --- a/src/s2_model.cpp +++ b/src/s2_model.cpp @@ -183,21 +183,31 @@ static bool allocate_weight_buffers(ggml_backend_t backend, return true; } -SlowARModel::SlowARModel() {} +SlowARModel::SlowARModel() + : fast_gallocr_(nullptr) +{} SlowARModel::~SlowARModel() { + for (auto & slot : fast_slots_) { + if (slot.allocr) { ggml_gallocr_free(slot.allocr); slot.allocr = nullptr; } + if (slot.ctx) { ggml_free(slot.ctx); slot.ctx = nullptr; } + } + fast_slots_.clear(); + + if (fast_slot_.allocr) { ggml_gallocr_free(fast_slot_.allocr); fast_slot_.allocr = nullptr; } + if (fast_slot_.ctx) { ggml_free(fast_slot_.ctx); fast_slot_.ctx = nullptr; } + fast_slot_.valid = false; + + if (fast_gallocr_) { ggml_gallocr_free(fast_gallocr_); fast_gallocr_ = nullptr; } + if (fast_sched_) ggml_backend_sched_free(fast_sched_); if (sched_) ggml_backend_sched_free(sched_); - - if (kv_buf_) ggml_backend_buffer_free(kv_buf_); + if (kv_buf_) ggml_backend_buffer_free(kv_buf_); free_backend_buffers(weights_.model_bufs_gpu); free_backend_buffers(weights_.model_bufs_cpu); - if (backend_gpu_) ggml_backend_free(backend_gpu_); if (backend_cpu_) ggml_backend_free(backend_cpu_); - if (ctx_kv_) ggml_free(ctx_kv_); - weights_.ctx_w = nullptr; } @@ -1223,35 +1233,362 @@ bool SlowARModel::fast_decode(const std::vector & hidden_in, ggml_build_forward_expand(gf, logits); ggml_backend_cpu_set_n_threads(backend_cpu_, resolve_n_threads(n_threads)); - ggml_backend_sched_reset(fast_sched_); - if (!ggml_backend_sched_alloc_graph(fast_sched_, gf)) { - std::fprintf(stderr, "[fast_decode] sched alloc failed\n"); - ggml_backend_sched_reset(fast_sched_); - ggml_free(ctx0); - return false; - } + // If no GPU layers are offloaded, use gallocr cached buffer pool. + // If GPU offload is active, use the scheduler to safely handle PCIe transfers. + const bool use_cpu_path = fast_decoder_cpu_ || + (n_gpu_layers_ == 0 || backend_gpu_ == nullptr); + + if (use_cpu_path) { + if (!fast_gallocr_) { + fast_gallocr_ = ggml_gallocr_new(ggml_backend_get_default_buffer_type(backend_cpu_)); + if (!fast_gallocr_) { + std::fprintf(stderr, "[fast_decode] gallocr creation failed\n"); + ggml_free(ctx0); + return false; + } + } - ggml_backend_tensor_set(hidden0, hidden_in.data(), 0, hidden_in.size() * sizeof(float)); - ggml_backend_tensor_set(positions, pos_vals.data(), 0, pos_vals.size() * sizeof(int32_t)); - if (prefix_ids) { - ggml_backend_tensor_set(prefix_ids, prefix_tokens.data(), 0, - prefix_tokens.size() * sizeof(int32_t)); - } + if (!ggml_gallocr_alloc_graph(fast_gallocr_, gf)) { + std::fprintf(stderr, "[fast_decode] gallocr alloc failed\n"); + ggml_free(ctx0); + return false; + } - if (ggml_backend_sched_graph_compute(fast_sched_, gf) != GGML_STATUS_SUCCESS) { - std::fprintf(stderr, "[fast_decode] sched compute failed\n"); + ggml_backend_tensor_set(hidden0, hidden_in.data(), 0, hidden_in.size() * sizeof(float)); + ggml_backend_tensor_set(positions, pos_vals.data(), 0, pos_vals.size() * sizeof(int32_t)); + if (prefix_ids) { + ggml_backend_tensor_set(prefix_ids, prefix_tokens.data(), 0, prefix_tokens.size() * sizeof(int32_t)); + } + + if (ggml_backend_graph_compute(backend_cpu_, gf) != GGML_STATUS_SUCCESS) { + std::fprintf(stderr, "[fast_decode] cpu compute failed\n"); + ggml_free(ctx0); + return false; + } + } else { ggml_backend_sched_reset(fast_sched_); - ggml_free(ctx0); - return false; + + if (!ggml_backend_sched_alloc_graph(fast_sched_, gf)) { + std::fprintf(stderr, "[fast_decode] sched alloc failed\n"); + ggml_backend_sched_reset(fast_sched_); + ggml_free(ctx0); + return false; + } + + ggml_backend_tensor_set(hidden0, hidden_in.data(), 0, hidden_in.size() * sizeof(float)); + ggml_backend_tensor_set(positions, pos_vals.data(), 0, pos_vals.size() * sizeof(int32_t)); + if (prefix_ids) { + ggml_backend_tensor_set(prefix_ids, prefix_tokens.data(), 0, prefix_tokens.size() * sizeof(int32_t)); + } + + if (ggml_backend_sched_graph_compute(fast_sched_, gf) != GGML_STATUS_SUCCESS) { + std::fprintf(stderr, "[fast_decode] sched compute failed\n"); + ggml_backend_sched_reset(fast_sched_); + ggml_free(ctx0); + return false; + } } logits_out.resize(hparams_.codebook_size); ggml_backend_tensor_get(logits, logits_out.data(), 0, hparams_.codebook_size * sizeof(float)); - ggml_backend_sched_reset(fast_sched_); + if (!use_cpu_path) { + ggml_backend_sched_reset(fast_sched_); + } + ggml_free(ctx0); return true; } +bool SlowARModel::fast_decode_batch( + const std::vector & hidden_in, + int32_t semantic_code, + int32_t n_threads, + const SamplerParams & sparams, + std::vector & codebooks_out) +{ + if (!hparams_.has_fast_decoder) return false; + if (static_cast(hidden_in.size()) != hparams_.embedding_length) return false; + + const int32_t num_cb = hparams_.num_codebooks; // 10 + const int32_t n_residual = num_cb - 1; // 9 + const int32_t max_prefix = num_cb; // 10 prefix tokens + const int32_t max_n_tok = max_prefix + 1; // 11 + const int32_t cb_size = hparams_.codebook_size; + + const bool use_cpu_fast = fast_decoder_cpu_ || + (n_gpu_layers_ == 0 || backend_gpu_ == nullptr); + if (!use_cpu_fast) { + codebooks_out.clear(); + codebooks_out.reserve(n_residual); + std::vector prefix; + prefix.reserve(num_cb); + prefix.push_back(semantic_code); + std::vector logits; + for (int32_t cb = 1; cb < num_cb; ++cb) { + if (!fast_decode(hidden_in, prefix, n_threads, logits)) { + for (int32_t r = cb; r < num_cb; ++r) codebooks_out.push_back(0); + return false; + } + int32_t cb_token = sample_token( + logits.data(), static_cast(logits.size()), sparams); + codebooks_out.push_back(cb_token); + prefix.push_back(cb_token); + } + return true; + } + + const int32_t fast_dim = hparams_.fast_embedding_length; + const int32_t n_head = hparams_.fast_head_count; + const int32_t n_head_kv = hparams_.fast_head_count_kv; + const int32_t head_dim = (hparams_.fast_head_dim > 0) + ? hparams_.fast_head_dim + : fast_dim / n_head; + const int32_t q_size = n_head * head_dim; + const int32_t kv_size = n_head_kv * head_dim; + const float attn_scale = 1.0f / std::sqrt(static_cast(head_dim)); + + auto build_fast_body = [&](ggml_context * ctx0, + int32_t n_tok, + ggml_tensor * hidden0, + ggml_tensor * prefix_ids, + ggml_tensor * positions) -> ggml_tensor * + { + const int32_t n_prefix = n_tok - 1; + + ggml_tensor * projected = (weights_.fast_project_in != nullptr) + ? mul_mat_checked(ctx0, weights_.fast_project_in, hidden0, "mul_mat:fast_project_in") + : hidden0; + if (projected->type != GGML_TYPE_F32) + projected = ggml_cast(ctx0, projected, GGML_TYPE_F32); + + ggml_tensor * x = projected; + if (n_prefix > 0) { + ggml_tensor * prefix_emb = ggml_get_rows(ctx0, weights_.fast_embeddings, prefix_ids); + if (prefix_emb->type != GGML_TYPE_F32) + prefix_emb = ggml_cast(ctx0, prefix_emb, GGML_TYPE_F32); + x = ggml_concat(ctx0, x, prefix_emb, 1); + } + + for (int32_t il = 0; il < hparams_.fast_block_count; ++il) { + const auto & layer = weights_.fast_layers[il]; + + ggml_tensor * attn_in = rms_norm_weighted(ctx0, x, layer.attention_norm, hparams_.fast_rms_norm_eps); + ggml_tensor * qkv = mul_mat_checked(ctx0, layer.wqkv, attn_in, "mul_mat:fast_wqkv"); + const size_t es = ggml_element_size(qkv); + + ggml_tensor * q2d = ggml_view_2d(ctx0, qkv, q_size, n_tok, qkv->nb[1], 0); + ggml_tensor * k2d = ggml_view_2d(ctx0, qkv, kv_size, n_tok, qkv->nb[1], q_size * es); + ggml_tensor * v2d = ggml_view_2d(ctx0, qkv, kv_size, n_tok, qkv->nb[1], (q_size + kv_size) * es); + + ggml_tensor * q = ggml_reshape_3d(ctx0, ggml_cont(ctx0, q2d), head_dim, n_head, n_tok); + ggml_tensor * k = ggml_reshape_3d(ctx0, ggml_cont(ctx0, k2d), head_dim, n_head_kv, n_tok); + ggml_tensor * v = ggml_reshape_3d(ctx0, ggml_cont(ctx0, v2d), head_dim, n_head_kv, n_tok); + + if (hparams_.fast_attention_qk_norm) { + q = rms_norm_weighted(ctx0, q, layer.q_norm, hparams_.fast_rms_norm_eps); + k = rms_norm_weighted(ctx0, k, layer.k_norm, hparams_.fast_rms_norm_eps); + } + + q = ggml_rope_ext(ctx0, q, positions, nullptr, head_dim, 0, + hparams_.fast_context_length, hparams_.fast_rope_freq_base, + 1.0f, 0.0f, 1.0f, 1.0f, 1.0f); + k = ggml_rope_ext(ctx0, k, positions, nullptr, head_dim, 0, + hparams_.fast_context_length, hparams_.fast_rope_freq_base, + 1.0f, 0.0f, 1.0f, 1.0f, 1.0f); + + ggml_tensor * k_rep = repeat_interleave_heads(ctx0, k, n_head / n_head_kv); + ggml_tensor * v_rep = repeat_interleave_heads(ctx0, v, n_head / n_head_kv); + + ggml_tensor * Q = ggml_permute(ctx0, q, 0, 2, 1, 3); + ggml_tensor * K = ggml_permute(ctx0, k_rep, 0, 2, 1, 3); + ggml_tensor * KQ = mul_mat_checked(ctx0, K, Q, "mul_mat:fast_kq"); + ggml_tensor * KQs = ggml_scale(ctx0, KQ, attn_scale); + ggml_tensor * KQm = ggml_diag_mask_inf(ctx0, KQs, 0); + ggml_tensor * KQf = ggml_soft_max(ctx0, KQm); + + ggml_tensor * V = ggml_cont(ctx0, ggml_permute(ctx0, v_rep, 1, 2, 0, 3)); + ggml_tensor * KQV = mul_mat_checked(ctx0, V, KQf, "mul_mat:fast_kqv"); + ggml_tensor * KQVm = ggml_permute(ctx0, KQV, 0, 2, 1, 3); + ggml_tensor * attn_cur = ggml_cpy(ctx0, KQVm, + ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, q_size, n_tok)); + + ggml_tensor * attn_out = mul_mat_checked(ctx0, layer.wo, attn_cur, "mul_mat:fast_wo"); + ggml_tensor * h = ggml_add(ctx0, x, attn_out); + ggml_tensor * ff_in = rms_norm_weighted(ctx0, h, layer.ffn_norm, hparams_.fast_rms_norm_eps); + ggml_tensor * gate = mul_mat_checked(ctx0, layer.w1, ff_in, "mul_mat:fast_w1"); + ggml_tensor * up = mul_mat_checked(ctx0, layer.w3, ff_in, "mul_mat:fast_w3"); + ggml_tensor * ff_h = ggml_swiglu_split(ctx0, gate, up); + ggml_tensor * ff_out = mul_mat_checked(ctx0, layer.w2, ff_h, "mul_mat:fast_w2"); + x = ggml_add(ctx0, h, ff_out); + } + + return x; + }; + + if (!fast_slot_.valid) { + fast_slot_buf_size_ = 16u * 1024u * 1024u; + fast_slot_buf_.resize(fast_slot_buf_size_); + + ggml_init_params ip = { fast_slot_buf_size_, fast_slot_buf_.data(), true }; + fast_slot_.ctx = ggml_init(ip); + if (!fast_slot_.ctx) return false; + + ggml_context * ctx0 = fast_slot_.ctx; + ggml_cgraph * gf = ggml_new_graph_custom(ctx0, 32768, false); + + ggml_tensor * hidden0 = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, hparams_.embedding_length, 1); + ggml_tensor * prefix_ids = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, max_prefix); + ggml_tensor * positions = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, max_n_tok); + + ggml_tensor * x = build_fast_body(ctx0, max_n_tok, hidden0, prefix_ids, positions); + + ggml_tensor * out = rms_norm_weighted(ctx0, x, weights_.fast_norm, hparams_.fast_rms_norm_eps); + ggml_tensor * cont = ggml_cont(ctx0, out); + ggml_tensor * logits_all = mul_mat_checked(ctx0, weights_.fast_output, cont, "mul_mat:fast_logits_all"); + ggml_build_forward_expand(gf, logits_all); + + if (!fast_slot_.allocr) { + fast_slot_.allocr = ggml_gallocr_new( + ggml_backend_get_default_buffer_type(backend_cpu_)); + } + if (!fast_slot_.allocr || !ggml_gallocr_alloc_graph(fast_slot_.allocr, gf)) { + ggml_free(fast_slot_.ctx); + fast_slot_.ctx = nullptr; + return false; + } + + fast_slot_.gf = gf; + fast_slot_.hidden0 = hidden0; + fast_slot_.prefix_ids = prefix_ids; + fast_slot_.positions = positions; + fast_slot_.logits_all = logits_all; + fast_slot_.n_tokens = max_n_tok; + fast_slot_.valid = true; + } + + auto ensure_slot = [&](int32_t n_tok) -> FastGraphSlot * { + if (static_cast(fast_slots_.size()) <= n_tok) + fast_slots_.resize(n_tok + 1); + + auto & slot = fast_slots_[n_tok]; + if (slot.valid) return &slot; + + const int32_t n_prefix = n_tok - 1; + + ggml_init_params ip = { 8u * 1024u * 1024u, nullptr, true }; + slot.ctx = ggml_init(ip); + if (!slot.ctx) return nullptr; + + ggml_context * ctx0 = slot.ctx; + ggml_cgraph * gf = ggml_new_graph_custom(ctx0, 32768, false); + + ggml_tensor * hidden0 = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, hparams_.embedding_length, 1); + ggml_tensor * prefix_ids = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, n_prefix); + ggml_tensor * positions = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, n_tok); + + ggml_tensor * x = build_fast_body(ctx0, n_tok, hidden0, prefix_ids, positions); + + ggml_tensor * out = rms_norm_weighted(ctx0, x, weights_.fast_norm, hparams_.fast_rms_norm_eps); + ggml_tensor * cont = ggml_cont(ctx0, out); + ggml_tensor * last = ggml_cpy(ctx0, + last_token_view(ctx0, cont, n_tok), + ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, fast_dim, 1)); + ggml_tensor * logits = mul_mat_checked(ctx0, weights_.fast_output, last, "mul_mat:fast_logits"); + ggml_build_forward_expand(gf, logits); + + slot.allocr = ggml_gallocr_new( + ggml_backend_get_default_buffer_type(backend_cpu_)); + if (!slot.allocr || !ggml_gallocr_alloc_graph(slot.allocr, gf)) { + ggml_free(slot.ctx); + slot.ctx = nullptr; + return nullptr; + } + + slot.gf = gf; + slot.hidden0 = hidden0; + slot.prefix_ids = prefix_ids; + slot.positions = positions; + slot.logits = logits; + slot.n_tokens = n_tok; + slot.valid = true; + return &slot; + }; + + codebooks_out.resize(n_residual); + std::vector prefix_vals(max_prefix, 0); + prefix_vals[0] = semantic_code; + + ggml_backend_cpu_set_n_threads(backend_cpu_, resolve_n_threads(n_threads)); + + for (int32_t cb = 0; cb < n_residual; ++cb) { + const int32_t n_tok = cb + 2; // 2, 3, ..., 10 + const int32_t n_prefix = n_tok - 1; + const int32_t logit_pos = cb + 1; + + FastGraphSlot * slot = ensure_slot(n_tok); + + if (slot) { + std::vector pos_vals(n_tok); + for (int32_t i = 0; i < n_tok; ++i) pos_vals[i] = i; + + ggml_backend_tensor_set(slot->hidden0, + hidden_in.data(), 0, hidden_in.size() * sizeof(float)); + ggml_backend_tensor_set(slot->prefix_ids, + prefix_vals.data(), 0, n_prefix * sizeof(int32_t)); + ggml_backend_tensor_set(slot->positions, + pos_vals.data(), 0, n_tok * sizeof(int32_t)); + + if (ggml_backend_graph_compute(backend_cpu_, slot->gf) != GGML_STATUS_SUCCESS) + return false; + + std::vector logits(cb_size); + ggml_backend_tensor_get(slot->logits, + logits.data(), 0, cb_size * sizeof(float)); + + codebooks_out[cb] = sample_token(logits.data(), cb_size, sparams); + } else { + std::vector pos_vals(max_n_tok); + for (int32_t i = 0; i < max_n_tok; ++i) pos_vals[i] = i; + + ggml_backend_tensor_set(fast_slot_.hidden0, + hidden_in.data(), 0, hidden_in.size() * sizeof(float)); + ggml_backend_tensor_set(fast_slot_.prefix_ids, + prefix_vals.data(), 0, max_prefix * sizeof(int32_t)); + ggml_backend_tensor_set(fast_slot_.positions, + pos_vals.data(), 0, max_n_tok * sizeof(int32_t)); + + if (ggml_backend_graph_compute(backend_cpu_, fast_slot_.gf) != GGML_STATUS_SUCCESS) + return false; + + std::vector logits(cb_size); + const size_t byte_offset = + static_cast(logit_pos) * cb_size * sizeof(float); + ggml_backend_tensor_get(fast_slot_.logits_all, + logits.data(), byte_offset, cb_size * sizeof(float)); + + codebooks_out[cb] = sample_token(logits.data(), cb_size, sparams); + } + + if (cb + 1 < max_prefix) + prefix_vals[cb + 1] = codebooks_out[cb]; + } + + bool all_slots_ready = true; + for (int32_t n = 2; n <= max_n_tok; ++n) { + if (n >= static_cast(fast_slots_.size()) || !fast_slots_[n].valid) { + all_slots_ready = false; + break; + } + } + if (all_slots_ready && fast_slot_.valid) { + if (fast_slot_.allocr) { ggml_gallocr_free(fast_slot_.allocr); fast_slot_.allocr = nullptr; } + if (fast_slot_.ctx) { ggml_free(fast_slot_.ctx); fast_slot_.ctx = nullptr; } + fast_slot_.valid = false; + } + + return true; +} + }