// src/core/mtp.cpp - see include/strata/core/mtp.hpp. #include "strata/core/mtp.hpp" #include "strata/core/coupled_draft.hpp" #include "strata/core/on_device.hpp" #include "strata/core/native_head.hpp" #include "strata/core/peer_experts.hpp" #include "strata/kernels/bf16_gemv.hpp" #include "strata/kernels/cpu/expert.hpp" #include "strata/kernels/elementwise.hpp" #include "strata/kernels/fused_gr.hpp" #include "strata/kernels/gr.hpp" #include "strata/kernels/kv_q4.hpp" #include "strata/kernels/kv_q8.hpp" #include "strata/kernels/native_moe.hpp" #include "strata/kernels/native_mmvq.hpp" #include "strata/kernels/native_qsa.hpp" #include "strata/kernels/native_rope.hpp" #include "strata/kernels/native_router.hpp" #include "strata/kernels/qsa.hpp" #include "strata/kernels/qsa_decode_attn.hpp" #include "strata/kernels/quantize_act.hpp" #include "strata/kernels/rope.hpp" #include "strata/kernels/router_top10.hpp" #include "strata/kernels/s2_expert_grouped.hpp" #include "strata/kernels/sampler.hpp" #include "strata/kernels/shared_expert.hpp" #include "strata/kernels/verify_kernels.hpp" #include #include #include #include #include #include #include #include #include #include namespace strata::core { namespace { constexpr float EPS = 1e-6f; constexpr int GGML_Q8_0 = 8; using Clock = std::chrono::steady_clock; double ms_since(Clock::time_point t) { return std::chrono::duration(Clock::now() - t).count(); } struct Bump { uint8_t* base = nullptr; uint64_t used = 0; template T* take(uint64_t n) { T* p = base ? (T*) (base + used) : nullptr; used += (n * sizeof(T) + 255) & ~255ull; return p; } }; bool mapped(size_t bytes, void** h, void** d) { if (cudaHostAlloc(h, bytes, cudaHostAllocMapped | (peer_portable() ? cudaHostAllocPortable : 0)) != cudaSuccess) return false; std::memset(*h, 0, bytes); return cudaHostGetDevicePointer(d, *h, 0) == cudaSuccess; } strata::kernels::QsaShapes shapes_of(const ModelGeometry& g) { strata::kernels::QsaShapes s = strata::kernels::qsa_real_shapes(); s.n_head = g.n_head; s.n_head_kv = g.n_head_kv; s.head_dim = g.head_dim; s.idx_n_head = g.idx_q_heads; s.idx_dim = g.idx_key_dim; return s; } // 64-bit seek/tell on a `FILE*`: `fseek`/`ftell` take a 32-bit `long` on Windows and would wrap past 2 GiB. #if defined(_WIN32) #define STRATA_FILE_SEEK64(f, o, w) _fseeki64((f), (long long) (o), (w)) #define STRATA_FILE_TELL64(f) _ftelli64(f) #else #define STRATA_FILE_SEEK64(f, o, w) fseeko((f), (off_t) (o), (w)) #define STRATA_FILE_TELL64(f) ftello(f) #endif bool read_file(const std::string& path, std::vector& out) { // Loader fix (0.1.15+loaderfix.2): `ifstream::read` reaches the disk as 4095-byte reads on MSVC - the same // split `load_experts_ranges` had - so the drafter's 111 MiB dense blob paid 4 KiB per operation on a cold // start. `fread` passes a request bigger than the stream buffer straight to `ReadFile()`. FILE* f = std::fopen(path.c_str(), "rb"); if (f == nullptr) return false; // A `FILE*` has no destructor and the short-read path below returns early: the guard closes on every path. struct Closer { FILE* f; ~Closer() { if (f != nullptr) std::fclose(f); } } closer{f}; if (STRATA_FILE_SEEK64(f, 0, SEEK_END) != 0) return false; const long long n = (long long) STRATA_FILE_TELL64(f); if (n < 0 || STRATA_FILE_SEEK64(f, 0, SEEK_SET) != 0) return false; out.resize((size_t) n); return n == 0 || std::fread(out.data(), 1, (size_t) n, f) == (size_t) n; } } // namespace MtpDrafter::~MtpDrafter() { if (cs_) cudaStreamSynchronize(cs_); for (auto& e : prefill_exec_) if (e) cudaGraphExecDestroy(e); for (auto& e : prefill_dev_exec_) if (e) cudaGraphExecDestroy(e); if (pf_dev_) cudaFree(pf_dev_); for (auto& e : round_exec_) if (e) cudaGraphExecDestroy(e); for (auto& e : step_exec_) if (e) cudaGraphExecDestroy(e); for (auto& e : round_exec_c_) if (e) cudaGraphExecDestroy(e); for (auto& e : step_exec_c_) if (e) cudaGraphExecDestroy(e); if (cparams_) cudaFree(cparams_); if (cring_) cudaFree(cring_); if (dinv_) cudaFree(dinv_); if (cscratch_) cudaFree(cscratch_); if (h_cparams_) cudaFreeHost(h_cparams_); if (h_chist_) cudaFreeHost(h_chist_); if (cs_) cudaStreamDestroy(cs_); if (dense_) cudaFree(dense_); if (experts_) cudaFree(experts_); if (state_arena_) cudaFree(state_arena_); if (arena_) cudaFree(arena_); if (head_logits_) cudaFree(head_logits_); if (dhead_) cudaFree(dhead_); if (dvocab_) cudaFree(dvocab_); void* hosts[] = {h_tok_, h_step_, h_pos_, h_row_, h_out_, h_prob_}; for (void* h : hosts) if (h) cudaFreeHost(h); } const float* MtpDrafter::f32(const char* name) const { for (const auto& t : tensors_) if (t.name == name && t.kind == "f32") return (const float*) (dense_ + t.off); return nullptr; } const uint16_t* MtpDrafter::bf16(const char* name) const { for (const auto& t : tensors_) if (t.name == name && t.kind == "bf16") return (const uint16_t*) (dense_ + t.off); return nullptr; } const void* MtpDrafter::q8(const char* name) const { for (const auto& t : tensors_) if (t.name == name && t.kind == "q8_0") return dense_ + t.off; return nullptr; } bool MtpDrafter::load(const std::string& rt_dir, const ModelGeometry& g, SessionState& ss, int max_t, std::string& err, int64_t window) { cudaGetDevice(&device_); // a layer split's last stage on another GPU: the drafter lives there g_ = &g; ss_ = &ss; max_t_ = max_t; rt_dir_ = rt_dir; if (max_t < 1 || max_t > strata::kernels::kVerifyMaxT) { err = "mtp: max_t out of range"; return false; } // Loader fix (0.1.15+loaderfix.2): the two reads below are the whole “drafter files” cost; reporting // them apart from the rest of the stage is what makes the next regression visible. const auto t_files = std::chrono::steady_clock::now(); // ---- the index and the dense weights { std::ifstream idx(rt_dir + "/dense.txt"); if (!idx) { err = "mtp: cannot open " + rt_dir + "/dense.txt (run tools/mtp_rt.py)"; return false; } std::string line; while (std::getline(idx, line)) { if (line.empty()) continue; std::istringstream is(line); Tensor t; is >> t.name >> t.kind >> t.rows >> t.cols >> t.off >> t.bytes; if (!is) { err = "mtp: malformed dense.txt line: " + line; return false; } tensors_.push_back(t); } std::vector blob; if (!read_file(rt_dir + "/dense.bin", blob)) { err = "mtp: cannot read dense.bin"; return false; } const cudaError_t alloc = cudaMalloc((void**) &dense_, blob.size()); if (alloc != cudaSuccess) { size_t free_bytes = 0, total_bytes = 0; const cudaError_t info = cudaMemGetInfo(&free_bytes, &total_bytes); err = "mtp: dense weights allocation failed (" + std::string(cudaGetErrorString(alloc)) + "), requested " + std::to_string(blob.size() >> 20) + " MiB, CUDA0 free " + (info == cudaSuccess ? std::to_string(free_bytes >> 20) + " MiB" : "unknown"); return false; } cudaMemcpy(dense_, blob.data(), blob.size(), cudaMemcpyHostToDevice); vram_ += blob.size(); } // ---- the 512 routed experts, one blob each { const uint64_t bytes = (uint64_t) g.n_expert * strata::kernels::cpu::BLOB; // Loader fix (0.1.15+loaderfix.2): each 64 MiB read below reached the disk as ~16k 4095-byte reads under // MSVC's `basic_filebuf::xsgetn`, which is what made 675 MiB of drafter experts take minutes. FILE* f = std::fopen((rt_dir + "/experts.bin").c_str(), "rb"); if (f == nullptr) { err = "mtp: cannot open experts.bin"; return false; } struct Closer { FILE* f; ~Closer() { if (f != nullptr) std::fclose(f); } } closer{f}; if (cudaMalloc((void**) &experts_, bytes) != cudaSuccess) { err = "mtp: the 512 experts do not fit in VRAM"; return false; } std::vector chunk(64u << 20); for (uint64_t off = 0; off < bytes;) { const uint64_t n = std::min(chunk.size(), bytes - off); if (std::fread(chunk.data(), 1, (size_t) n, f) != (size_t) n) { err = "mtp: experts.bin is truncated"; return false; } cudaMemcpy(experts_ + off, chunk.data(), n, cudaMemcpyHostToDevice); off += n; } vram_ += bytes; } const char* required[] = {"fc_embedding.weight", "fc_hidden.weight", "self_attn.q_proj.weight", "self_attn.k_proj.weight", "self_attn.v_proj.weight", "self_attn.o_proj.weight", "mlp.shared_expert.gate_proj.weight", "mlp.shared_expert.up_proj.weight", "mlp.shared_expert.down_proj.weight"}; for (const char* n : required) if (!q8(n)) { err = std::string("mtp: ") + n + " is missing (q8_0)"; return false; } // ---- the layer's own K/V (dense attention: no indexer state is read) const strata::kernels::QsaShapes s = shapes_of(g); const int64_t max_cells = ss.qsa_states[ss.qsa_primary()].max_cells; // KV streaming: the drafter only reads its last `window` cells, so with streaming on its K/V is a ring of the // window (plus the cells a round writes ahead of its queries) over a host copy, refilled on a resume. The host copy // is pinned after the expert arena has pinned what it could: if it does not fit, the K/V stays whole in VRAM. int64_t ring = (window > 0 && window < max_cells) ? window + 4 * (int64_t) max_t + 64 : 0; // K8V4 never applies to the drafter: its own attention paths (below, and verify.cpp) handle whole formats // only, whatever ring shape it takes (0, a window, or the -1 fully-resident fallback). const bool kv_hybrid_was = qsa_kv_hybrid(); const bool kv_int8_was = qsa_kv_int8(); qsa_set_kv_hybrid(false); if (kv_hybrid_was) qsa_set_kv_int8(true); // the drafter under --kv k8v4: plain INT8 uint64_t sb = qsa_state_bytes(g, max_cells, false, ring); if (cudaMalloc(&state_arena_, sb) != cudaSuccess) { err = "mtp: the K/V state does not fit"; return false; } if (qsa_state_init(g, max_cells, state_arena_, st_, &ss.qsa_states[ss.qsa_primary()], ring) == 0) { if (st_.kv_mode == 0) { err = "mtp: state init failed"; return false; } std::fprintf(stderr, "strata mtp: no pinned RAM left for the draft layer's K/V copy; keeping it in VRAM\n"); cudaGetLastError(); cudaFree(state_arena_); st_ = QsaState{}; ring = -1; // fully resident sb = qsa_state_bytes(g, max_cells, false, ring); if (cudaMalloc(&state_arena_, sb) != cudaSuccess) { err = "mtp: the K/V state does not fit"; return false; } if (qsa_state_init(g, max_cells, state_arena_, st_, &ss.qsa_states[ss.qsa_primary()], ring) == 0) { err = "mtp: state init failed"; return false; } } qsa_set_kv_int8(kv_int8_was); qsa_set_kv_hybrid(kv_hybrid_was); qsa_state_zero(st_, g, nullptr); cudaDeviceSynchronize(); vram_ += sb; // ---- buffers window_ = (window > 0 && window < max_cells) ? window : 0; cap_ = (((window_ > 0 ? window_ : max_cells) + 63) / 64) * 64; attn_scratch_floats_ = (int64_t) strata::kernels::qsa_decode_attn_scratch_floats(cap_, s); const uint64_t T = (uint64_t) max_t, N = (uint64_t) g.n_embd, HC = (uint64_t) g.hc, K = (uint64_t) ss.k; const uint64_t NH = (uint64_t) g.n_head, HD = (uint64_t) g.head_dim, NKV = (uint64_t) g.n_head_kv; const uint64_t R2 = 2 * T; // step/pos rows: T catch-up rows + up to T-2 chain steps bool ok = mapped(T * 4 + 64, (void**) &h_tok_, (void**) &m_tok_) && mapped(R2 * 4 * 4 + 64, (void**) &h_step_, (void**) &m_step_) && mapped(R2 * NH * 4 + 64, (void**) &h_pos_, (void**) &m_pos_) && mapped(64, (void**) &h_row_, (void**) &m_row_) && mapped(T * 4 + 64, (void**) &h_out_, (void**) &m_out_) && mapped(T * 4 + 64, (void**) &h_prob_, (void**) &m_prob_); if (!ok) { err = "mtp: mapped staging failed"; return false; } auto carve = [&](Bump& b) { tok_ = b.take(T); step_ = b.take(R2 * 4); pos_ = b.take(R2 * NH); row_ = b.take(4); ident_ = b.take(T * (uint64_t) cap_); Rin_ = b.take(T * HC * N); R_ = b.take(T * HC * N); emb_ = b.take(T * N); en_ = b.take(T * N); e2_ = b.take(T * N); hn_ = b.take(T * HC * N); h2_ = b.take(T * HC * N); mixed_ = b.take(T * N); inj_ = b.take(T * HC); inj2_ = b.take(T * HC); lo_ = b.take(T * (uint64_t) g.hc_lr); rs_ = b.take(T * HC); bo_ = b.take(T * N); xn_ = b.take(T * HC * N); xq_ = b.take(strata::kernels::native_q8_1_bytes((int) (NH * HD), 8)); qfull_ = b.take(T * NH * 2 * HD); qcur_ = b.take(T * NH * HD); kcur_ = b.take(T * NKV * HD); vcur_ = b.take(T * NKV * HD); attn_ = b.take(T * NH * HD); attn32_ = b.take(T * NH * HD); attn_scratch_ = b.take((uint64_t) attn_scratch_floats_); // the full layer runs one row at a time logits_ = b.take(T * (uint64_t) g.n_expert); w_ = b.take(T * K); ids_ = b.take(T * K); shared_ = b.take(T * N); parts_ = b.take(T * K * N); y_ = b.take(T * N); sample_ = b.take(T * N); hit_slot_ = b.take(T * K); hit_dst_ = b.take(T * K); hit_count_ = b.take(4); grp_ptr_ = b.take(T * K); grp_start_ = b.take(T * K + 1); grp_counts_ = b.take(4); hit_xq_ = b.take(T * (N / 32) * 34); hit_xs_ = b.take(T * (N / 32)); hit_scratch_ = b.take(strata::kernels::moe_hit_grouped_scratch_bytes((int64_t) (T * K), g.n_embd, g.n_ff)); sh_scratch_ = (float*) b.take(strata::kernels::shared_expert_scratch_bytes(g.n_ff)); x_bf16_ = b.take(N); out_ids_ = b.take(T + 4); probs_ = b.take(T + 4); dummy_inj_ = b.take(HC); }; Bump count; carve(count); if (cudaMalloc(&arena_, count.used) != cudaSuccess) { err = "mtp: buffers do not fit"; return false; } cudaMemset(arena_, 0, count.used); Bump real; real.base = (uint8_t*) arena_; carve(real); vram_ += count.used; { std::vector id((size_t) (T * (uint64_t) cap_)); for (uint64_t t = 0; t < T; ++t) for (int64_t i = 0; i < cap_; ++i) id[(size_t) (t * (uint64_t) cap_ + (uint64_t) i)] = (int32_t) i; cudaMemcpy(ident_, id.data(), id.size() * 4, cudaMemcpyHostToDevice); } if (cudaStreamCreateWithFlags(&cs_, cudaStreamNonBlocking) != cudaSuccess) { err = "mtp: stream"; return false; } const double files_s = std::chrono::duration(std::chrono::steady_clock::now() - t_files).count(); std::fprintf(stderr, "strata mtp: draft layer loaded, %.0f MiB of VRAM (experts %.0f, dense %.0f), files read in %.2f s (%.0f MiB/s)\n", (double) vram_ / 1048576.0, (double) g.n_expert * strata::kernels::cpu::BLOB / 1048576.0, (double) tensors_.back().off / 1048576.0, files_s, files_s > 0 ? ((double) g.n_expert * strata::kernels::cpu::BLOB + (double) tensors_.back().off) / 1048576.0 / files_s : 0.0); return true; } uint64_t MtpDrafter::bind_bytes(uint64_t head_row_bytes, int64_t n_vocab) const { uint64_t bytes = head_logits_ ? 0 : (uint64_t) max_t_ * (uint64_t) n_vocab * sizeof(float); if (dhead_ == nullptr) { if (FILE* f = std::fopen((rt_dir_ + "/draft_vocab.bin").c_str(), "rb")) { std::fseek(f, 0, SEEK_END); const long size = std::ftell(f); std::fclose(f); if (size >= 4 && size % 4 == 0) bytes += (uint64_t) (size / 4) * head_row_bytes + (uint64_t) size; } } // coupled draft sampling (STRATA_SPEC_COUPLED=1 only): an upper bound - the id -> subset map, the penalty ring, // the split scratch (64 lists of 64 entries at most) and the parameters if (coupled_draft_env() && cparams_ == nullptr) bytes += (uint64_t) n_vocab * 4 + (uint64_t) (kCoupledHistCap + max_t_) * 4 + 64ull * 64ull * 8ull + 4096; return bytes; } bool MtpDrafter::setup_coupled(std::string& err) { const int64_t nv = dhead_ != nullptr ? n_dvocab_ : n_vocab_; const size_t scratch = strata::kernels::coupled_draft_scratch_bytes((int) nv); if (scratch == 0) { std::fprintf(stderr, "strata mtp: STRATA_SPEC_COUPLED: %lld draft logits are too wide for the coupled sampler; " "argmax drafts\n", (long long) nv); return true; } const size_t ring = (size_t) (kCoupledHistCap + max_t_) * sizeof(int32_t); if (cudaMalloc((void**) &cparams_, sizeof(strata::kernels::SamplerParams)) != cudaSuccess || cudaMalloc((void**) &cring_, ring) != cudaSuccess || cudaMalloc(&cscratch_, scratch) != cudaSuccess || !mapped(sizeof(strata::kernels::SamplerParams), (void**) &h_cparams_, (void**) &m_cparams_) || !mapped((size_t) kCoupledHistCap * sizeof(int32_t), (void**) &h_chist_, (void**) &m_chist_)) { err = "mtp: the coupled draft sampler's buffers do not fit"; return false; } cudaMemset(cring_, 0xff, ring); // -1: no token std::fill(h_chist_, h_chist_ + kCoupledHistCap, -1); vram_ += sizeof(strata::kernels::SamplerParams) + ring + scratch; if (dhead_ != nullptr) { // token id -> subset index (-1: not in the draft head), for the penalties std::vector sub((size_t) n_dvocab_), inv((size_t) n_vocab_, -1); cudaMemcpy(sub.data(), dvocab_, sub.size() * sizeof(int32_t), cudaMemcpyDeviceToHost); for (size_t i = 0; i < sub.size(); ++i) if (sub[i] >= 0 && sub[i] < n_vocab_ && inv[(size_t) sub[i]] < 0) inv[(size_t) sub[i]] = (int32_t) i; if (cudaMalloc((void**) &dinv_, inv.size() * sizeof(int32_t)) != cudaSuccess) { err = "mtp: the coupled draft sampler's token map does not fit"; return false; } cudaMemcpy(dinv_, inv.data(), inv.size() * sizeof(int32_t), cudaMemcpyHostToDevice); vram_ += inv.size() * sizeof(int32_t); } if (cudaDeviceSynchronize() != cudaSuccess) { err = std::string("mtp: coupled setup: ") + cudaGetErrorString(cudaGetLastError()); return false; } coupled_ok_ = true; std::fprintf(stderr, "strata mtp: coupled draft sampling on (STRATA_SPEC_COUPLED): sampled requests draft with the " "target's chain and Philox draw over %lld tokens\n", (long long) nv); return true; } void MtpDrafter::set_draft_sampling(const strata::kernels::SamplerParams& sp) { coupled_active_ = coupled_ok_ && !sp.greedy && sp.temperature > 0.0f; if (!coupled_active_) return; *h_cparams_ = sp; // read by the next round graph (after the previous one has synced) std::fill(h_chist_, h_chist_ + kCoupledHistCap, -1); } void MtpDrafter::set_draft_history(const int32_t* tail, int64_t n_tail, int32_t next) { if (!coupled_active_) return; const int h = coupled_hist_len(h_cparams_->penalty_last_n, kCoupledHistCap); coupled_hist_base(tail, n_tail, next, h, h_chist_ + (kCoupledHistCap - h)); } namespace { // #474: what to change when the draft head's token subset does not fit the VRAM left - a smaller subset that setup // ships (data/draft_vocab_*.bin; their token counts below), and how to pick it. The start used to stop at "the draft // head does not fit" with no hint. Text only, after the failure: nothing changes for a start that fits. void draft_head_hint(int64_t n_tokens, int64_t row_bytes) { size_t free_b = 0, total_b = 0; const bool have_free = cudaMemGetInfo(&free_b, &total_b) == cudaSuccess; auto mib = [&](int64_t n) { return (double) (n * row_bytes) / 1048576.0; }; std::string free_s; if (have_free) { char b[64]; std::snprintf(b, sizeof(b), " and %.0f MiB is free", (double) free_b / 1048576.0); free_s = b; } std::fprintf(stderr, "strata mtp: the draft head over %lld tokens needs %.0f MiB of VRAM%s.\n", (long long) n_tokens, mib(n_tokens), free_s.c_str()); struct Subset { const char* name; int64_t tokens; const char* what; }; static const Subset smaller[] = {{"cyrillic", 58963, "English, code and the Cyrillic script"}, {"en", 40525, "English and code"}}; std::string opts; for (const Subset& s : smaller) if (s.tokens < n_tokens) { char b[160]; std::snprintf(b, sizeof(b), "%s--draft-vocab %s (%s, ~%.0f MiB)", opts.empty() ? "" : " or ", s.name, s.what, mib(s.tokens)); opts += b; } if (!opts.empty()) std::fprintf(stderr, "strata mtp: hint: a smaller draft vocabulary needs less VRAM: %s. Start once with it - " "START-HERE.bat --draft-vocab en (Windows) or ./setup.sh --draft-vocab en - and the model " "keeps it (\"draft_vocab\" in its strata-*.json config); or a smaller --context in " "setup.\n", opts.c_str()); else std::fprintf(stderr, "strata mtp: hint: this is already the smallest shipped draft vocabulary: a smaller " "--context (or closing what else uses the GPU) leaves it room.\n"); } } // namespace bool MtpDrafter::bind(const WeightTable& wt, const NativeHead* head, const float* window_R, std::string& err) { const OnDevice on_device(device_); wt_ = &wt; head_ = head; window_R_ = window_R; const WeightRef* wo = wt.find("output.weight"); if (!wo) { err = "mtp: output.weight is missing"; return false; } n_vocab_ = wo->ne1; if (head == nullptr || !head->loaded()) { err = "mtp: the draft layer needs the native head (--native)"; return false; } if (head_logits_ == nullptr && cudaMalloc((void**) &head_logits_, (size_t) max_t_ * (size_t) n_vocab_ * sizeof(float)) != cudaSuccess) { err = "mtp: the draft logits do not fit"; return false; } // the draft head's token subset, when tools/draft_vocab.py wrote one if (dhead_ == nullptr) { std::vector raw; if (read_file(rt_dir_ + "/draft_vocab.bin", raw) && raw.size() >= 4 && raw.size() % 4 == 0) { n_dvocab_ = (int64_t) (raw.size() / 4); const int64_t row_bytes = (int64_t) head->row_bytes(); // a vocabulary row of the native head if (cudaMalloc((void**) &dvocab_, raw.size()) != cudaSuccess || cudaMalloc((void**) &dhead_, (size_t) (n_dvocab_ * row_bytes)) != cudaSuccess) { err = "mtp: the draft head does not fit"; draft_head_hint(n_dvocab_, row_bytes); return false; } cudaMemcpy(dvocab_, raw.data(), raw.size(), cudaMemcpyHostToDevice); strata::kernels::gather_rows((const uint8_t*) head->weights(), row_bytes, dvocab_, n_dvocab_, dhead_, nullptr); cudaDeviceSynchronize(); vram_ += (uint64_t) (n_dvocab_ * row_bytes) + raw.size(); std::fprintf(stderr, "strata mtp: draft head over %lld tokens (%.1f MiB)\n", (long long) n_dvocab_, (double) (n_dvocab_ * row_bytes) / 1048576.0); } } if (coupled_draft_env() && cparams_ == nullptr && !setup_coupled(err)) return false; return true; } // The layer for T rows. full = false stops after the K/V append (the prompt only needs the cache). bool MtpDrafter::record_forward(int T, int step_row0, cudaStream_t cs, std::string& err) { using namespace strata::kernels; const ModelGeometry& g = *g_; SessionState& ss = *ss_; // step_row0 >= 0: the full layer on step rows [step_row0, +T); step_row0 < 0: K/V only on rows [-1 - step_row0, +T) const bool full = step_row0 >= 0; const int row0 = full ? step_row0 : -1 - step_row0; if (full && T != 1) { err = "mtp: the full layer runs one row at a time (its attention scratch is sized for one)"; return false; } const int64_t N = g.n_embd, HC = g.hc, K = ss.k, NH = g.n_head, HD = g.head_dim, NKV = g.n_head_kv; const QsaShapes s = shapes_of(g); const GrShapes gs{g.n_embd, g.hc, g.hc_lr}; const int32_t* step = step_ + row0 * 4; const int32_t* pos = pos_ + row0 * NH; try { // ---- the two input branches const WeightRef* we = wt_->find("token_embd.weight"); if (!we) { err = "mtp: token_embd.weight is missing"; return false; } if (const NativeEmbed* ne = native_embed()) { // plan v0.3 P6: the GGUF-form table ne->gather_dev(tok_, T, emb_, cs); } else { const auto* codes = (const uint8_t*) we->data; const auto* scales = (const float*) (codes + we->codes_bytes); const auto* offsets = we->has_offset ? (const float*) (codes + we->codes_bytes + we->scales_bytes) : nullptr; embedding_gather_dev(codes, scales, offsets, tok_, T, we->ne0, we->code_bits, we->code_bias, we->group_elems, (uint64_t) (we->ne0 / (8 / we->code_bits)), (uint64_t) (we->ne0 / we->group_elems), emb_, cs); } native_qsa_rms_norm_weighted(emb_, f32("pre_fc_norm_embedding.weight"), en_, (int) N, T, EPS, cs); native_quantize_q8_1(en_, xq_, (int) N, T, cs); native_mmvq(GGML_Q8_0, q8("fc_embedding.weight"), xq_, e2_, (int) N, (int) N, T, cs); native_qsa_rms_norm_weighted(Rin_, f32("pre_fc_norm_hidden.weight"), hn_, (int) (HC * N), T, EPS, cs); for (int c0 = 0; c0 < T * HC; c0 += 8) { const int nc = (int) std::min(8, T * HC - c0); native_quantize_q8_1(hn_ + (size_t) c0 * N, xq_, (int) N, nc, cs); native_mmvq(GGML_Q8_0, q8("fc_hidden.weight"), xq_, h2_ + (size_t) c0 * N, (int) N, (int) N, nc, cs); } add_streams_broadcast(h2_, e2_, R_, N, (int) HC, T, cs); // ---- the attention hyper-connection { FusedGrArgs fa[kFusedGrMaxT]; for (int t = 0; t < T; ++t) { fa[t].R = R_ + (size_t) t * HC * N; fa[t].R_out = R_ + (size_t) t * HC * N; fa[t].apply = false; fa[t].w_norm = f32("attn_hyper_connection.hc_norm.weight"); fa[t].w_down = bf16("attn_hyper_connection.input_mix_weight_down.weight"); fa[t].w_up = bf16("attn_hyper_connection.input_mix_weight_up.weight"); fa[t].w_inject = bf16("attn_hyper_connection.block_inject_weight.weight"); fa[t].eps = EPS; fa[t].lo = lo_ + t * g.hc_lr; fa[t].rs = rs_ + t * HC; fa[t].inject_out = inj_ + t * HC; fa[t].mixed = mixed_ + t * N; } fused_gr_read_multi(fa, T, xn_, cs); } // ---- attention: K/V into the layer's own cache, then (full) dense attention over every cell auto norm_rope = [&](float* data, const float* gamma, int rows, int cols, const int32_t* p) { native_qsa_rms_norm_weighted(data, gamma, data, cols, rows, EPS, cs); if (native_rope_enabled()) native_rope_apply(data, data, rows, cols, (int) s.n_rot, rope_scaling(), p, cs); else rope_neox_apply(data, data, rows, cols, (int) s.n_rot, st_.cos_tab, st_.sin_tab, p, cs); }; native_quantize_q8_1(mixed_, xq_, (int) N, T, cs); native_mmvq(GGML_Q8_0, q8("self_attn.k_proj.weight"), xq_, kcur_, (int) N, (int) (NKV * HD), T, cs); native_mmvq(GGML_Q8_0, q8("self_attn.v_proj.weight"), xq_, vcur_, (int) N, (int) (NKV * HD), T, cs); for (int t = 0; t < T; ++t) { norm_rope(kcur_ + t * NKV * HD, f32("self_attn.k_norm.weight"), (int) NKV, (int) HD, pos + t * NH); if (st_.kv_rot) { // rotated K and V (kv_q4.hpp): Q4_0, and INT8 with STRATA_KV_ROT=1 fwht256_inplace_cuda(kcur_ + t * NKV * HD, NKV, cs); fwht256_inplace_cuda(vcur_ + t * NKV * HD, NKV, cs); } // stored in the state's own format (#293 appended rotated INT8 K/V as Q4_0, into pools INT8 never has) if (st_.kv_q4) kv_append_q4_step(st_.k_q4, st_.v_q4, st_.page_table, step + t * 4, kcur_ + t * NKV * HD, vcur_ + t * NKV * HD, s, cs, &st_.host); else if (st_.kv_int8) kv_append_q8_step(st_.k_q, st_.v_q, st_.k_scale, st_.v_scale, st_.page_table, step + t * 4, kcur_ + t * NKV * HD, vcur_ + t * NKV * HD, s, cs, &st_.host); else kv_append_step(st_.k_pool, st_.v_pool, st_.page_table, step + t * 4, kcur_ + t * NKV * HD, vcur_ + t * NKV * HD, s, cs, &st_.host); } if (!full) return true; native_mmvq(GGML_Q8_0, q8("self_attn.q_proj.weight"), xq_, qfull_, (int) N, (int) (NH * 2 * HD), T, cs); for (int t = 0; t < T; ++t) { float* qc = qcur_ + t * NH * HD; if (cudaMemcpy2DAsync(qc, (size_t) HD * 4, qfull_ + t * NH * 2 * HD, (size_t) HD * 2 * 4, (size_t) HD * 4, (size_t) NH, cudaMemcpyDeviceToDevice, cs) != cudaSuccess) { err = "mtp: q split failed"; return false; } norm_rope(qc, f32("self_attn.q_norm.weight"), (int) NH, (int) HD, pos + t * NH); if (st_.kv_rot) fwht256_inplace_cuda(qc, NH, cs); } const QsaAttnPools pools = qsa_attn_pools(st_); if (window_ > 0) window_ids(const_cast(step), T, (int) window_, ident_, cap_, cs); qsa_decode_attn_batch(qcur_, pools, ident_, step, cap_, s, attn_scratch_, attn_, T, cs); if (st_.kv_rot) fwht256_inplace_cuda(attn_, (int64_t) T * NH, cs); for (int t = 0; t < T; ++t) native_qsa_gate_apply(attn_ + t * NH * HD, qfull_ + t * NH * 2 * HD, attn32_ + t * NH * HD, (int) NH, (int) HD, cs); native_quantize_q8_1(attn32_, xq_, (int) (NH * HD), T, cs); native_mmvq(GGML_Q8_0, q8("self_attn.o_proj.weight"), xq_, bo_, (int) (NH * HD), (int) N, T, cs); // ---- the MLP hyper-connection (the attention write folded in) { FusedGrArgs fa[kFusedGrMaxT]; for (int t = 0; t < T; ++t) { fa[t].R = R_ + (size_t) t * HC * N; fa[t].R_out = R_ + (size_t) t * HC * N; fa[t].apply = true; fa[t].bo_prev = bo_ + t * N; fa[t].inj_prev = inj_ + t * HC; fa[t].w_norm = f32("mlp_hyper_connection.hc_norm.weight"); fa[t].w_down = bf16("mlp_hyper_connection.input_mix_weight_down.weight"); fa[t].w_up = bf16("mlp_hyper_connection.input_mix_weight_up.weight"); fa[t].w_inject = bf16("mlp_hyper_connection.block_inject_weight.weight"); fa[t].eps = EPS; fa[t].lo = lo_ + t * g.hc_lr; fa[t].rs = rs_ + t * HC; fa[t].inject_out = inj2_ + t * HC; fa[t].mixed = mixed_ + t * N; } fused_gr_read_multi(fa, T, xn_, cs); } // ---- MoE: router, the 512 resident experts, the shared expert, the combine, the write for (int t = 0; t < T; ++t) { bf16_gemv_fp32_mmvf(mixed_ + t * N, bf16("mlp.gate.weight"), logits_ + t * g.n_expert, (int) N, (int) g.n_expert, cs); if (native_router_enabled()) native_router_top10(logits_ + t * g.n_expert, ids_ + t * K, w_ + t * K, cs); else router_top10(logits_ + t * g.n_expert, 1, (int) g.n_expert, (int) K, ids_ + t * K, w_ + t * K, cs); } moe_group_resident(ids_, (int) (T * K), (int) K, experts_, (int64_t) strata::kernels::cpu::BLOB, grp_ptr_, grp_start_, grp_counts_, hit_dst_, hit_slot_, cs); quantize_q8_0_scaled(mixed_, hit_xq_, hit_xs_, (int64_t) T * N, cs); moe_grouped_s2(grp_ptr_, grp_start_, grp_counts_, hit_dst_, hit_slot_, (int64_t) T * K, (int64_t) T * K, hit_xq_, hit_xs_, hit_scratch_, parts_, cs); NativeSharedWeights nsw; nsw.gate_type = GGML_Q8_0; nsw.gate_data = q8("mlp.shared_expert.gate_proj.weight"); nsw.up_type = GGML_Q8_0; nsw.up_data = q8("mlp.shared_expert.up_proj.weight"); nsw.down_type = GGML_Q8_0; nsw.down_data = q8("mlp.shared_expert.down_proj.weight"); nsw.q8_1 = xq_; const SForm none{}; for (int t = 0; t < T; ++t) { f32_to_bf16_bulk(mixed_ + t * N, x_bf16_, N, cs); shared_expert(nullptr, nullptr, x_bf16_, none, nullptr, nullptr, nullptr, none, nullptr, nullptr, nullptr, none, nullptr, nullptr, nullptr, bf16("mlp.shared_expert_gate.weight"), sh_scratch_, shared_ + t * N, N, g.n_ff, 32, cs, mixed_ + t * N, &nsw); if (native_moe_combine_enabled()) native_moe_combine(parts_ + (size_t) t * K * N, w_ + t * K, shared_ + t * N, y_ + t * N, N, K, cs); else moe_combine(parts_ + (size_t) t * K * N, w_ + t * K, shared_ + t * N, y_ + t * N, N, K, cs); gr_write(R_ + (size_t) t * HC * N, y_ + t * N, inj2_ + t * HC, gs, R_ + (size_t) t * HC * N, cs); } // ---- the final mixer and the main model's head for (int t = 0; t < T; ++t) gr_read(R_ + (size_t) t * HC * N, f32("hyper_connection_mixer.hc_norm.weight"), bf16("hyper_connection_mixer.input_mix_weight_down.weight"), bf16("hyper_connection_mixer.input_mix_weight_up.weight"), nullptr, EPS, gs, ss.block.gr, sample_ + t * N, dummy_inj_, cs); native_quantize_q8_1(sample_, xq_, (int) N, T, cs); const bool sub = dhead_ != nullptr; const int64_t nv = sub ? n_dvocab_ : n_vocab_; native_mmvq(head_->type(), sub ? dhead_ : head_->weights(), xq_, head_logits_, (int) N, (int) nv, T, cs); if (coupled_rec_) { // coupled draft sampling: the target's chain and Philox draw for the row that will verify this draft // (counter = this cell + 1, from its step record), penalties over the ring; T == 1 (full layer) coupled_draft_sample(head_logits_, (int) nv, sub ? dvocab_ : nullptr, sub ? dinv_ : nullptr, (int) n_vocab_, cparams_, cring_, kCoupledHistCap, coupled_j_, step, cscratch_, out_ids_, probs_, cs); return true; } SamplerParams sp; sp.greedy = true; sp.temperature = 0.0f; sample_tokens(head_logits_, T, (int) nv, nullptr, 0, sp, out_ids_, cs); row_top_prob(head_logits_, T, (int) nv, out_ids_, probs_, cs); if (sub) map_ids(out_ids_, dvocab_, T, cs); } catch (const std::exception& e) { err = std::string("mtp: ") + e.what(); return false; } return true; } namespace { bool finish_capture(cudaStream_t cs, bool ok, cudaGraphExec_t& exec, const char* what, std::string& err) { cudaGraph_t graph = nullptr; const cudaError_t ce = cudaStreamEndCapture(cs, &graph); if (!ok) { if (graph) cudaGraphDestroy(graph); return false; } if (ce != cudaSuccess || cudaGraphInstantiate(&exec, graph, 0) != cudaSuccess) { if (graph) cudaGraphDestroy(graph); err = std::string("mtp: ") + what + " capture: " + cudaGetErrorString(ce); return false; } cudaGraphDestroy(graph); // an explicit upload: the first launch's implicit one blocked behind a device-side spin (verify.cpp) cudaGraphUpload(exec, cs); cudaStreamSynchronize(cs); return true; } } // namespace bool MtpDrafter::capture_prefill(int T, std::string& err) { if (prefill_exec_[T]) return true; using namespace strata::kernels; if (cudaStreamBeginCapture(cs_, cudaStreamCaptureModeThreadLocal) != cudaSuccess) { err = "mtp: begin capture"; return false; } copy_i32_from_mapped(tok_, m_tok_, T, cs_); copy_i32_from_mapped(step_, m_step_, (int64_t) T * 4, cs_); copy_i32_from_mapped(pos_, m_pos_, (int64_t) T * g_->n_head, cs_); const bool ok = record_forward(T, -1, cs_, err); // K/V only, rows [0, T) return finish_capture(cs_, ok, prefill_exec_[T], "prefill", err); } bool MtpDrafter::capture_prefill_dev(int T, std::string& err) { if (prefill_dev_exec_[T]) return true; if (cudaStreamBeginCapture(cs_, cudaStreamCaptureModeThreadLocal) != cudaSuccess) { err = "mtp: begin capture"; return false; } const bool ok = record_forward(T, -1, cs_, err); // K/V only, rows [0, T); tok_/step_/pos_ filled before launch return finish_capture(cs_, ok, prefill_dev_exec_[T], "prefill (device inputs)", err); } bool MtpDrafter::capture_round(int T, bool coupled, std::string& err) { cudaGraphExec_t& exec = coupled ? round_exec_c_[T] : round_exec_[T]; if (exec) return true; using namespace strata::kernels; const int64_t HCN = g_->hc * g_->n_embd; if (cudaStreamBeginCapture(cs_, cudaStreamCaptureModeThreadLocal) != cudaSuccess) { err = "mtp: begin capture"; return false; } bool ok = true; // coupled: the request's chain and the penalty history's base, for this round's drafts if (coupled) coupled_draft_stage(m_cparams_, m_chist_, cparams_, cring_, kCoupledHistCap, cs_); copy_i32_from_mapped(tok_, m_tok_, T, cs_); copy_i32_from_mapped(step_, m_step_, (int64_t) 2 * T * 4, cs_); copy_i32_from_mapped(pos_, m_pos_, (int64_t) 2 * T * g_->n_head, cs_); copy_i32_from_mapped(row_, m_row_, 2, cs_); copy_from_mapped(Rin_, window_R_, (int64_t) T * HCN, cs_); // the catch-up: K/V for the window's T cells, then the full layer for row a only (its cell's K/V is written // again, identically), staged by the host in step row 2*max_t - 1; the draft chain is one graph per step // (`capture_step`) so the host can stop it when a draft is unlikely const int ra = 2 * max_t_ - 1; ok = record_forward(T, -1, cs_, err); if (ok) mtp_select(Rin_, HCN, tok_, row_, Rin_, tok_, nullptr, 0, cs_); if (ok) { copy_i32_from_mapped(step_ + ra * 4, m_step_ + ra * 4, 4, cs_); copy_i32_from_mapped(pos_ + ra * g_->n_head, m_pos_ + ra * g_->n_head, g_->n_head, cs_); coupled_rec_ = coupled; coupled_j_ = 0; ok = record_forward(1, ra, cs_, err); coupled_rec_ = false; } if (ok) mtp_select(R_, HCN, out_ids_, row_ + 1, Rin_, tok_, m_out_, 0, cs_, probs_, m_prob_); return finish_capture(cs_, ok, exec, coupled ? "round (coupled)" : "round", err); } // Chain step j (1..max_t-2): one row at the cell staged in step row `max_t + j - 1`, from the previous step's // residual and token (left in Rin_[0] / tok_[0] by mtp_select); draft j and its probability to the mapped outputs. bool MtpDrafter::capture_step(int j, bool coupled, std::string& err) { cudaGraphExec_t& exec = coupled ? step_exec_c_[j] : step_exec_[j]; if (exec) return true; using namespace strata::kernels; const int64_t HCN = g_->hc * g_->n_embd; const int row = max_t_ + j - 1; if (cudaStreamBeginCapture(cs_, cudaStreamCaptureModeThreadLocal) != cudaSuccess) { err = "mtp: begin capture"; return false; } copy_i32_from_mapped(step_ + row * 4, m_step_ + row * 4, 4, cs_); copy_i32_from_mapped(pos_ + row * g_->n_head, m_pos_ + row * g_->n_head, g_->n_head, cs_); coupled_rec_ = coupled; coupled_j_ = j; bool ok = record_forward(1, row, cs_, err); coupled_rec_ = false; if (ok) mtp_select(R_, HCN, out_ids_, row_ + 1, Rin_, tok_, m_out_, j, cs_, probs_, m_prob_); return finish_capture(cs_, ok, exec, coupled ? "step (coupled)" : "step", err); } void MtpDrafter::kv_restore(int64_t upto) { const OnDevice on_device(device_); if (st_.kv_mode != 2 || upto <= 0) return; // the ring's blocks below `upto`, from the host copy: a checkpoint resume may have left later cells in them const strata::kernels::QsaShapes s = shapes_of(*g_); const int64_t b1 = (upto + s.page_size - 1) / s.page_size, b0 = std::max(0, b1 - st_.n_slots); strata::kernels::kv_ring_restore(qsa_attn_pools(st_), st_.host, qsa_kv_format(st_), b0, b1, st_.n_slots, s, cs_); cudaStreamSynchronize(cs_); } bool MtpDrafter::prefill(const float* R_rows, const int32_t* next_tokens, int64_t n, int64_t cell0, std::string& err) { const OnDevice on_device(device_); const Clock::time_point t0 = Clock::now(); const int64_t HCN = g_->hc * g_->n_embd; // cells the window can never reach again need no K/V const int64_t first_needed = (window_ > 0 && prompt_len_ > 0) ? prompt_len_ - window_ - 64 : 0; // E-4: every group's token / step / position records uploaded at once; each group is then device copies and a // graph on the one stream, with a single sync at the end (a group of <= max_t rows used to be staged in mapped // memory and synced before the next: ~5,500 host round trips on a 32K prompt). The same work in the same // order. STRATA_MTP_PREFILL_SYNC=1 keeps the old loop, =0 forces E-4. // HIP defaults to the old loop. This pass runs whenever E-9 (Prefill::draft_kv) declines, which it does for the // drafter's ring (KV streaming, --kv-resident), and on gfx1201 / ROCm 7.2.4 the E-4 queue (a graph launch and four // device copies per group, ~2,000 groups per 8192-token chunk, no sync) sometimes never completes: the prompt hangs // in hipGraphLaunch or the final sync until the watchdog ends the engine. R9700, full IQ3_XXS, --kv-resident // 32768, mixed load (chats, 32K and 7K prompts, deep follow-ups): E-4 hung in the first round on 2 of 2 tries, // the old loop ran 30 of 30 rounds clean, prompt speed unchanged. static const bool per_group_sync = [] { if (const char* v = std::getenv("STRATA_MTP_PREFILL_SYNC")) return std::atoi(v) != 0; #if defined(STRATA_USE_HIP) return true; #else return false; #endif }(); const int64_t NHp = g_->n_head, per_row = 1 + 4 + NHp; if (!per_group_sync && n > 0) { if (pf_cap_ < n * per_row) { if (pf_dev_) cudaFree(pf_dev_); pf_dev_ = nullptr; pf_cap_ = 0; if (cudaMalloc((void**) &pf_dev_, (size_t) (n * per_row) * sizeof(int32_t)) != cudaSuccess) { err = "mtp prefill: the input records do not fit"; return false; } pf_cap_ = n * per_row; } std::vector rec((size_t) (n * per_row)); int32_t* tk = rec.data(); int32_t* stp = tk + n; int32_t* ps = stp + 4 * n; for (int64_t i = 0; i < n; ++i) { const int64_t cell = cell0 + i; tk[i] = next_tokens[i]; stp[i * 4 + 0] = (int32_t) cell; stp[i * 4 + 1] = (int32_t) (cell + 1); stp[i * 4 + 2] = (int32_t) ((cell + 1) / 4); stp[i * 4 + 3] = (int32_t) (cell + 1); for (int64_t h = 0; h < NHp; ++h) ps[i * NHp + h] = (int32_t) cell; } if (cudaMemcpyAsync(pf_dev_, rec.data(), rec.size() * sizeof(int32_t), cudaMemcpyHostToDevice, cs_) != cudaSuccess) { err = "mtp prefill: the input records upload failed"; return false; } const int32_t* d_tk = pf_dev_; const int32_t* d_stp = d_tk + n; const int32_t* d_ps = d_stp + 4 * n; for (int64_t c = 0; c < n; c += max_t_) { const int T = (int) std::min(max_t_, n - c); if (cell0 + c + T <= first_needed) continue; if (!capture_prefill_dev(T, err)) return false; if (cudaMemcpyAsync(tok_, d_tk + c, (size_t) T * 4, cudaMemcpyDeviceToDevice, cs_) != cudaSuccess || cudaMemcpyAsync(step_, d_stp + c * 4, (size_t) T * 16, cudaMemcpyDeviceToDevice, cs_) != cudaSuccess || cudaMemcpyAsync(pos_, d_ps + c * NHp, (size_t) (T * NHp) * 4, cudaMemcpyDeviceToDevice, cs_) != cudaSuccess || cudaMemcpyAsync(Rin_, R_rows + (size_t) c * HCN, (size_t) T * HCN * sizeof(float), cudaMemcpyDeviceToDevice, cs_) != cudaSuccess || cudaGraphLaunch(prefill_dev_exec_[T], cs_) != cudaSuccess) { err = std::string("mtp prefill: ") + cudaGetErrorString(cudaGetLastError()); return false; } } if (cudaStreamSynchronize(cs_) != cudaSuccess) { err = std::string("mtp prefill: ") + cudaGetErrorString(cudaGetLastError()); return false; } ms_prefill += ms_since(t0); return true; } for (int64_t c = 0; c < n; c += max_t_) { const int T = (int) std::min(max_t_, n - c); if (cell0 + c + T <= first_needed) continue; if (!capture_prefill(T, err)) return false; for (int t = 0; t < T; ++t) { const int64_t cell = cell0 + c + t; h_tok_[t] = next_tokens[c + t]; h_step_[t * 4 + 0] = (int32_t) cell; h_step_[t * 4 + 1] = (int32_t) (cell + 1); h_step_[t * 4 + 2] = (int32_t) ((cell + 1) / 4); h_step_[t * 4 + 3] = (int32_t) (cell + 1); for (int64_t h = 0; h < g_->n_head; ++h) h_pos_[t * g_->n_head + h] = (int32_t) cell; } if (cudaMemcpyAsync(Rin_, R_rows + (size_t) c * HCN, (size_t) T * HCN * sizeof(float), cudaMemcpyDeviceToDevice, cs_) != cudaSuccess || cudaGraphLaunch(prefill_exec_[T], cs_) != cudaSuccess || cudaStreamSynchronize(cs_) != cudaSuccess) { err = std::string("mtp prefill: ") + cudaGetErrorString(cudaGetLastError()); return false; } } ms_prefill += ms_since(t0); return true; } bool MtpDrafter::draft(int T, const int32_t* tokens, int64_t p, int a, int32_t* drafts, std::string& err, float* probs, float min_p, int* n_drafts) { const OnDevice on_device(device_); if (T < 1 || T > max_t_ || a < 0 || a >= T) { err = "mtp: draft arguments out of range"; return false; } const bool cp = coupled_active_; // coupled draft sampling for this request: its own graphs if (!capture_round(T, cp, err)) return false; const Clock::time_point t0 = Clock::now(); const int64_t NH = g_->n_head; auto put = [&](int row, int64_t cell) { h_step_[row * 4 + 0] = (int32_t) cell; h_step_[row * 4 + 1] = (int32_t) (cell + 1); h_step_[row * 4 + 2] = (int32_t) ((cell + 1) / 4); h_step_[row * 4 + 3] = (int32_t) (cell + 1); for (int64_t h = 0; h < NH; ++h) h_pos_[row * NH + h] = (int32_t) cell; }; for (int t = 0; t < T; ++t) { h_tok_[t] = tokens[t]; put(t, p + t); } put(2 * max_t_ - 1, coupled_draft_cell(p, a, 0)); // p + a: draft 0's cell h_row_[0] = a; h_row_[1] = 0; std::atomic_thread_fence(std::memory_order_seq_cst); if (cudaGraphLaunch(cp ? round_exec_c_[T] : round_exec_[T], cs_) != cudaSuccess || cudaStreamSynchronize(cs_) != cudaSuccess) { err = std::string("mtp draft: ") + cudaGetErrorString(cudaGetLastError()); return false; } drafts[0] = ((volatile int32_t*) h_out_)[0]; float pj = ((volatile float*) h_prob_)[0]; if (probs) probs[0] = pj; int n = 1; // the chain continues while the last draft is likely enough to be verified for (int j = 1; j < std::min(max_t_ - 1, max_drafts_) && pj >= min_p; ++j) { if (!capture_step(j, cp, err)) return false; put(max_t_ + j - 1, coupled_draft_cell(p, a, j)); // p + a + j; coupled: drawn at counter cell + 1 std::atomic_thread_fence(std::memory_order_seq_cst); if (cudaGraphLaunch(cp ? step_exec_c_[j] : step_exec_[j], cs_) != cudaSuccess || cudaStreamSynchronize(cs_) != cudaSuccess) { err = std::string("mtp draft step: ") + cudaGetErrorString(cudaGetLastError()); return false; } drafts[j] = ((volatile int32_t*) h_out_)[j]; pj = ((volatile float*) h_prob_)[j]; if (probs) probs[j] = pj; ++n; } for (int j = n; j < max_t_ - 1; ++j) { drafts[j] = 0; if (probs) probs[j] = 0.0f; } if (n_drafts) *n_drafts = n; ms_draft += ms_since(t0); ++rounds; return true; } bool MtpDrafter::draft_first(int T, const float* R_row, int32_t token, int64_t cell, int32_t* drafts, std::string& err, float* probs, float min_p, int* n_drafts) { const OnDevice on_device(device_); // row 0 is the real pair; rows 1.. repeat it and only write cells the next round overwrites const int64_t HCN = g_->hc * g_->n_embd; for (int t = 0; t < T; ++t) if (cudaMemcpy((void*) (window_R_ + (size_t) t * HCN), R_row, (size_t) HCN * sizeof(float), cudaMemcpyDeviceToDevice) != cudaSuccess) { err = "mtp: staging the first residual failed"; return false; } std::vector toks((size_t) T, token); return draft(T, toks.data(), cell, 0, drafts, err, probs, min_p, n_drafts); } } // namespace strata::core