File size: 49,919 Bytes
bbb6388 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 496 497 498 499 500 501 502 503 504 505 506 507 508 509 510 511 512 513 514 515 516 517 518 519 520 521 522 523 524 525 526 527 528 529 530 531 532 533 534 535 536 537 538 539 540 541 542 543 544 545 546 547 548 549 550 551 552 553 554 555 556 557 558 559 560 561 562 563 564 565 566 567 568 569 570 571 572 573 574 575 576 577 578 579 580 581 582 583 584 585 586 587 588 589 590 591 592 593 594 595 596 597 598 599 600 601 602 603 604 605 606 607 608 609 610 611 612 613 614 615 616 617 618 619 620 621 622 623 624 625 626 627 628 629 630 631 632 633 634 635 636 637 638 639 640 641 642 643 644 645 646 647 648 649 650 651 652 653 654 655 656 657 658 659 660 661 662 663 664 665 666 667 668 669 670 671 672 673 674 675 676 677 678 679 680 681 682 683 684 685 686 687 688 689 690 691 692 693 694 695 696 697 698 699 700 701 702 703 704 705 706 707 708 709 710 711 712 713 714 715 716 717 718 719 720 721 722 723 724 725 726 727 728 729 730 731 732 733 734 735 736 737 738 739 740 741 742 743 744 745 746 747 748 749 750 751 752 753 754 755 756 757 758 759 760 761 762 763 764 765 766 767 768 769 770 771 772 773 774 775 776 777 778 779 780 781 782 783 784 785 786 787 788 789 790 791 792 793 794 795 796 797 798 799 800 801 802 803 804 805 806 807 808 809 810 811 812 813 814 815 816 817 818 819 820 821 822 823 824 825 826 827 828 829 830 831 832 833 834 835 836 837 838 839 840 841 842 843 844 845 846 847 848 849 850 851 852 853 854 855 856 857 858 859 860 861 862 863 864 865 866 867 868 869 870 871 872 873 874 875 876 877 878 879 880 881 882 883 884 885 886 887 888 889 890 891 892 893 894 895 896 | // 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 <algorithm>
#include <atomic>
#include <chrono>
#include <cstdio>
#include <cstdlib>
#include <cstring>
#include <exception>
#include <fstream>
#include <sstream>
#include <vector>
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<double, std::milli>(Clock::now() - t).count(); }
struct Bump {
uint8_t* base = nullptr;
uint64_t used = 0;
template <typename T> 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<uint8_t>& 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<uint8_t> 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<uint8_t> chunk(64u << 20);
for (uint64_t off = 0; off < bytes;) {
const uint64_t n = std::min<uint64_t>(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<int32_t>(T); step_ = b.take<int32_t>(R2 * 4); pos_ = b.take<int32_t>(R2 * NH); row_ = b.take<int32_t>(4);
ident_ = b.take<int32_t>(T * (uint64_t) cap_);
Rin_ = b.take<float>(T * HC * N); R_ = b.take<float>(T * HC * N);
emb_ = b.take<float>(T * N); en_ = b.take<float>(T * N); e2_ = b.take<float>(T * N);
hn_ = b.take<float>(T * HC * N); h2_ = b.take<float>(T * HC * N);
mixed_ = b.take<float>(T * N); inj_ = b.take<float>(T * HC); inj2_ = b.take<float>(T * HC);
lo_ = b.take<float>(T * (uint64_t) g.hc_lr); rs_ = b.take<float>(T * HC); bo_ = b.take<float>(T * N);
xn_ = b.take<float>(T * HC * N);
xq_ = b.take<uint8_t>(strata::kernels::native_q8_1_bytes((int) (NH * HD), 8));
qfull_ = b.take<float>(T * NH * 2 * HD); qcur_ = b.take<float>(T * NH * HD);
kcur_ = b.take<float>(T * NKV * HD); vcur_ = b.take<float>(T * NKV * HD);
attn_ = b.take<float>(T * NH * HD); attn32_ = b.take<float>(T * NH * HD);
attn_scratch_ = b.take<float>((uint64_t) attn_scratch_floats_); // the full layer runs one row at a time
logits_ = b.take<float>(T * (uint64_t) g.n_expert); w_ = b.take<float>(T * K); ids_ = b.take<int32_t>(T * K);
shared_ = b.take<float>(T * N); parts_ = b.take<float>(T * K * N); y_ = b.take<float>(T * N);
sample_ = b.take<float>(T * N);
hit_slot_ = b.take<int32_t>(T * K); hit_dst_ = b.take<int32_t>(T * K); hit_count_ = b.take<int32_t>(4);
grp_ptr_ = b.take<unsigned long long>(T * K); grp_start_ = b.take<int32_t>(T * K + 1);
grp_counts_ = b.take<int32_t>(4);
hit_xq_ = b.take<uint8_t>(T * (N / 32) * 34); hit_xs_ = b.take<float>(T * (N / 32));
hit_scratch_ = b.take<uint8_t>(strata::kernels::moe_hit_grouped_scratch_bytes((int64_t) (T * K), g.n_embd, g.n_ff));
sh_scratch_ = (float*) b.take<uint8_t>(strata::kernels::shared_expert_scratch_bytes(g.n_ff));
x_bf16_ = b.take<uint16_t>(N);
out_ids_ = b.take<int32_t>(T + 4);
probs_ = b.take<float>(T + 4);
dummy_inj_ = b.take<float>(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<int32_t> 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<double>(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<int32_t> 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<uint8_t> 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<int64_t>(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<int32_t*>(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<int64_t>(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<int32_t> 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<int64_t>(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<int64_t>(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<int32_t> toks((size_t) T, token);
return draft(T, toks.data(), cell, 0, drafts, err, probs, min_p, n_drafts);
}
} // namespace strata::core
|