diff --git a/python/freetoken/engine/config.py b/python/freetoken/engine/config.py index 543012f39..bcbe6bcf2 100644 --- a/python/freetoken/engine/config.py +++ b/python/freetoken/engine/config.py @@ -23,6 +23,8 @@ class EngineConfig: moe_backend: str = "auto" # NVFP4 routed-expert GEMM backend (--nvfp4-backend): auto|marlin|flashinfer|triton. nvfp4_backend: str = "triton" + # PLE table backend: "disk" (default) reads rows from the checkpoint files per fill, "pinned" preloads the table into page-locked host RAM. + ple_backend: str = "disk" # Expert-bank host load (--expert-load): auto|serial|parallel. "auto" reads scattered # experts in parallel but falls back to serial when free RAM can't cover the banks + the # parallel reader's extra (non-reclaimable) whole-shard buffer; "serial" forces the diff --git a/python/freetoken/engine/engine.py b/python/freetoken/engine/engine.py index 73dc7688d..33b0d5495 100644 --- a/python/freetoken/engine/engine.py +++ b/python/freetoken/engine/engine.py @@ -917,11 +917,9 @@ def rebuild_runtime_cache( def forward_batch(self, batch: Batch, args: BatchSamplingArgs) -> ForwardOutput: assert torch.cuda.current_stream() == self.stream - with self.ctx.forward_batch(batch): - if self.graph_runner.can_use_cuda_graph(batch): - logits = self.graph_runner.replay(batch) - else: - logits = self.model.forward() + use_graph = self.graph_runner.can_use_cuda_graph(batch) + with self.ctx.forward_batch(batch), self.model.forward_host_ctx(batch, use_graph): + logits = self.graph_runner.replay(batch) if use_graph else self.model.forward() if self.cpu_moe_executor is not None: # One pinned read: surfaces a fired flag-handshake watchdog (dead coordinator # -> stale expert outputs) as a loud error instead of silent corruption. diff --git a/python/freetoken/kernel/csrc/ple_store/ple_store_ext.cpp b/python/freetoken/kernel/csrc/ple_store/ple_store_ext.cpp new file mode 100644 index 000000000..cddaad960 --- /dev/null +++ b/python/freetoken/kernel/csrc/ple_store/ple_store_ext.cpp @@ -0,0 +1,613 @@ +// Disk-backed PLE row store: rows read straight from the checkpoint's fp8 shard tensors +// through an extent table (PleRowSource in ple_ssd.py). Engine-thread only, no locks. +// Duplicate rows in one fill dedup into ONE batched read round; no RAM cache and no +// per-sequence state. Hash reference: tests/models/qwen4_exp/test_ple_disk.py. +// Platform seams: TableFile (O_DIRECT+pread; Win: NO_BUFFERING), BatchReader (io_uring, +// pread-pool fallback = the portable shape), cumemop_* (dlopen libcuda; Win: nvcuda). + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include +#include +#include +#include + +#if defined(__linux__) && __has_include() +#include +#include +#include +#define PLE_HAS_IO_URING 1 +#else +#define PLE_HAS_IO_URING 0 +#endif + +#include + +namespace py = pybind11; + +namespace { + +constexpr int64_t kPage = 4096; +constexpr int64_t kSpanMax = 2 * kPage; // a row is <= one page, so it spans at most two +constexpr unsigned kBatchEntries = 64; +// fio on this class of disk: pread saturates at ~16 threads; more only adds latency +constexpr unsigned kReaderThreads = 16; + +// ---- portability shims ---- + +void release_store_i64(int64_t *ptr, int64_t value) { + __atomic_store_n(ptr, value, __ATOMIC_RELEASE); +} + +uint8_t *page_aligned_alloc(size_t bytes) { + void *p = nullptr; + if (posix_memalign(&p, kPage, bytes) != 0) throw std::bad_alloc(); + return static_cast(p); +} + +// Stream memops for the flag-sync fast path, resolved from the driver at runtime. +using CuMemOp64Fn = int (*)(void *stream, unsigned long long addr, unsigned long long value, + unsigned int flags); +CuMemOp64Fn g_cu_write64 = nullptr; +CuMemOp64Fn g_cu_wait64 = nullptr; +constexpr unsigned kCuWaitValueGeq = 0x0; +constexpr unsigned kCuWriteDefault = 0x0; + +bool cumemop_resolve() { + static bool resolved = [] { + void *h = dlopen("libcuda.so.1", RTLD_LAZY | RTLD_LOCAL); + if (h == nullptr) h = dlopen("libcuda.so", RTLD_LAZY | RTLD_LOCAL); + if (h == nullptr) return false; + g_cu_write64 = reinterpret_cast(dlsym(h, "cuStreamWriteValue64_v2")); + if (g_cu_write64 == nullptr) + g_cu_write64 = reinterpret_cast(dlsym(h, "cuStreamWriteValue64")); + g_cu_wait64 = reinterpret_cast(dlsym(h, "cuStreamWaitValue64_v2")); + if (g_cu_wait64 == nullptr) + g_cu_wait64 = reinterpret_cast(dlsym(h, "cuStreamWaitValue64")); + return g_cu_write64 != nullptr && g_cu_wait64 != nullptr; + }(); + return resolved; +} + +int memop_write(uintptr_t stream, uintptr_t addr, uint64_t value) { + if (!cumemop_resolve()) return -1; + return g_cu_write64(reinterpret_cast(stream), addr, value, kCuWriteDefault); +} + +int memop_wait_geq(uintptr_t stream, uintptr_t addr, uint64_t value) { + if (!cumemop_resolve()) return -1; + return g_cu_wait64(reinterpret_cast(stream), addr, value, kCuWaitValueGeq); +} + +// WAIT(>=1) then RESET: resetting first would race a fast host signal and deadlock the stream. +void memop_wait_reset(uintptr_t stream, uintptr_t flag_addr) { + if (memop_wait_geq(stream, flag_addr, 1) != 0 || memop_write(stream, flag_addr, 0) != 0) + throw std::runtime_error("stream memops rejected in capture; set FREETOKEN_PLE_SYNC=gate"); +} + +void signal_flag(uintptr_t flag_addr) { + release_store_i64(reinterpret_cast(flag_addr), 1); +} + +// ---- TableFile: platform seam for on-disk files ---- + +// Read up to len (EOF may short it); fewer than need bytes is an error. +void pread_min(int fd, uint8_t *buf, int64_t len, int64_t need, int64_t off) { + int64_t done = 0; + while (done < len) { + ssize_t got = ::pread(fd, buf + done, len - done, off + done); + if (got < 0) { + if (errno == EINTR) continue; + throw std::runtime_error(std::string("pread: ") + std::strerror(errno)); + } + if (got == 0) break; + done += got; + } + if (done < need) + throw std::runtime_error("short read at offset " + std::to_string(off) + ": got " + + std::to_string(done) + " of " + std::to_string(need)); +} + +class TableFile { + public: + explicit TableFile(const std::string &path) { + fd_ = ::open(path.c_str(), O_RDONLY | O_CLOEXEC | O_DIRECT); + direct_ = fd_ >= 0; + if (fd_ < 0) { + fd_ = ::open(path.c_str(), O_RDONLY | O_CLOEXEC); + direct_ = false; + } + if (fd_ < 0) throw std::runtime_error(path + ": " + std::strerror(errno)); + struct stat st{}; + if (fstat(fd_, &st) != 0) { + ::close(fd_); + throw std::runtime_error(path + ": fstat: " + std::strerror(errno)); + } + size_ = st.st_size; + if (!direct_) posix_fadvise(fd_, 0, 0, POSIX_FADV_RANDOM); + } + + ~TableFile() { + if (fd_ >= 0) ::close(fd_); + } + + TableFile(const TableFile &) = delete; + TableFile &operator=(const TableFile &) = delete; + + bool direct_io() const { return direct_; } + int native_fd() const { return fd_; } + int64_t size() const { return size_; } + + // keep buffered fallback reads out of the page cache; a no-op under direct I/O + void discard_cache(int64_t off, int64_t len) const { + if (!direct_) posix_fadvise(fd_, off, len, POSIX_FADV_DONTNEED); + } + + private: + int fd_ = -1; + bool direct_ = false; + int64_t size_ = 0; +}; + +// ---- BatchReader: platform seam for parallel positioned reads ---- + +// Pipelined: at most capacity() reads in flight; wait_one() returns a finished tag to refill. +class BatchReader { + public: + virtual ~BatchReader() = default; + virtual const char *name() const = 0; + virtual unsigned capacity() const = 0; + virtual void submit(unsigned tag, int fd, uint8_t *buf, int64_t len, int64_t need, + int64_t off) = 0; + virtual unsigned wait_one() = 0; + // reap every in-flight read so stale completions cannot leak into the next fill + virtual void drain() noexcept = 0; +}; + +class ThreadPoolBatchReader final : public BatchReader { + public: + ThreadPoolBatchReader() { + const unsigned n = std::max(1u, std::min(kReaderThreads, std::thread::hardware_concurrency())); + for (unsigned i = 0; i < n; i++) workers_.emplace_back([this] { work(); }); + } + + ~ThreadPoolBatchReader() override { + { + std::lock_guard lock(mu_); + stop_ = true; + } + work_cv_.notify_all(); + for (auto &w : workers_) w.join(); + } + + const char *name() const override { return "pread-pool"; } + unsigned capacity() const override { return kBatchEntries; } + + void submit(unsigned tag, int fd, uint8_t *buf, int64_t len, int64_t need, + int64_t off) override { + { + std::lock_guard lock(mu_); + queue_.push_back(Req{tag, fd, buf, len, need, off}); + in_flight_++; + } + work_cv_.notify_one(); + } + + unsigned wait_one() override { + std::unique_lock lock(mu_); + done_cv_.wait(lock, [this] { return !done_.empty(); }); + Done d = std::move(done_.front()); + done_.pop_front(); + in_flight_--; + if (!d.error.empty()) throw std::runtime_error(d.error); + return d.tag; + } + + void drain() noexcept override { + // wait out all in-flight reads: a late worker write must not race the slot's reuse + std::unique_lock lock(mu_); + done_cv_.wait(lock, [this] { return done_.size() == in_flight_; }); + in_flight_ = 0; + done_.clear(); + } + + private: + struct Req { + unsigned tag; + int fd; + uint8_t *buf; + int64_t len; + int64_t need; + int64_t off; + }; + struct Done { + unsigned tag; + std::string error; + }; + + void work() { + for (;;) { + Req r; + { + std::unique_lock lock(mu_); + work_cv_.wait(lock, [this] { return stop_ || !queue_.empty(); }); + if (stop_) return; + r = queue_.front(); + queue_.pop_front(); + } + Done d{r.tag, {}}; + try { + pread_min(r.fd, r.buf, r.len, r.need, r.off); + } catch (const std::exception &e) { + d.error = e.what(); + } + { + std::lock_guard lock(mu_); + done_.push_back(std::move(d)); + } + done_cv_.notify_one(); + } + } + + std::vector workers_; + std::mutex mu_; + std::condition_variable work_cv_, done_cv_; + std::deque queue_; + std::deque done_; + size_t in_flight_ = 0; + bool stop_ = false; +}; + +#if PLE_HAS_IO_URING + +// Minimal single-issuer io_uring: submit up to `entries` reads, wait for all. +class IoUringBatchReader final : public BatchReader { + public: + IoUringBatchReader() = default; + + bool init(unsigned entries) { + struct io_uring_params p{}; + fd_ = (int)syscall(__NR_io_uring_setup, entries, &p); + if (fd_ < 0) return false; + sq_size_ = p.sq_off.array + p.sq_entries * sizeof(uint32_t); + cq_size_ = p.cq_off.cqes + p.cq_entries * sizeof(io_uring_cqe); + if (p.features & IORING_FEAT_SINGLE_MMAP) sq_size_ = cq_size_ = std::max(sq_size_, cq_size_); + sq_ptr_ = mmap(nullptr, sq_size_, PROT_READ | PROT_WRITE, MAP_SHARED | MAP_POPULATE, fd_, + IORING_OFF_SQ_RING); + if (sq_ptr_ == MAP_FAILED) return false; + cq_ptr_ = (p.features & IORING_FEAT_SINGLE_MMAP) + ? sq_ptr_ + : mmap(nullptr, cq_size_, PROT_READ | PROT_WRITE, MAP_SHARED | MAP_POPULATE, + fd_, IORING_OFF_CQ_RING); + if (cq_ptr_ == MAP_FAILED) return false; + sqes_size_ = p.sq_entries * sizeof(io_uring_sqe); + sqes_ = (io_uring_sqe *)mmap(nullptr, sqes_size_, PROT_READ | PROT_WRITE, + MAP_SHARED | MAP_POPULATE, fd_, IORING_OFF_SQES); + if (sqes_ == MAP_FAILED) return false; + + auto at = [&](void *base, uint32_t off) { return (uint8_t *)base + off; }; + sq_tail_ = (uint32_t *)at(sq_ptr_, p.sq_off.tail); + sq_mask_ = (uint32_t *)at(sq_ptr_, p.sq_off.ring_mask); + sq_array_ = (uint32_t *)at(sq_ptr_, p.sq_off.array); + cq_head_ = (uint32_t *)at(cq_ptr_, p.cq_off.head); + cq_tail_ = (uint32_t *)at(cq_ptr_, p.cq_off.tail); + cq_mask_ = (uint32_t *)at(cq_ptr_, p.cq_off.ring_mask); + cqes_ = (io_uring_cqe *)at(cq_ptr_, p.cq_off.cqes); + entries_ = p.sq_entries; + lens_.assign(entries_, 0); + sq_shadow_tail_ = *sq_tail_; + return true; + } + + ~IoUringBatchReader() override { + if (sqes_ && sqes_ != MAP_FAILED) munmap(sqes_, sqes_size_); + if (cq_ptr_ && cq_ptr_ != MAP_FAILED && cq_ptr_ != sq_ptr_) munmap(cq_ptr_, cq_size_); + if (sq_ptr_ && sq_ptr_ != MAP_FAILED) munmap(sq_ptr_, sq_size_); + if (fd_ >= 0) ::close(fd_); + } + + const char *name() const override { return "io_uring"; } + unsigned capacity() const override { return entries_; } + + void submit(unsigned tag, int fd, uint8_t *buf, int64_t len, int64_t need, + int64_t off) override { + io_uring_sqe *sqe = &sqes_[sq_shadow_tail_ & *sq_mask_]; + std::memset(sqe, 0, sizeof(*sqe)); + sqe->opcode = IORING_OP_READ; + sqe->fd = fd; + sqe->addr = (uint64_t)(uintptr_t)buf; + sqe->len = (uint32_t)len; + sqe->off = (uint64_t)off; + sqe->user_data = tag; + sq_array_[sq_shadow_tail_ & *sq_mask_] = sq_shadow_tail_ & *sq_mask_; + sq_shadow_tail_++; + __atomic_store_n(sq_tail_, sq_shadow_tail_, __ATOMIC_RELEASE); + lens_[tag] = need; + to_submit_++; + in_flight_++; + } + + unsigned wait_one() override { + for (;;) { + uint32_t head = *cq_head_; + uint32_t ctail = __atomic_load_n(cq_tail_, __ATOMIC_ACQUIRE); + if (head != ctail) { + const io_uring_cqe &cqe = cqes_[head & *cq_mask_]; + const unsigned tag = (unsigned)cqe.user_data; + const int res = cqe.res; + __atomic_store_n(cq_head_, head + 1, __ATOMIC_RELEASE); + in_flight_--; + if (res < 0) + throw std::runtime_error(std::string("io_uring read: ") + std::strerror(-res)); + if (res < lens_[tag]) throw std::runtime_error("io_uring short read"); + return tag; + } + const unsigned to_submit = to_submit_; + long rc = syscall(__NR_io_uring_enter, fd_, to_submit, 1, IORING_ENTER_GETEVENTS, nullptr, 0); + if (rc < 0) { + if (errno == EINTR) continue; + throw std::runtime_error(std::string("io_uring_enter: ") + std::strerror(errno)); + } + // partial submission is legal (signal, transient alloc); the rest stay in the ring + to_submit_ = to_submit - (unsigned)rc; + } + } + + void drain() noexcept override { + while (in_flight_ > 0) { + uint32_t head = *cq_head_; + uint32_t ctail = __atomic_load_n(cq_tail_, __ATOMIC_ACQUIRE); + if (head != ctail) { + __atomic_store_n(cq_head_, head + 1, __ATOMIC_RELEASE); + in_flight_--; + continue; + } + const unsigned to_submit = to_submit_; + long rc = syscall(__NR_io_uring_enter, fd_, to_submit, 1, IORING_ENTER_GETEVENTS, nullptr, 0); + if (rc < 0) { + if (errno != EINTR) return; + continue; + } + to_submit_ = to_submit - (unsigned)rc; + } + } + + private: + int fd_ = -1; + void *sq_ptr_ = nullptr, *cq_ptr_ = nullptr; + io_uring_sqe *sqes_ = nullptr; + size_t sq_size_ = 0, cq_size_ = 0, sqes_size_ = 0; + uint32_t *sq_tail_ = nullptr, *sq_mask_ = nullptr, *sq_array_ = nullptr; + uint32_t *cq_head_ = nullptr, *cq_tail_ = nullptr, *cq_mask_ = nullptr; + io_uring_cqe *cqes_ = nullptr; + unsigned entries_ = 0; + uint32_t sq_shadow_tail_ = 0; + unsigned to_submit_ = 0; + unsigned in_flight_ = 0; + std::vector lens_; +}; + +#endif // PLE_HAS_IO_URING + +std::unique_ptr make_batch_reader(bool use_io_uring) { +#if PLE_HAS_IO_URING + if (use_io_uring) { + auto ring = std::make_unique(); + if (ring->init(kBatchEntries)) return ring; + } +#else + (void)use_io_uring; +#endif + return std::make_unique(); +} + +// ---- row store (platform-free) ---- + +int64_t wrap_mul(int64_t a, int64_t b) { + return (int64_t)((uint64_t)a * (uint64_t)b); +} + +int64_t pos_mod(int64_t v, int64_t m) { + int64_t r = v % m; + return r < 0 ? r + m : r; +} + +class PleStore { + struct Extent { + const TableFile *file; + int64_t base; + }; + + public: + PleStore(std::vector paths, std::vector extent_file, + std::vector extent_base, int64_t rows_per_extent, int64_t row_bytes, + int64_t row_stride, std::vector multipliers, std::vector head_vocab_sizes, + std::vector head_offsets, int64_t eos_token_id, bool use_io_uring) + : row_bytes_(row_bytes), + row_stride_(row_stride), + rows_per_extent_(rows_per_extent), + mult_(std::move(multipliers)), + sizes_(std::move(head_vocab_sizes)), + offsets_(std::move(head_offsets)), + eos_(eos_token_id) { + if (mult_.size() != 3 || sizes_.size() != offsets_.size() || sizes_.empty()) + throw std::runtime_error("PLE hash geometry: want 3 multipliers and equal-length head tables"); + if (row_bytes_ > kPage) + throw std::runtime_error("PLE row_bytes " + std::to_string(row_bytes_) + + " exceeds a page; bounce slots assume one-page rows"); + for (const std::string &p : paths) + files_.push_back(std::make_unique(p)); + const int64_t extent_bytes = (rows_per_extent_ - 1) * row_stride_ + row_bytes_; + for (size_t e = 0; e < extent_file.size(); e++) { + const size_t fi = (size_t)extent_file.at(e); + const int64_t base = extent_base.at(e); + if (base + extent_bytes > files_.at(fi)->size()) + throw std::runtime_error(paths[fi] + ": extent needs " + + std::to_string(base + extent_bytes) + " bytes, file has " + + std::to_string(files_[fi]->size())); + extents_.push_back(Extent{files_[fi].get(), base}); + } + reader_ = make_batch_reader(use_io_uring); + bounce_ = page_aligned_alloc((size_t)reader_->capacity() * kSpanMax); + } + + // reader first: a still-running read must not land in freed bounce memory + ~PleStore() { + reader_.reset(); + free(bounce_); + } + + PleStore(const PleStore &) = delete; + PleStore &operator=(const PleStore &) = delete; + + // Row ids for the token at w[2] with context (w[0], w[1]); mirrors NGramEmbedding.row_ids incl. the eos barrier. + void hash_rows(const int64_t *w, int64_t *rows) { + const int64_t prev1 = w[1]; + const int64_t prev2 = prev1 == eos_ ? eos_ : w[0]; + const int64_t bigram = wrap_mul(w[2], mult_[0]) ^ wrap_mul(prev1, mult_[1]); + const int64_t trigram = bigram ^ wrap_mul(prev2, mult_[2]); + const size_t half = sizes_.size() / 2; + for (size_t h = 0; h < sizes_.size(); h++) + rows[h] = pos_mod(h < half ? bigram : trigram, sizes_[h]) + offsets_[h]; + } + + // Hash and queue one run: tokens_addr holds n+2 ids, the leading two are context. No I/O until flush(). + void stage(uintptr_t tokens_addr, int64_t n, uintptr_t staging_addr) { + const int64_t *tokens = reinterpret_cast(tokens_addr); + uint8_t *staging = reinterpret_cast(staging_addr); + const size_t heads = sizes_.size(); + std::vector rows(heads); + for (int64_t i = 0; i < n; i++) { + hash_rows(tokens + i, rows.data()); + for (size_t h = 0; h < heads; h++) + request_row(rows[h], staging + ((size_t)i * heads + h) * row_bytes_); + } + } + + // One batched disk round for everything staged; signals even when nothing was. + void flush(uintptr_t signal_addr) { + flush_pending(); + if (signal_addr) signal_flag(signal_addr); + } + + std::string io_backend() const { + size_t direct = 0; + for (const auto &f : files_) direct += f->direct_io() ? 1 : 0; + std::string s = reader_->name(); + if (direct == files_.size()) return s + ", O_DIRECT"; + return s + ", buffered " + std::to_string(files_.size() - direct) + "/" + + std::to_string(files_.size()) + " files"; + } + + private: + struct Pending { + const TableFile *file; + int64_t read_off; + int64_t read_len; + int64_t row_off; // row payload start inside the read buffer + std::vector dsts; + }; + + // Queue dst on this fill's pending batch; duplicate rows fan out from one read. + void request_row(int64_t row_id, uint8_t *dst) { + auto pit = pending_index_.find(row_id); + if (pit != pending_index_.end()) { + pending_[pit->second].dsts.push_back(dst); + return; + } + const Extent &ext = extents_[row_id / rows_per_extent_]; + const int64_t off = ext.base + (row_id % rows_per_extent_) * row_stride_; + Pending p{ext.file, off, row_bytes_, 0, {dst}}; + if (ext.file->direct_io()) { + // full aligned span even past EOF; truncating would break direct-I/O alignment + p.read_off = off & ~(kPage - 1); + p.row_off = off - p.read_off; + p.read_len = ((off + row_bytes_ + kPage - 1) & ~(kPage - 1)) - p.read_off; + } + pending_index_.emplace(row_id, pending_.size()); + pending_.push_back(std::move(p)); + } + + // Read every pending row in reader-capacity batches and fan out the copies. + void flush_pending() { + if (pending_.empty()) return; + struct Cleanup { + PleStore *s; + ~Cleanup() { + s->pending_.clear(); + s->pending_index_.clear(); + } + } cleanup{this}; + + const unsigned cap = reader_->capacity(); + const size_t total = pending_.size(); + std::vector tag_pending(cap); + size_t next = 0; + auto submit_slot = [&](unsigned tag) { + const Pending &p = pending_[next]; + tag_pending[tag] = next++; + reader_->submit(tag, p.file->native_fd(), bounce_ + (size_t)tag * kSpanMax, p.read_len, + p.row_off + row_bytes_, p.read_off); + }; + try { + for (unsigned tag = 0; tag < std::min((size_t)cap, total); tag++) submit_slot(tag); + for (size_t completed = 0; completed < total; completed++) { + const unsigned tag = reader_->wait_one(); + const Pending &p = pending_[tag_pending[tag]]; + const uint8_t *row = bounce_ + (size_t)tag * kSpanMax + p.row_off; + for (uint8_t *dst : p.dsts) std::memcpy(dst, row, row_bytes_); + p.file->discard_cache(p.read_off, p.read_len); + if (next < total) submit_slot(tag); + } + } catch (...) { + reader_->drain(); + throw; + } + } + + int64_t row_bytes_, row_stride_, rows_per_extent_; + std::vector mult_, sizes_, offsets_; + int64_t eos_; + std::vector> files_; + std::vector extents_; + std::unique_ptr reader_; + uint8_t *bounce_ = nullptr; + + std::vector pending_; + std::unordered_map pending_index_; +}; + +} // namespace + +PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { + py::class_(m, "PleStore") + .def(py::init, std::vector, std::vector, + int64_t, int64_t, int64_t, std::vector, + std::vector, std::vector, int64_t, bool>(), + py::arg("paths"), py::arg("extent_file"), py::arg("extent_base"), + py::arg("rows_per_extent"), py::arg("row_bytes"), py::arg("row_stride"), + py::arg("multipliers"), + py::arg("head_vocab_sizes"), py::arg("head_offsets"), py::arg("eos_token_id"), + py::arg("use_io_uring") = true) + .def("stage", &PleStore::stage, py::arg("tokens_addr"), py::arg("n"), + py::arg("staging_addr"), py::call_guard()) + .def("flush", &PleStore::flush, py::arg("signal_addr") = 0, + py::call_guard()) + .def("io_backend", &PleStore::io_backend); + m.def("memop_write", &memop_write, py::arg("stream"), py::arg("addr"), py::arg("value")); + m.def("memop_wait_geq", &memop_wait_geq, py::arg("stream"), py::arg("addr"), py::arg("value")); + m.def("memop_wait_reset", &memop_wait_reset, py::arg("stream"), py::arg("flag_addr")); + m.def("signal_flag", &signal_flag, py::arg("flag_addr")); +} diff --git a/python/freetoken/models/blocks.py b/python/freetoken/models/blocks.py index ca5b19723..e811ca806 100644 --- a/python/freetoken/models/blocks.py +++ b/python/freetoken/models/blocks.py @@ -1,6 +1,7 @@ from __future__ import annotations from abc import ABC, abstractmethod +from contextlib import contextmanager from typing import TYPE_CHECKING from freetoken.layers import ( @@ -16,6 +17,8 @@ if TYPE_CHECKING: import torch + from freetoken.core import Batch + from .config import ModelConfig @@ -23,6 +26,11 @@ class BaseLLMModel(ABC, BaseOP): @abstractmethod def forward(self) -> torch.Tensor: ... + @contextmanager + def forward_host_ctx(self, batch: Batch, use_graph: bool): + """Around one forward dispatch: enter before it is enqueued, exit right after. A backend that feeds the forward from host memory overrides this.""" + yield + class GatedMLP(BaseOP): def __init__(self, config: ModelConfig): diff --git a/python/freetoken/models/qwen4_exp/model.py b/python/freetoken/models/qwen4_exp/model.py index e5b365aa3..19c31e2fc 100644 --- a/python/freetoken/models/qwen4_exp/model.py +++ b/python/freetoken/models/qwen4_exp/model.py @@ -169,6 +169,36 @@ def load_host_tables(self, engine_config) -> int: emb.attach_table(ZeroTable(offsets[-1] + sizes[-1], args.ngram_head_dim)) return 0 + if engine_config.ple_backend == "disk": + from freetoken.utils import download_hf_weight + + from .ple_disk import DiskRowTable, resolve_row_source + + folder = download_hf_weight(engine_config.model_path) + # one WAIT node per captured graph: the flag protocol supports a single consume + assert len(ple_layers) == 1, "disk PLE backend expects exactly one PLE layer" + emb, args = ple_layers[0].ple_embedding, ple_layers[0].args + # hash with the state-dict-loaded constants, the same source the pinned path reads + constants = { + "num_ngram_heads": args.num_ngram_heads, + "layer_multipliers": emb.layer_multipliers.tolist(), + "per_head_vocab_sizes": emb.ngram_heads_vocab_sizes.tolist(), + "per_head_offsets": emb.ngram_heads_offsets.tolist(), + "eos_token_id": args.ngram_boundary_token_id, + } + disk_table = DiskRowTable( + resolve_row_source(folder), + constants, + max_graph_rows=max(256, engine_config.cuda_graph_max_bs or 0), + max_extend_tokens=engine_config.max_extend_tokens, + ) + self._ple_table = disk_table + for ple in ple_layers: + ple.ple_embedding.attach_table(disk_table) + # engine enters this around every dispatch; the graph itself never waits on the disk + self.forward_host_ctx = disk_table.forward_host_ctx + return 0 + from .weight import load_ple_table table = load_ple_table(engine_config.model_path, self._config.qwen4_args) diff --git a/python/freetoken/models/qwen4_exp/ple_disk.py b/python/freetoken/models/qwen4_exp/ple_disk.py new file mode 100644 index 000000000..691546c3f --- /dev/null +++ b/python/freetoken/models/qwen4_exp/ple_disk.py @@ -0,0 +1,270 @@ +"""Disk-backed PLE table (--ple-backend disk): the C++ store hashes n-gram windows and batch-reads rows from the checkpoint's fp8 shard tensors into pinned staging; the captured ``lookup`` is a fixed-shape H2D copy + dequant. + +Hash windows are pure functions of ``req.input_ids`` + ``device_len`` (prefix hits, restores and COW forks need no bookkeeping); the decode input token lives device-side under overlap scheduling and is read back here. +""" + +from __future__ import annotations + +import os +from contextlib import contextmanager +from dataclasses import dataclass +from typing import Sequence + +import safetensors +import torch + +from freetoken.core import Batch +from freetoken.kernel.pinned import alloc_pinned_tensor +from freetoken.utils import init_logger + +from .weight import ( + _PLE_SCALE_SUFFIX, + _PLE_SHARD_RE, + _PLE_ST_DTYPE, + _ple_table_files, + _safetensors_header, +) + +_IO_URING_ENV = "FREETOKEN_PLE_IO_URING" +_SYNC_ENV = "FREETOKEN_PLE_SYNC" # auto | wait | gate + +logger = init_logger(__name__) + + +def _context(ids: torch.Tensor, position: int, eos: int) -> list[int]: + """The two token ids before ``position``; eos pads past the start.""" + return [int(ids[position - 2]) if position >= 2 else eos, + int(ids[position - 1]) if position >= 1 else eos] + + +@dataclass(frozen=True) +class PleRowSource: + """On-disk row layout: equal extents, row i of an extent at ``base + i * row_stride`` (a repacked flat file is one extent with its own stride).""" + + paths: list[str] + extent_file: list[int] + extent_base: list[int] + rows_per_extent: int + row_bytes: int + row_stride: int + scale: float + + @property + def total_rows(self) -> int: + return len(self.extent_base) * self.rows_per_extent + + +def source_from_safetensors(folder: str) -> PleRowSource: + """Map the checkpoint's ``ngram_embedding.shard_`` tensors in place: one extent per shard, no copy.""" + rows = cols = 0 + scale: torch.Tensor | None = None + paths: list[str] = [] + path_idx: dict[str, int] = {} + shards: dict[int, tuple[int, int]] = {} + for path in _ple_table_files(folder): + header, base = _safetensors_header(path) + for key, meta in header.items(): + if key == "__metadata__": + continue + if key.endswith(_PLE_SCALE_SUFFIX): + with safetensors.safe_open(path, framework="pt", device="cpu") as f: + scale = f.get_tensor(key).reshape(()) + continue + match = _PLE_SHARD_RE.search(key) + if match is None: + continue + if meta["dtype"] != _PLE_ST_DTYPE: + raise ValueError(f"PLE shard {key} has dtype {meta['dtype']}, expected {_PLE_ST_DTYPE}") + if rows and tuple(meta["shape"]) != (rows, cols): + raise ValueError(f"PLE shard {key} is {meta['shape']}, expected {[rows, cols]}") + rows, cols = meta["shape"] + if path not in path_idx: + path_idx[path] = len(paths) + paths.append(path) + idx = int(match.group("shard")) + if idx in shards: + raise ValueError(f"duplicate PLE shard {idx} in {path}") + shards[idx] = (path_idx[path], base + meta["data_offsets"][0]) + if sorted(shards) != list(range(len(shards))) or not shards: + raise ValueError(f"PLE shard indices are not contiguous 0..N-1: {sorted(shards)[:8]}") + if scale is None: + raise ValueError("PLE table has no weight_scale") + order = [shards[i] for i in range(len(shards))] + return PleRowSource(paths, [f for f, _ in order], [b for _, b in order], rows, cols, cols, float(scale)) + + +def resolve_row_source(folder: str) -> PleRowSource: + """Pick the row source for a checkpoint; the seam where a repacked format would plug in.""" + return source_from_safetensors(folder) + + +class DiskRowTable: + """``PLETableBackend`` whose rows are read from disk per fill (--ple-backend disk).""" + + def __init__( + self, + source: PleRowSource, + hash_constants: dict, + *, + max_graph_rows: int = 256, + max_extend_tokens: int = 8192, + dtype: torch.dtype = torch.bfloat16, + ) -> None: + from freetoken.kernel import _ple_store + + self.num_rows = source.total_rows + self.head_dim = source.row_bytes # fp8: one byte per element + self.dtype = dtype + self.heads = int(hash_constants["num_ngram_heads"]) + self.scale = source.scale + self.eos_token_id = int(hash_constants["eos_token_id"]) + sizes = [int(x) for x in hash_constants["per_head_vocab_sizes"]] + offsets = [int(x) for x in hash_constants["per_head_offsets"]] + need = max(o + s for o, s in zip(offsets, sizes)) + if need > source.total_rows: + raise ValueError( + f"PLE row source holds {source.total_rows} rows but the hash addresses {need}; incomplete checkpoint?" + ) + self._store = _ple_store.PleStore( + paths=list(source.paths), + extent_file=list(source.extent_file), + extent_base=list(source.extent_base), + rows_per_extent=source.rows_per_extent, + row_bytes=source.row_bytes, + row_stride=source.row_stride, + multipliers=[int(x) for x in hash_constants["layer_multipliers"]], + head_vocab_sizes=sizes, + head_offsets=offsets, + eos_token_id=self.eos_token_id, + use_io_uring=os.getenv(_IO_URING_ENV, "1") != "0", + ) + self._device = torch.device("cuda", torch.cuda.current_device()) + self._token_bytes = self.heads * self.head_dim + # allocated up front: pinned alloc inside stream capture is illegal; one replay consumes it at a time + self._graph_pinned = alloc_pinned_tensor(max_graph_rows * self._token_bytes, dtype=torch.uint8) + self._graph_pinned.zero_() # padded decode lanes read whatever sits here + # outlives any one graph: a cache rebuild recaptures against the same pointer + self._graph_dev = torch.empty( + max_graph_rows * self._token_bytes, dtype=torch.uint8, device=self._device + ) + eager_bytes = max_extend_tokens * self._token_bytes + self._eager_pinned = alloc_pinned_tensor(eager_bytes, dtype=torch.uint8) + self._eager_pinned.zero_() # the warmup prefill stages nothing and reads whatever sits here + self._eager_dev = torch.empty(eager_bytes, dtype=torch.uint8, device=self._device) + # probe picks flag-sync (graph WAITs at the consume, host fills then signals) or launch-gating + self._wait_sync = self._probe_wait_sync(os.getenv(_SYNC_ENV, "auto")) + # one flag for all graphs: the readback event orders a fill after the previous graph, so signals never overlap + self._flag = alloc_pinned_tensor(1, dtype=torch.int64) + self._flag.zero_() + self._token_readback = alloc_pinned_tensor(max_graph_rows, dtype=torch.int32) + self._readback_event = torch.cuda.Event() + sync = "wait-sync" if self._wait_sync else "launch-gating" + logger.info_rank0(f"PLE disk backend: {self._store.io_backend()}, {sync}") + + def _probe_wait_sync(self, mode: str) -> bool: + from freetoken.kernel import _ple_store + + if mode == "gate": + return False + scratch = alloc_pinned_tensor(1, dtype=torch.int64) + scratch.zero_() + stream = torch.cuda.current_stream(self._device) + ok = ( + _ple_store.memop_write(stream.cuda_stream, scratch.data_ptr(), 7) == 0 + and _ple_store.memop_wait_geq(stream.cuda_stream, scratch.data_ptr(), 7) == 0 + ) + if ok: + stream.synchronize() + ok = int(scratch[0]) == 7 + if mode == "wait" and not ok: + raise RuntimeError("FREETOKEN_PLE_SYNC=wait but stream memops are unavailable") + return ok + + # ---------------- host side (engine thread, before the forward launches) ---------------- + + def fill(self, runs: Sequence[torch.Tensor], *, graph: bool) -> None: + """Stage per-request token runs (two context ids, then the new tokens) in batch order.""" + pinned = self._graph_pinned if graph else self._eager_pinned + offset = 0 + for run in runs: + self._store.stage(run.data_ptr(), run.numel() - 2, pinned.data_ptr() + offset * self._token_bytes) + offset += run.numel() - 2 + self._store.flush(self._flag.data_ptr() if graph and self._wait_sync else 0) + + def host_fill_batch(self, batch: Batch, use_graph: bool): + """Stage this batch's rows; returns the post-dispatch fill callable under flag-sync, else None.""" + eos = self.eos_token_id + if batch.is_decode: + reqs = list(batch.reqs) + if use_graph and self._wait_sync: + bs = batch.size + self._token_readback[:bs].copy_(batch.input_ids, non_blocking=True) + self._readback_event.record(torch.cuda.current_stream(self._device)) + + def _complete() -> None: + try: + self._readback_event.synchronize() + tokens = self._token_readback[:bs].to(torch.int64).tolist() + runs = [torch.tensor([*_context(r.input_ids, r.device_len - 1, eos), t], dtype=torch.int64) + for r, t in zip(reqs, tokens)] + self.fill(runs, graph=True) + except BaseException: + from freetoken.kernel import _ple_store + + # unblock the stream before surfacing; the step's output is discarded + _ple_store.signal_flag(self._flag.data_ptr()) + raise + + return _complete + # launch-gating: this D2H is the step's readback and orders the fill after sampling + tokens = batch.input_ids.to("cpu").to(torch.int64).tolist() + runs = [torch.tensor([*_context(r.input_ids, r.device_len - 1, eos), t], dtype=torch.int64) + for r, t in zip(reqs, tokens)] + self.fill(runs, graph=use_graph) + return None + runs = [ + torch.cat(( + torch.tensor(_context(req.input_ids, req.cached_len, eos), dtype=torch.int64), + req.input_ids[req.cached_len : req.device_len].to(torch.int64), + )) + for req in batch.padded_reqs + ] + self.fill(runs, graph=False) + return None + + @contextmanager + def forward_host_ctx(self, batch: Batch, use_graph: bool): + """Around one dispatch: stage on enter, run the deferred fill+signal on exit.""" + deferred = self.host_fill_batch(batch, use_graph) + yield + # no try/finally: a failed launch leaves no WAIT pending, so the fill must not run + if deferred is not None: + deferred() + + # ---------------- device side (PLETableBackend protocol) ---------------- + + def lookup(self, row_ids: torch.Tensor, out: torch.Tensor | None = None) -> torch.Tensor: + rows = row_ids.shape[0] + capturing = torch.cuda.is_current_stream_capturing() + if capturing and self._wait_sync: + from freetoken.kernel import _ple_store + + _ple_store.memop_wait_reset( + torch.cuda.current_stream(self._device).cuda_stream, self._flag.data_ptr() + ) + pinned, dev = ( + (self._graph_pinned, self._graph_dev) if capturing else (self._eager_pinned, self._eager_dev) + ) + nbytes = rows * self._token_bytes + dev[:nbytes].copy_(pinned[:nbytes], non_blocking=True) + values = dev[:nbytes].view(torch.float8_e4m3fn).to(self.dtype) + if self.scale != 1.0: + values = values * self.scale + values = values.view(*row_ids.shape[:-1], -1) + if out is None: + return values + out.copy_(values) + return out + + def prefetch(self, row_ids: torch.Tensor) -> None: + return None diff --git a/python/freetoken/server/args.py b/python/freetoken/server/args.py index 6696f65dd..5b4db587d 100644 --- a/python/freetoken/server/args.py +++ b/python/freetoken/server/args.py @@ -477,6 +477,16 @@ def _infer_reasoning_parser(model_path: str) -> str | None: ), ) + parser.add_argument( + "--ple-backend", + default=ServerArgs.ple_backend, + choices=["pinned", "disk"], + help=( + "Where a PLE n-gram table lives. 'disk' (default) reads rows straight from the " + "checkpoint files; 'pinned' preloads the whole table into page-locked host RAM." + ), + ) + parser.add_argument( "--nvfp4-backend", default=ServerArgs.nvfp4_backend, diff --git a/setup.py b/setup.py index cfe41b7d8..d8e705ab6 100644 --- a/setup.py +++ b/setup.py @@ -3,6 +3,8 @@ import importlib.util from pathlib import Path +import sys + from setuptools import setup from torch.utils.cpp_extension import BuildExtension, CUDA_HOME, CppExtension @@ -62,6 +64,16 @@ def _cuda_runtime_paths() -> tuple[list[str], list[str]]: libraries=["cudart"], extra_compile_args=["-O3", "-std=c++17", "-pthread"], ), + # --ple-backend disk row store; Linux-only until the TableFile/BatchReader seams grow Windows bodies + *([ + CppExtension( + name="freetoken.kernel._ple_store", + sources=[ + "python/freetoken/kernel/csrc/ple_store/ple_store_ext.cpp", + ], + extra_compile_args=["-O3", "-std=c++17"], + ) + ] if sys.platform == "linux" else []), ], cmdclass={"build_ext": BuildExtension.with_options(use_ninja=True)}, ) diff --git a/tests/models/qwen4_exp/test_ple_disk.py b/tests/models/qwen4_exp/test_ple_disk.py new file mode 100644 index 000000000..fa1180739 --- /dev/null +++ b/tests/models/qwen4_exp/test_ple_disk.py @@ -0,0 +1,334 @@ +"""Disk PLE backend, module level: store byte fidelity, on-disk layouts and errors, table vs GPU oracle, and the CUDA-graph sync protocol.""" + +from __future__ import annotations + +from types import SimpleNamespace + +import pytest +import torch + +from freetoken.models.qwen4_exp.config import parse_config +from freetoken.models.qwen4_exp.ple import GpuResidentTable, NGramEmbedding + +from .common import EOS, hash_constants, requires_cuda, toy_hf_config +from .test_ple import _meta + +_ple_store = pytest.importorskip("freetoken.kernel._ple_store") + +_KEY_PREFIX = "model.layers.1.ple.ple_embedding.ngram_embedding" + + +def _embedding() -> NGramEmbedding: + args = parse_config(toy_hf_config()).qwen4_args + emb = NGramEmbedding(args) + multipliers, sizes, offsets = hash_constants(args) + emb.layer_multipliers.copy_(multipliers) + emb.ngram_heads_vocab_sizes.copy_(sizes) + emb.ngram_heads_offsets.copy_(offsets) + return emb + + +def _bitwise_equal(got: torch.Tensor, want: torch.Tensor) -> bool: + # random table bytes include fp8 NaN encodings, and NaN != NaN under torch.equal + return torch.equal(got.view(torch.int16), want.view(torch.int16)) + + +def _make_store(tmp_path, *, write=True, use_io_uring=True): + args = parse_config(toy_hf_config()).qwen4_args + multipliers, sizes, offsets = hash_constants(args) + total_rows = int(offsets[-1] + sizes[-1]) + cols = args.ngram_head_dim + gen = torch.Generator().manual_seed(5) + table = torch.randint(0, 256, (total_rows, cols), dtype=torch.uint8, generator=gen) + path = tmp_path / "ple-table.bin" + if write: + path.write_bytes(table.numpy().tobytes()) + store = _ple_store.PleStore( + paths=[str(path)], + extent_file=[0], + extent_base=[0], + rows_per_extent=total_rows, + row_bytes=cols, + row_stride=cols, + multipliers=multipliers.tolist(), + head_vocab_sizes=sizes.tolist(), + head_offsets=offsets.tolist(), + eos_token_id=EOS, + use_io_uring=use_io_uring, + ) + return store, table, args + + +def _fill(store, args, window, tokens): + ctx = torch.tensor([window[0], window[1], *tokens], dtype=torch.int64) + staging = torch.empty(len(tokens) * args.num_ngram_heads * args.ngram_head_dim, dtype=torch.uint8) + store.stage(ctx.data_ptr(), len(tokens), staging.data_ptr()) + store.flush(0) + return staging + + +def _write_checkpoint(tmp_path, table, n_shards): + from safetensors.torch import save_file + + per = table.shape[0] // n_shards + tensors = { + f"{_KEY_PREFIX}.shard_{i}.weight": table[i * per : (i + 1) * per].view(torch.float8_e4m3fn) + for i in range(n_shards) + } + tensors[f"{_KEY_PREFIX}.weight_scale"] = torch.tensor(0.03125, dtype=torch.bfloat16) + save_file(tensors, str(tmp_path / "model.safetensors")) + + +def _make_table(tmp_path): + from freetoken.models.qwen4_exp.ple_disk import DiskRowTable, source_from_safetensors + + args = parse_config(toy_hf_config()).qwen4_args + multipliers, sizes, offsets = hash_constants(args) + total_rows = int(offsets[-1] + sizes[-1]) + gen = torch.Generator().manual_seed(9) + table = torch.randint(0, 256, (total_rows, args.ngram_head_dim), dtype=torch.uint8, generator=gen) + n_shards = next(k for k in (4, 2, 1) if total_rows % k == 0) + _write_checkpoint(tmp_path, table, n_shards) + constants = { + "num_ngram_heads": args.num_ngram_heads, + "layer_multipliers": multipliers.tolist(), + "per_head_vocab_sizes": sizes.tolist(), + "per_head_offsets": offsets.tolist(), + "eos_token_id": EOS, + } + disk = DiskRowTable(source_from_safetensors(str(tmp_path)), constants) + oracle = GpuResidentTable(table.cuda().view(torch.float8_e4m3fn), scale=0.03125) + return disk, oracle, args + + +def _decode_batch(history, token): + req = SimpleNamespace( + input_ids=torch.tensor(history, dtype=torch.int32), + device_len=len(history) + 1, + cached_len=len(history), + ) + return SimpleNamespace( + is_decode=True, + input_ids=torch.tensor([token], dtype=torch.int32, device="cuda"), + reqs=[req], + size=1, + padded_size=1, + ) + + +def test_store_stages_bitwise_rows(tmp_path): + store, table, args = _make_store(tmp_path) + emb = _embedding() + row = args.num_ngram_heads * args.ngram_head_dim + + # full run with mid-sequence eos vs production row ids + seq = [3, 4, EOS, 5, EOS, EOS, 8, 9, 3, 4] + whole = _fill(store, args, (EOS, EOS), seq) + ids = emb.row_ids(_meta([seq], [[EOS, EOS]])) + assert torch.equal(whole, table[ids.reshape(-1)].reshape(-1)), "prefill vs oracle" + + # decode = many 1-token stages; must reproduce the same bytes + parts, window = [], (EOS, EOS) + for t in seq: + parts.append(_fill(store, args, window, [t])) + window = (window[1], t) + assert torch.equal(torch.cat(parts), whole), "split stages vs one stage" + + # several lanes merged into one flush stay independent + contexts = [(3, 4), (EOS, EOS), (7, EOS)] + ctx = torch.tensor([[o, nw, 9] for o, nw in contexts], dtype=torch.int64) + staging = torch.empty(3 * row, dtype=torch.uint8) + for i in range(len(contexts)): + store.stage(ctx.data_ptr() + 24 * i, 1, staging.data_ptr() + i * row) + store.flush(0) + for i, context in enumerate(contexts): + assert torch.equal(staging[i * row : (i + 1) * row], _fill(store, args, context, [9])), f"lane {i}" + + # hundreds of deduped reads through the 64-deep pipeline + gen = torch.Generator().manual_seed(23) + big = torch.randint(0, EOS, (150,), generator=gen, dtype=torch.int64) + got = _fill(store, args, (EOS, EOS), big.tolist()) + ids = emb.row_ids(_meta([big.tolist()], [[EOS, EOS]])) + assert torch.equal(got, table[ids.reshape(-1)].reshape(-1)), "pipeline vs oracle" + + # flush signals the flag, even when nothing was staged + flag = torch.zeros(1, dtype=torch.int64) + store.flush(flag.data_ptr()) + assert int(flag[0]) == 1, "empty flush must still signal" + + +def test_layouts_readers_and_errors(tmp_path): + # 4 extents in 2 files, out of order, unaligned junk between; the last extent ends at EOF + sizes, offsets = [500, 400, 300, 800], [0, 500, 900, 1200] + total, per, cols, eos = 2000, 500, 24, 90 + gen = torch.Generator().manual_seed(11) + table = torch.randint(0, 256, (total, cols), dtype=torch.uint8, generator=gen) + shard = lambda i: table[i * per : (i + 1) * per].numpy().tobytes() # noqa: E731 + nb = per * cols + flat = tmp_path / "flat.bin" + flat.write_bytes(table.numpy().tobytes()) + fa, fb = tmp_path / "a.bin", tmp_path / "b.bin" + fa.write_bytes(b"j" * 1231 + shard(0) + b"k" * 77 + shard(2)) + fb.write_bytes(shard(1) + b"m" * 4095 + shard(3)) + kwargs = dict( + rows_per_extent=per, row_bytes=cols, row_stride=cols, + multipliers=[3, 5, 7], head_vocab_sizes=sizes, head_offsets=offsets, + eos_token_id=eos, + ) + ref = _ple_store.PleStore( + paths=[str(flat)], extent_file=[0, 0, 0, 0], extent_base=[0, nb, 2 * nb, 3 * nb], **kwargs + ) + multi = _ple_store.PleStore( + paths=[str(fa), str(fb)], extent_file=[0, 1, 0, 1], + extent_base=[1231, 0, 1231 + nb + 77, nb + 4095], **kwargs, + ) + tokens = torch.randint(0, eos, (40,), generator=gen, dtype=torch.int64) + ctx = torch.cat((torch.tensor([eos, eos], dtype=torch.int64), tokens)) + + def run(store): + staging = torch.empty(40 * 4 * cols, dtype=torch.uint8) + store.stage(ctx.data_ptr(), 40, staging.data_ptr()) + store.flush(0) + return staging + + assert torch.equal(run(multi), run(ref)), "multi-extent vs flat" + + # thread-pool fallback must produce the same bytes as io_uring + ring, _, args = _make_store(tmp_path) + pool, _, _ = _make_store(tmp_path, write=False, use_io_uring=False) + gen2 = torch.Generator().manual_seed(31) + seq = torch.randint(0, EOS, (150,), generator=gen2, dtype=torch.int64).tolist() + assert torch.equal(_fill(pool, args, (EOS, EOS), seq), _fill(ring, args, (EOS, EOS), seq)), "pool vs ring" + + # geometry that exceeds the file is rejected at construction + del ring, pool + with open(tmp_path / "ple-table.bin", "r+b") as fh: + fh.truncate(1000) + with pytest.raises(Exception, match="extent needs"): + _make_store(tmp_path, write=False) + + # checkpoint scan guards + from safetensors.torch import save_file + + from freetoken.models.qwen4_exp.ple_disk import source_from_safetensors + + save_file( + {f"{_KEY_PREFIX}.shard_0.weight": torch.zeros(8, 4, dtype=torch.uint8), + f"{_KEY_PREFIX}.weight_scale": torch.tensor(1.0, dtype=torch.bfloat16)}, + str(tmp_path / "model.safetensors"), + ) + with pytest.raises(ValueError, match="dtype"): + source_from_safetensors(str(tmp_path)) + save_file( + {f"{_KEY_PREFIX}.shard_1.weight": torch.zeros(8, 4, dtype=torch.float8_e4m3fn), + f"{_KEY_PREFIX}.weight_scale": torch.tensor(1.0, dtype=torch.bfloat16)}, + str(tmp_path / "model.safetensors"), + ) + with pytest.raises(ValueError, match="contiguous"): + source_from_safetensors(str(tmp_path)) + save_file( + {f"{_KEY_PREFIX}.shard_0.weight": torch.zeros(8, 4, dtype=torch.float8_e4m3fn), + f"{_KEY_PREFIX}.weight_scale": torch.tensor(1.0, dtype=torch.bfloat16)}, + str(tmp_path / "model.safetensors"), + ) + save_file( + {f"{_KEY_PREFIX}.shard_0.weight": torch.zeros(8, 4, dtype=torch.float8_e4m3fn)}, + str(tmp_path / "model-2.safetensors"), + ) + with pytest.raises(ValueError, match="duplicate"): + source_from_safetensors(str(tmp_path)) + (tmp_path / "model-2.safetensors").unlink() + + # truncated checkpoint: a contiguous shard prefix passes the scan, init rejects the row count + from freetoken.models.qwen4_exp.ple_disk import DiskRowTable + + args = parse_config(toy_hf_config()).qwen4_args + multipliers, vocab, offs = hash_constants(args) + rows = int(offs[-1] + vocab[-1]) + _write_checkpoint(tmp_path, torch.zeros(rows // 2, args.ngram_head_dim, dtype=torch.uint8), 1) + constants = { + "num_ngram_heads": args.num_ngram_heads, "layer_multipliers": multipliers.tolist(), + "per_head_vocab_sizes": vocab.tolist(), "per_head_offsets": offs.tolist(), "eos_token_id": EOS, + } + with pytest.raises(ValueError, match="hash addresses"): + DiskRowTable(source_from_safetensors(str(tmp_path)), constants) + + +@requires_cuda +def test_disk_table_matches_oracle(tmp_path): + disk, oracle, args = _make_table(tmp_path) + emb = _embedding() + + # prefill: two segments, one fresh and one mid-sequence window + seqs = [[3, 4, EOS, 5, 6, 8], [2, EOS, 11, 12, 13, 14]] + disk.fill([torch.tensor([EOS, EOS, *seqs[0]]), torch.tensor([21, 22, *seqs[1]])], graph=False) + row_ids = emb.row_ids(_meta(seqs, [[EOS, EOS], [21, 22]])).cuda() + assert _bitwise_equal(disk.lookup(row_ids), oracle.lookup(row_ids)), "prefill" + + # decode steps with a rolling window, plus the out= contract + older, newer = 41, EOS + for token in (7, 9, 13): + disk.fill([torch.tensor([older, newer, token])], graph=False) + ids = emb.row_ids(_meta([[token]], [[older, newer]], decode=True)).cuda() + out = torch.empty((1, ids.shape[-1] * disk.head_dim), dtype=disk.dtype, device="cuda") + assert disk.lookup(ids, out) is out and _bitwise_equal(out, oracle.lookup(ids)), f"token {token}" + older, newer = newer, token + + # the engine hook end to end: eager decode, then fresh + continuation prefill + disk.host_fill_batch(_decode_batch([3, 4, EOS, 5], 9), use_graph=False) + ids = emb.row_ids(_meta([[9]], [[EOS, 5]], decode=True)).cuda() + assert _bitwise_equal(disk.lookup(ids), oracle.lookup(ids)), "hook decode" + + prompt = [3, 4, EOS, 5, 6, 8] + fresh = SimpleNamespace(input_ids=torch.tensor(prompt[:4], dtype=torch.int32), device_len=4, cached_len=0) + cont = SimpleNamespace(input_ids=torch.tensor(prompt, dtype=torch.int32), device_len=6, cached_len=4) + disk.host_fill_batch(SimpleNamespace(is_decode=False, padded_reqs=[fresh, cont]), use_graph=False) + ids = emb.row_ids(_meta([prompt[:4], prompt[4:]], [[EOS, EOS], [prompt[2], prompt[3]]])).cuda() + assert _bitwise_equal(disk.lookup(ids), oracle.lookup(ids)), "hook prefill" + + +@requires_cuda +def test_graph_sync_protocol(tmp_path, monkeypatch): + disk, oracle, args = _make_table(tmp_path) + emb = _embedding() + rows = 1 + row_ids = torch.zeros((rows, args.num_ngram_heads), dtype=torch.int64, device="cuda") + out = torch.empty((rows, args.num_ngram_heads * disk.head_dim), dtype=disk.dtype, device="cuda") + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(stream): + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=stream): + disk.lookup(row_ids, out) + torch.cuda.current_stream().wait_stream(stream) + + # plain fill-then-replay + disk.fill([torch.tensor([3, 4, 7])], graph=True) + graph.replay() + torch.cuda.synchronize() + ids = emb.row_ids(_meta([[7]], [[3, 4]], decode=True)).cuda() + assert _bitwise_equal(out, oracle.lookup(ids)), "capture+replay" + + if disk._wait_sync: + # launch first, fill after: an early WAIT pass would surface the previous step's bytes + older, newer = 3, 4 + for token in (7, 9, 13): + graph.replay() + disk.fill([torch.tensor([older, newer, token])], graph=True) + torch.cuda.synchronize() + ids = emb.row_ids(_meta([[token]], [[older, newer]], decode=True)).cuda() + assert _bitwise_equal(out, oracle.lookup(ids)), f"wait-sync token {token}" + older, newer = newer, token + + # the engine-shaped seam: replay inside the context, deferred fill on exit + with disk.forward_host_ctx(_decode_batch([3, 4], 7), use_graph=True): + graph.replay() + torch.cuda.synchronize() + ids = emb.row_ids(_meta([[7]], [[3, 4]], decode=True)).cuda() + assert _bitwise_equal(out, oracle.lookup(ids)), "forward_host_ctx deferred" + + # gate mode: the hook fills inline and returns no deferred + monkeypatch.setenv("FREETOKEN_PLE_SYNC", "gate") + gated, _, _ = _make_table(tmp_path) + assert not gated._wait_sync + assert gated.host_fill_batch(_decode_batch([3, 4], 5), use_graph=True) is None