diff --git a/AGENTS.md b/AGENTS.md index 581d9af..7373a59 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -29,6 +29,7 @@ Before modifying a subsystem, read its corresponding documentation. | `ARCHITECTURE.md` | High-level system architecture | | `ROADMAP.md` | Development roadmap and milestones | | `MODEL_CONTRACTS.md` | ONNX model interfaces and tensor contracts | +| `KV.md` | Decoder KV cache (host vs CUDA) and GPU decode plan | Avoid duplicating documentation across multiple files. High-level concepts belong in `ARCHITECTURE.md`, while subsystem-specific implementation details belong in their dedicated documents. diff --git a/README.md b/README.md index 84f8294..bead619 100644 --- a/README.md +++ b/README.md @@ -245,8 +245,10 @@ cargo run --features load-dynamic,cuda --example inference -- \ models/lightonocr examples/SROIE-receipt.jpeg default cuda ``` -Optional 5th argument is the CUDA `device_id` (default `0`). On CUDA, decode keeps -KV past/present on the GPU via IoBinding; sampling still runs on the host. +Optional 5th argument is the CUDA `device_id` (default `0`). On CUDA, decode +defaults to device-resident KV via IoBinding (`CudaKVCache`); set +`FAST_LIGHTONOCR_CUDA_HOST_KV` to force the host `KVCache` path. Sampling still +runs on the host. See [`docs/KV.md`](docs/KV.md). To compare CPU thread settings on your machine: @@ -307,7 +309,7 @@ The Python bindings use the same native library and build infrastructure. - ✅ CPU performance work (KV-cache reuse, top-k/top-p, ORT session tuning, decode host reuse, inference bench) - 🚧 Generation parity and deterministic seeded generation - 🚧 Broader processor parity coverage -- ✅ CUDA execution provider + device-resident decoder KV (IoBinding) +- ✅ CUDA execution provider + `KVCacheBackend` (`KVCache` / `CudaKVCache`) - 🚧 CoreML / DirectML execution providers - 🚧 Python exposure of runtime / EP options diff --git a/bindings/python/README.md b/bindings/python/README.md index decce38..272aa4e 100644 --- a/bindings/python/README.md +++ b/bindings/python/README.md @@ -1,6 +1,6 @@ # fast-lightonocr -> ⚡ Native Python bindings for the Rust **Fast LightOnOCR** inference engine. +> Native Python bindings for the Rust **Fast LightOnOCR** inference engine. `fast-lightonocr` provides high-performance OCR for documents and images using Baidu's **LightOnOCR** model. Model inference runs entirely in native Rust, @@ -9,25 +9,25 @@ document parsing. --- -## ✨ Features +## Features -- 🚀 Native Rust inference engine -- 🧠 ONNX Runtime backend -- 📄 OCR for documents and images -- 📝 Structured Markdown output -- 📊 Structured HTML table extraction -- 🎨 Configurable table rendering -- 🎛️ Multiple model presets (`default`, `fp16`, `q4`) +- Native Rust inference engine +- ONNX Runtime backend +- OCR for documents and images +- Structured Markdown output +- Structured HTML table extraction +- Configurable table rendering +- Multiple model presets (`default`, `fp16`, `q4`) --- -## 📦 Installation +## Installation Install with the matching extra for your backend. Published wheels target **Linux x86_64** and **macOS arm64** (macOS Intel is not published: ONNX Runtime 1.28 has no compatible wheel there). -### CPU (default) +### CPU ```bash pip install "fast-lightonocr[cpu]" @@ -37,32 +37,20 @@ CPU wheels bundle ONNX Runtime. No extra environment setup is required. ### CUDA +Published CUDA wheels are a dedicated build profile (default PyPI wheels stay +CPU). Install a CUDA-profile package plus the extra: + ```bash pip install "fast-lightonocr[cuda]" ``` -Requires a CUDA-enabled package build and a compatible NVIDIA driver. The -`cuda` extra pulls in `onnxruntime-gpu` (CUDA 13 / cuDNN) and `nvidia-cublas`. - -Select CUDA at load time: - -```python -model = LightOnOCR.from_pretrained( - "onnx-community/LightOnOCR-2-1B-ONNX", - runtime_kwargs={ - "execution_provider": "cuda", - "device_id": 0, - }, -) -``` - -When `execution_provider="cuda"`, `from_pretrained` preloads the pip NVIDIA -CUDA/cuDNN libraries (`onnxruntime.preload_dlls`), so `LD_LIBRARY_PATH` is -usually unnecessary. CPU loads never take that path. +Requires a compatible NVIDIA driver. The `cuda` extra pulls in +`onnxruntime-gpu` (CUDA 13 / cuDNN) and `nvidia-cublas`. Select CUDA at load +time with `runtime_kwargs` — see [Runtime options](#runtime-options). ### Building from source -Source installs use the project build backend. It discovers ONNX Runtime from +Run these from `bindings/python`. The build backend discovers ONNX Runtime from `ORT_DYLIB_PATH` when set, otherwise from the profile’s Python ORT package, validates ONNX Runtime 1.28.x (C API level 27), and bundles the native runtime into the wheel. @@ -70,18 +58,16 @@ into the wheel. #### CPU ```bash +cd bindings/python pip install -v ".[cpu]" -``` - # or explicitly: - -```bash BUILD_PROFILE=cpu pip install -v ".[cpu]" ``` #### CUDA ```bash +cd bindings/python BUILD_PROFILE=cuda pip install -v ".[cuda]" ``` @@ -90,9 +76,11 @@ ORT CUDA provider plugins (`libonnxruntime_providers_{shared,cuda}`) into the wheel. The `[cuda]` extra installs the CUDA 13 / cuDNN / cublas user libraries used at runtime. +For editable/`maturin develop` workflows, see [Development](#development). + --- -## 🚀 Quick Start +## Quick Start ```python from fast_lightonocr import LightOnOCR @@ -109,80 +97,95 @@ them locally. --- -## 📄 OCR Results - -The raw model output is available through `result.text`. +## Configuration -```python -print(result.text) -``` +`from_pretrained()` accepts a model preset plus two override dicts: +`runtime_kwargs` (ONNX Runtime sessions) and `generation_kwargs` (decode). -The Python bindings also expose a parsed document representation that extracts -embedded HTML tables while preserving the original document structure. +### Model presets ```python -print(result.document) +model = LightOnOCR.from_pretrained( + "onnx-community/LightOnOCR-2-1B-ONNX", + preset="q4", +) ``` -Tables can be accessed directly: +Available presets: -```python -for table in result.tables: - print(table.text_rows) -``` +- `default` +- `fp16` +- `q4` ---- +### Runtime options -## 📋 Table Rendering +Override ONNX Runtime session settings at load time with `runtime_kwargs`. +Unknown keys raise `ValueError`. These options are applied **before** sessions +are created and cannot be changed after load. -By default, tables are rendered using ASCII borders. +Supported keys: -```python -result = model.process( - "receipt.jpg", - table_format="grid", -) -``` +| Key | Type | Default | Notes | +| --- | --- | --- | --- | +| `execution_provider` | `"cpu"` \| `"cuda"` | `"cpu"` | `"cuda"` requires a CUDA-enabled build and `[cuda]` extra | +| `device_id` | `int` | `0` | CUDA device index | +| `intra_threads` | `int` | host parallelism | Intra-op threads (no effect if ORT is built with OpenMP; use `OMP_NUM_THREADS`) | +| `inter_threads` | `int` | `1` | Used only when `parallel_execution` is `True` | +| `parallel_execution` | `bool` | `False` | ORT parallel execution mode | -Markdown tables are also supported. +CUDA: ```python -result = model.process( - "receipt.jpg", - table_format="github", +from fast_lightonocr import LightOnOCR + +model = LightOnOCR.from_pretrained( + "onnx-community/LightOnOCR-2-1B-ONNX", + preset="q4", + runtime_kwargs={ + "execution_provider": "cuda", + "device_id": 0, + }, + generation_kwargs={ + "max_new_tokens": 1024, + "do_sample": False, + }, ) -``` -Any table format supported by `tabulate` may be used. +result = model.process("receipt.jpg") +print(result.text) +``` ---- +When `execution_provider="cuda"`, `from_pretrained` preloads the pip NVIDIA +CUDA/cuDNN libraries (`onnxruntime.preload_dlls`). CPU loads never take that +path. Autoregressive decode keeps KV past/present on the GPU after the first +step (IoBinding); token sampling still runs on the host. -## ⚙️ Model Presets +If CUDA EP registration fails with a missing `libcublasLt` / provider `.so`, +add the pip `nvidia/*/lib` directories and the driver (`libcuda`) to +`LD_LIBRARY_PATH` for that process (common on some notebook runtimes). -`from_pretrained()` supports three ONNX model presets. +CPU thread tuning: ```python model = LightOnOCR.from_pretrained( "onnx-community/LightOnOCR-2-1B-ONNX", - preset="q4", + runtime_kwargs={ + "execution_provider": "cpu", + "intra_threads": 8, + }, ) ``` -Available presets: - -- `default` -- `fp16` -- `q4` - ### Generation overrides Model defaults come from Hugging Face `generation_config.json` (typically -`do_sample=True`, `temperature=0.2`, `top_k=0`, `top_p=0.9`). +`do_sample=True`, `temperature=0.2`, `top_k=0`, `top_p=0.9`). -Override them at load time with `generation_kwargs` (merged onto the decoder config; unknown keys raise `ValueError`): +Override them at load time with `generation_kwargs` (merged onto the decoder +config; unknown keys raise `ValueError`): ```python -# Faster / deterministic OCR on CPU (greedy decoding) +# Faster / deterministic OCR (greedy decoding) model = LightOnOCR.from_pretrained( "onnx-community/LightOnOCR-2-1B-ONNX", preset="q4", @@ -220,13 +223,60 @@ Bare `max_new_tokens=` remains supported as a shorthand: model = LightOnOCR.from_pretrained("...", max_new_tokens=1024) ``` -> On CPU, prefer `do_sample=False` for throughput. +On CPU, prefer `do_sample=False` for throughput. If you need sampling, set a +modest `top_k` (for example `50`) instead of leaving the HF default `top_k=0`. + +--- + +## OCR Results + +The raw model output is available through `result.text`. + +```python +print(result.text) +``` + +The Python bindings also expose a parsed document representation that extracts +embedded HTML tables while preserving the original document structure. + +```python +print(result.document) +``` + +Tables can be accessed directly: + +```python +for table in result.tables: + print(table.text_rows) +``` + +--- + +## Table Rendering -> If you need sampling, set a modest `top_k` (for example `50`) instead of leaving the HF default `top_k=0`. +By default, tables are rendered using ASCII borders. + +```python +result = model.process( + "receipt.jpg", + table_format="grid", +) +``` + +Markdown tables are also supported. + +```python +result = model.process( + "receipt.jpg", + table_format="github", +) +``` + +Any table format supported by `tabulate` may be used. --- -## 🛠 Development +## Development Install the project and development dependencies: @@ -261,6 +311,8 @@ poetry run pip wheel . --wheel-dir dist # CUDA BUILD_PROFILE=cuda poetry run pip wheel . --wheel-dir dist +# then install the wheel with the CUDA extra, e.g. +# pip install "dist/fast_lightonocr--*.whl[cuda]" ``` > **Note** @@ -272,10 +324,10 @@ BUILD_PROFILE=cuda poetry run pip wheel . --wheel-dir dist --- -## 🙏 Acknowledgements +## Acknowledgements This package wraps the native Rust **Fast LightOnOCR** inference engine and uses the open-weight **LightOnOCR** model released by Baidu. -- 🤗 https://huggingface.co/onnx-community/LightOnOCR-2-1B-ONNX -- 💻 https://github.com/baidu/LightOnOCR +- https://huggingface.co/onnx-community/LightOnOCR-2-1B-ONNX +- https://github.com/baidu/LightOnOCR diff --git a/docs/ARCHITECTURE.md b/docs/ARCHITECTURE.md index a79fd32..35d1dc5 100644 --- a/docs/ARCHITECTURE.md +++ b/docs/ARCHITECTURE.md @@ -259,9 +259,11 @@ Responsibilities include: - token selection - stopping criteria -During autoregressive generation the decoder keeps one `KvCache` and overwrites -each layer's key/value buffers in place after every step, reusing allocation -capacity instead of allocating a fresh cache from the full `present.*` outputs. +During autoregressive generation the decoder selects a [`KVCacheBackend`](KV.md) +once (`KVCache` on CPU, `CudaKVCache` by default on the CUDA EP). The host +`KVCache` overwrites each layer's key/value buffers in place after every step, +reusing allocation capacity instead of allocating a fresh cache from the full +`present.*` outputs. See [`KV.md`](KV.md) for the CUDA IoBinding path. The generation loop also reuses a single-token `InputEmbeddings` buffer via `EmbeddingModel::embed_into`, copies only the final logits position into a @@ -291,9 +293,14 @@ Sessions are configured through `RuntimeOptions`: execution provider, intra-op t The default execution provider is CPU. CUDA is available when the crate is built with `--features cuda` and a CUDA-enabled ONNX Runtime is present. Selecting CUDA without that feature fails at session creation with a clear error; with the feature enabled, CUDA EP registration fails hard if the provider cannot initialize (no silent CPU fallback). -On the CUDA path, autoregressive decode still uses the host `KvCache` and `Session::run` by default (graphs may run on GPU via the CUDA EP; KV traffic crosses the host). An experimental IoBinding device-KV path is available behind `FAST_LIGHTONOCR_CUDA_DEVICE_KV=1`; empty past tensors start on the host and present outputs are promoted to device after the first step. Token sampling always runs on the host from final-position logits. CUDA helpers are compile-gated behind the `cuda` feature so CPU builds are unchanged. +On the CUDA path, autoregressive decode defaults to `CudaKVCache` with ORT +IoBinding so past/present stay on device after prefill (empty `seq=0` past +starts on the host; see [`KV.md`](KV.md)). Set `FAST_LIGHTONOCR_CUDA_HOST_KV` +to force the host `KVCache` + `Session::run` path for debugging. Token sampling +always runs on the host from final-position logits. CUDA helpers are +compile-gated behind the `cuda` feature so CPU builds are unchanged. -The runtime provides lightweight, strongly typed wrappers around the exported ONNX models, exposing Rust domain types such as `ImageTensor`, `ImageFeatures`, `InputEmbeddings`, `AttentionMask`, `Logits`, and `KvCache` rather than raw ONNX tensors. +The runtime provides lightweight, strongly typed wrappers around the exported ONNX models, exposing Rust domain types such as `ImageTensor`, `ImageFeatures`, `InputEmbeddings`, `AttentionMask`, `Logits`, and `KVCache` rather than raw ONNX tensors. Runtime behavior—including model selection, execution providers, and session configuration—is encapsulated behind the inference engine, allowing higher-level components to remain independent of ONNX Runtime implementation details. diff --git a/docs/KV.md b/docs/KV.md new file mode 100644 index 0000000..e4d067f --- /dev/null +++ b/docs/KV.md @@ -0,0 +1,115 @@ +# Decoder KV Cache (Host and CUDA) + +This document describes the host and CUDA KV-cache design for the decoder. + +High-level architecture still lives in [`ARCHITECTURE.md`](ARCHITECTURE.md). +ONNX tensor contracts live in [`MODEL_CONTRACTS.md`](MODEL_CONTRACTS.md). + +--- + +## Goals + +- Keep the **host** [`KVCache`](../src/model/decoder/kv_cache.rs) strategy + correct, simple, and the default for CPU. +- Provide a **CUDA-only** device-resident KV path that does not make + `KVCache` device-aware. +- Prefer **one generate loop** and a clear step abstraction over + copy-pasted control flow. +- Optimize GPU decode on the CUDA backend without changing host semantics. + +--- + +## Naming + +| Name | Kind | Role | +|------|------|------| +| `KVCacheBackend` | `pub(crate)` trait | Pluggable past/present strategy used by generate / step | +| `KVCache` | public struct | **Host-resident** backend (CPU path; also public `decode`). Implements `KVCacheBackend`. | +| `CudaKVCache` | `pub(crate)` struct in `cuda_backend` (`cfg(cuda)`) | **CUDA-resident** backend for IoBinding decode | +| `ActiveKVCache` | `pub(crate)` enum | Selected backend for one generate run (`Host` / `Cuda`) | +| `CudaIoContext` | `pub(crate)` struct in `cuda_backend` (`cfg(cuda)`) | CUDA/CPU `MemoryInfo` plus reusable embeds/mask staging | + +Factory: + +```text +create_kv_cache_backend(ep, batch) + Cuda EP and FAST_LIGHTONOCR_CUDA_HOST_KV unset → ActiveKVCache::Cuda(CudaKVCache) + otherwise → ActiveKVCache::Host(KVCache) +``` + +--- + +## Design + +### One `generate_streaming` loop + +```text +kv = create_kv_cache_backend(ep, batch) +cuda_io = Some(CudaIoContext) if Cuda backend else None +for step in 0..max_new_tokens: + logits = decode_step(input, mask, &mut kv, FinalPosition, cuda_io) + token = next_token(logits) + embed_into(token) ; update mask +``` + +### One step entry, two strategies + +- Shared: shape checks (host), logits extraction / final-position materialization. +- `KVCache` → `decode_step_host`: `Session::run` + present→past host buffer update. +- `CudaKVCache` → `decode_step_cuda`: IoBinding + promote device present; embeds/mask + copied through `CudaIoContext` staging buffers (capacity retained across steps). + +Public [`Decoder::decode`](../src/model/decoder/decoder.rs) always uses host `KVCache`. + +### Prefill contract (`CudaKVCache`) + +Empty past (`seq=0`) is allocated on the **host** on purpose: zero-length CUDA +past tensors + IoBinding have segfaulted on some stacks (e.g. Colab). After the +first step, `promote_present` keeps present outputs on device. + +```text +Prefill (step 0): host empty past ──IoBinding──► device present ──promote──► device past +Decode (step 1+): device past ──IoBinding──► device present ──promote──► device past + host embeds/mask host logits (sample) +``` + +Sampling always runs on the host from final-position logits. + +--- + +## Escape hatches + +| Env | Effect | +|-----|--------| +| `FAST_LIGHTONOCR_CUDA_HOST_KV` | Force host `KVCache` even when EP is CUDA (debug / parity) | + +--- + +## Validation matrix + +| Build | EP | KV backend | Expect | +|-------|----|------------|--------| +| no `cuda` feature | CPU | host | unchanged host path | +| `--features cuda` | CPU | host | host path | +| `--features cuda` | CUDA | `CudaKVCache` (default) | correct OCR; GPU memory in use | +| `--features cuda` + `FAST_LIGHTONOCR_CUDA_HOST_KV` | CUDA | host | correct OCR; KV traffic via host | + +--- + +## Follow-ons (not blocking) + +- Confirm vision/embed sessions stay on CUDA EP without surprise host copies. +- Device-side logits / sampling only if product needs it. +- Batch size > 1 for `CudaKVCache` if required later. +- Reuse a single IoBinding object across steps if ORT API allows safely. + +--- + +## References + +- Implementation: [`src/model/decoder/decoder.rs`](../src/model/decoder/decoder.rs), + [`src/model/decoder/kv_cache.rs`](../src/model/decoder/kv_cache.rs), + [`src/model/decoder/cuda_backend.rs`](../src/model/decoder/cuda_backend.rs) + (`--features cuda` only) +- Contracts: past/present shapes in [`MODEL_CONTRACTS.md`](MODEL_CONTRACTS.md) +- Roadmap EP section: [`ROADMAP.md`](ROADMAP.md) diff --git a/docs/MODEL_CONTRACTS.md b/docs/MODEL_CONTRACTS.md index fde38f1..2f7e126 100644 --- a/docs/MODEL_CONTRACTS.md +++ b/docs/MODEL_CONTRACTS.md @@ -158,6 +158,14 @@ Each decoder layer receives: The exported model contains **28 decoder layers**, each with one key tensor and one value tensor. +#### Prefill (`past_sequence_length == 0`) + +The first decoder call uses empty past tensors with sequence length `0`. On the +CUDA IoBinding path those empty pasts are allocated on the **host**; after the +first step, `present.*` outputs remain on device and become the next step's +past. See [`KV.md`](KV.md). Host `KVCache` keeps empty `Vec` buffers for +the same shape. + ## Outputs ### Logits diff --git a/docs/ROADMAP.md b/docs/ROADMAP.md index bd08b80..3ca7b0d 100644 --- a/docs/ROADMAP.md +++ b/docs/ROADMAP.md @@ -95,7 +95,10 @@ changing model contracts. ### Completed - ✅ CUDA EP registration and session wiring (`--features cuda`) -- ✅ CUDA EP generate path (default: host KV + `Session::run`; opt-in IoBinding device KV via `FAST_LIGHTONOCR_CUDA_DEVICE_KV`) +- ✅ `KVCacheBackend` abstraction (`KVCache` / `CudaKVCache`) with a single + `generate_streaming` loop (see [`KV.md`](KV.md)) +- ✅ CUDA EP generate defaults to IoBinding `CudaKVCache`; escape hatch + `FAST_LIGHTONOCR_CUDA_HOST_KV` - ✅ Expose EP / thread options through Python `runtime_kwargs` - ✅ Packaging notes for GPU / accelerator runtimes (wheels stay CPU-default) @@ -103,6 +106,8 @@ changing model contracts. - CoreML EP (macOS / Apple Silicon) - DirectML EP (Windows) +- CUDA follow-ons (non-blocking): vision/embed residency checks, IoBinding + object reuse across steps, batch>1 on `CudaKVCache` if needed --- diff --git a/examples/inference_bench.rs b/examples/inference_bench.rs index b6fbb41..3155e1f 100644 --- a/examples/inference_bench.rs +++ b/examples/inference_bench.rs @@ -6,6 +6,10 @@ //! cargo run --release --features load-dynamic --example inference_bench -- \ //! models/lightonocr examples/SROIE-receipt.jpeg q4,default,fp16 //! ``` +//! +//! For CUDA (device KV by default), rebuild with `--features load-dynamic,cuda` +//! and pass a CUDA provider through your usual runtime options / example args. +//! Compare against `FAST_LIGHTONOCR_CUDA_HOST_KV=1` to measure host-KV overhead. use std::cmp::Ordering; use std::path::Path; diff --git a/src/model/decoder/cuda_backend.rs b/src/model/decoder/cuda_backend.rs new file mode 100644 index 0000000..ad03782 --- /dev/null +++ b/src/model/decoder/cuda_backend.rs @@ -0,0 +1,406 @@ +//! CUDA-resident KV cache and IoBinding decode step. +//! +//! Compiled only with `--features cuda`. Host [`super::KVCache`] stays +//! device-unaware; this module is the CUDA implementation of +//! [`super::KVCacheBackend`]. See [`docs/KV.md`](../../../docs/KV.md). + +use ort::memory::{AllocationDevice, AllocatorType, MemoryInfo, MemoryType}; +use ort::session::Session; +use ort::value::{DynValue, Tensor, TensorElementType, ValueType}; + +use crate::model::InputEmbeddings; +use crate::{Error, Result}; + +use super::attention::AttentionMask; +use super::config::DecoderConfig; +use super::decoder::{ + DECODER_ATTENTION_MASK_NAME, DECODER_INPUT_EMBEDS_NAME, DECODER_LOGITS_NAME, + DECODER_USE_CACHE_BRANCH_NAME, LogitsSelection, extract_logits, past_key_name, past_value_name, + present_key_name, present_value_name, +}; +use super::kv_cache::KVCacheBackend; +use super::logits::Logits; + +/// CUDA-resident KV past tensors for IoBinding decode. +/// +/// Prefill starts with host-side empty (`seq=0`) past tensors; after the first +/// step, [`Self::promote_present`] keeps present outputs on device. +pub(crate) struct CudaKVCache { + past_keys: Vec, + past_values: Vec, + past_sequence_length: usize, + batch_size: usize, +} + +impl CudaKVCache { + /// Creates empty `(batch, kv_heads, 0, head_dim)` past tensors on the **host**. + /// + /// Zero-length past tensors are allocated on CPU on purpose: CUDA allocations + /// with a zero sequence dimension are unreliable with ORT IoBinding and have + /// been observed to segfault. After the first decode step, [`Self::promote_present`] + /// replaces these with device-resident present tensors. + pub(crate) fn empty(config: &DecoderConfig, batch_size: usize) -> Result { + if batch_size == 0 { + return Err(Error::InvalidKVCache { + reason: "batch size must be greater than zero".to_owned(), + }); + } + + let shape = [ + batch_size as i64, + config.num_key_value_heads as i64, + 0_i64, + config.head_dim as i64, + ]; + + let mut past_keys = Vec::with_capacity(config.num_hidden_layers); + let mut past_values = Vec::with_capacity(config.num_hidden_layers); + for _ in 0..config.num_hidden_layers { + let key = Tensor::::from_array((shape, Vec::::new())).map_err(|source| { + Error::OnnxRuntimeCompatibility { + reason: format!("failed to create empty host past key: {source}"), + } + })?; + let value = + Tensor::::from_array((shape, Vec::::new())).map_err(|source| { + Error::OnnxRuntimeCompatibility { + reason: format!("failed to create empty host past value: {source}"), + } + })?; + past_keys.push(key.into_dyn()); + past_values.push(value.into_dyn()); + } + + Ok(Self { + past_keys, + past_values, + past_sequence_length: 0, + batch_size, + }) + } + + pub(crate) fn batch_size(&self) -> usize { + self.batch_size + } + + pub(crate) fn past_sequence_length(&self) -> usize { + self.past_sequence_length + } + + pub(crate) fn is_empty(&self) -> bool { + self.past_sequence_length == 0 + } + + fn past_keys(&self) -> &[DynValue] { + &self.past_keys + } + + fn past_values(&self) -> &[DynValue] { + &self.past_values + } + + /// Replaces past buffers with present outputs and advances sequence length. + fn promote_present( + &mut self, + past_keys: Vec, + past_values: Vec, + total_sequence_length: usize, + ) { + self.past_keys = past_keys; + self.past_values = past_values; + self.past_sequence_length = total_sequence_length; + } +} + +impl KVCacheBackend for CudaKVCache { + fn batch_size(&self) -> usize { + self.batch_size + } + + fn past_sequence_length(&self) -> usize { + self.past_sequence_length + } + + fn is_empty(&self) -> bool { + self.past_sequence_length == 0 + } +} + +/// ORT memory infos and reusable host staging buffers for CUDA IoBinding decode. +pub(crate) struct CudaIoContext { + cuda_mem: MemoryInfo<'static>, + cpu_mem: MemoryInfo<'static>, + embeds_staging: Vec, + mask_staging: Vec, +} + +impl CudaIoContext { + pub(crate) fn new(device_id: i32) -> Result { + Ok(Self { + cuda_mem: cuda_memory_info(device_id)?, + cpu_mem: cpu_output_memory_info()?, + embeds_staging: Vec::new(), + mask_staging: Vec::new(), + }) + } + + /// Copies `data` into a reusable staging buffer, then builds an owned ORT tensor. + /// + /// Staging capacity is retained across steps so host reallocations stay rare. + fn embeds_tensor(&mut self, shape: [i64; 3], data: &[f32]) -> Result> { + self.embeds_staging.clear(); + self.embeds_staging.extend_from_slice(data); + Tensor::::from_array((shape, self.embeds_staging.clone())) + .map_err(|source| Error::DecoderTensorCreation { source }) + } + + /// Copies mask data into a reusable staging buffer, then builds an owned ORT tensor. + fn mask_tensor(&mut self, shape: [i64; 2], data: &[i64]) -> Result> { + self.mask_staging.clear(); + self.mask_staging.extend_from_slice(data); + Tensor::::from_array((shape, self.mask_staging.clone())) + .map_err(|source| Error::DecoderTensorCreation { source }) + } +} + +fn cuda_memory_info(device_id: i32) -> Result> { + MemoryInfo::new( + AllocationDevice::CUDA, + device_id, + AllocatorType::Device, + MemoryType::Default, + ) + .map_err(|source| Error::OnnxRuntimeCompatibility { + reason: format!("failed to create CUDA MemoryInfo: {source}"), + }) +} + +fn cpu_output_memory_info() -> Result> { + MemoryInfo::new( + AllocationDevice::CPU, + 0, + AllocatorType::Device, + MemoryType::CPUOutput, + ) + .map_err(|source| Error::OnnxRuntimeCompatibility { + reason: format!("failed to create CPU output MemoryInfo: {source}"), + }) +} + +fn validate_device_cache_output( + value: &DynValue, + name: &str, + batch_size: usize, + total_sequence_length: usize, + config: &DecoderConfig, +) -> Result<()> { + let ValueType::Tensor { ty, shape, .. } = value.dtype() else { + return Err(Error::InvalidDecoderOutput { + reason: format!("`{name}` is not a tensor"), + }); + }; + if *ty != TensorElementType::Float32 { + return Err(Error::InvalidDecoderOutput { + reason: format!("`{name}` has element type {ty:?}, expected Float32"), + }); + } + if shape.len() != 4 { + return Err(Error::InvalidDecoderOutput { + reason: format!("`{name}` has rank {}, expected 4", shape.len()), + }); + } + let expected = [ + batch_size as i64, + config.num_key_value_heads as i64, + total_sequence_length as i64, + config.head_dim as i64, + ]; + for (axis, expected_dim) in expected.into_iter().enumerate() { + if shape[axis] != expected_dim { + return Err(Error::InvalidDecoderOutput { + reason: format!( + "`{name}` dimension {axis} is {}, expected {expected_dim}", + shape[axis] + ), + }); + } + } + Ok(()) +} + +/// Inputs for one CUDA IoBinding decoder step. +pub(crate) struct CudaDecodeStep<'a> { + pub session: &'a mut Session, + pub config: &'a DecoderConfig, + pub has_cache_branch: bool, + pub input_embeddings: &'a InputEmbeddings, + pub attention_mask: &'a AttentionMask, + pub kv_cache: &'a mut CudaKVCache, + pub cuda_io: &'a mut CudaIoContext, + pub logits_selection: LogitsSelection, + pub logits_scratch: &'a mut Vec, +} + +/// IoBinding decode: past/present stay on device after prefill; logits on host. +pub(crate) fn decode_step(step: CudaDecodeStep<'_>) -> Result { + let CudaDecodeStep { + session, + config, + has_cache_branch, + input_embeddings, + attention_mask, + kv_cache, + cuda_io, + logits_selection, + logits_scratch, + } = step; + let (batch_size, sequence_length, hidden_size) = input_embeddings.shape(); + if batch_size != kv_cache.batch_size() { + return Err(Error::InvalidDecoderInput { + reason: format!( + "input embeddings batch size is {batch_size}, expected {}", + kv_cache.batch_size() + ), + }); + } + if hidden_size != config.hidden_size { + return Err(Error::InvalidDecoderInput { + reason: format!( + "input embeddings hidden size is {hidden_size}, expected {}", + config.hidden_size + ), + }); + } + + let total_sequence_length = kv_cache + .past_sequence_length() + .checked_add(sequence_length) + .ok_or_else(|| Error::InvalidDecoderInput { + reason: "total sequence length is too large".to_owned(), + })?; + let expected_attention_values = + batch_size + .checked_mul(total_sequence_length) + .ok_or_else(|| Error::InvalidDecoderInput { + reason: "attention mask shape is too large".to_owned(), + })?; + if attention_mask.len() != expected_attention_values { + return Err(Error::InvalidDecoderInput { + reason: format!( + "attention mask length is {}, expected {expected_attention_values}", + attention_mask.len() + ), + }); + } + + let use_cache_branch = !kv_cache.is_empty(); + + let embeds = cuda_io.embeds_tensor( + [ + batch_size as i64, + sequence_length as i64, + hidden_size as i64, + ], + input_embeddings.as_slice(), + )?; + let mask = cuda_io.mask_tensor( + [batch_size as i64, total_sequence_length as i64], + attention_mask.as_slice(), + )?; + + let mut binding = session + .create_binding() + .map_err(|source| Error::DecoderInference { source })?; + + binding + .bind_input(DECODER_INPUT_EMBEDS_NAME, &embeds) + .map_err(|source| Error::DecoderInference { source })?; + binding + .bind_input(DECODER_ATTENTION_MASK_NAME, &mask) + .map_err(|source| Error::DecoderInference { source })?; + + let cache_branch_tensor = if has_cache_branch { + let tensor = Tensor::from_array(((), vec![use_cache_branch])) + .map_err(|source| Error::DecoderTensorCreation { source })?; + binding + .bind_input(DECODER_USE_CACHE_BRANCH_NAME, &tensor) + .map_err(|source| Error::DecoderInference { source })?; + Some(tensor) + } else { + None + }; + let _cache_branch_tensor = cache_branch_tensor; + + for layer_index in 0..config.num_hidden_layers { + binding + .bind_input( + past_key_name(layer_index), + &kv_cache.past_keys()[layer_index], + ) + .map_err(|source| Error::DecoderInference { source })?; + binding + .bind_input( + past_value_name(layer_index), + &kv_cache.past_values()[layer_index], + ) + .map_err(|source| Error::DecoderInference { source })?; + } + + binding + .bind_output_to_device(DECODER_LOGITS_NAME, &cuda_io.cpu_mem) + .map_err(|source| Error::DecoderInference { source })?; + for layer_index in 0..config.num_hidden_layers { + binding + .bind_output_to_device(present_key_name(layer_index), &cuda_io.cuda_mem) + .map_err(|source| Error::DecoderInference { source })?; + binding + .bind_output_to_device(present_value_name(layer_index), &cuda_io.cuda_mem) + .map_err(|source| Error::DecoderInference { source })?; + } + + let mut outputs = session + .run_binding(&binding) + .map_err(|source| Error::DecoderInference { source })?; + + let logits = extract_logits( + &mut outputs, + batch_size, + sequence_length, + config.vocab_size, + logits_selection, + logits_scratch, + )?; + + let mut next_keys = Vec::with_capacity(config.num_hidden_layers); + let mut next_values = Vec::with_capacity(config.num_hidden_layers); + for layer_index in 0..config.num_hidden_layers { + let key_name = present_key_name(layer_index); + let key = outputs + .remove(key_name.as_str()) + .ok_or_else(|| Error::InvalidDecoderOutput { + reason: format!("missing `{key_name}` output"), + })?; + validate_device_cache_output(&key, &key_name, batch_size, total_sequence_length, config)?; + + let value_name = present_value_name(layer_index); + let value = + outputs + .remove(value_name.as_str()) + .ok_or_else(|| Error::InvalidDecoderOutput { + reason: format!("missing `{value_name}` output"), + })?; + validate_device_cache_output( + &value, + &value_name, + batch_size, + total_sequence_length, + config, + )?; + + next_keys.push(key); + next_values.push(value); + } + + kv_cache.promote_present(next_keys, next_values, total_sequence_length); + Ok(logits) +} diff --git a/src/model/decoder/decoder.rs b/src/model/decoder/decoder.rs index 90dbc85..fd42007 100644 --- a/src/model/decoder/decoder.rs +++ b/src/model/decoder/decoder.rs @@ -5,11 +5,6 @@ use std::path::{Path, PathBuf}; use ort::session::{Session, SessionInputValue}; use ort::value::{Outlet, TensorElementType, TensorRef, ValueType}; -#[cfg(feature = "cuda")] -use ort::memory::MemoryInfo; -#[cfg(feature = "cuda")] -use ort::value::Tensor; - use crate::model::InputEmbeddings; use crate::model::embedding_model::EmbeddingModel; use crate::util::{ExecutionProvider, RuntimeOptions}; @@ -17,26 +12,24 @@ use crate::{Error, Result}; use super::attention::AttentionMask; use super::config::{DecoderConfig, GenerationConfig}; -use super::generation::{self, FinishReason, GenerationOutput}; #[cfg(feature = "cuda")] -use super::kv_cache::{ - CudaKvState, cpu_output_memory_info, cuda_memory_info, validate_device_cache_output, -}; -use super::kv_cache::{KvCache, values_per_tensor}; +use super::cuda_backend::{self, CudaIoContext, CudaKVCache}; +use super::generation::{self, FinishReason, GenerationOutput}; +use super::kv_cache::{ActiveKVCache, KVCache, KVCacheBackend, values_per_tensor}; use super::logits::Logits; use super::output::DecoderOutput; -const DECODER_INPUT_EMBEDS_NAME: &str = "inputs_embeds"; -const DECODER_ATTENTION_MASK_NAME: &str = "attention_mask"; -const DECODER_USE_CACHE_BRANCH_NAME: &str = "use_cache_branch"; -const DECODER_LOGITS_NAME: &str = "logits"; +pub(crate) const DECODER_INPUT_EMBEDS_NAME: &str = "inputs_embeds"; +pub(crate) const DECODER_ATTENTION_MASK_NAME: &str = "attention_mask"; +pub(crate) const DECODER_USE_CACHE_BRANCH_NAME: &str = "use_cache_branch"; +pub(crate) const DECODER_LOGITS_NAME: &str = "logits"; const CONFIG_FILE: &str = "config.json"; const GENERATION_CONFIG_FILE: &str = "generation_config.json"; /// Which logits positions to materialize on the host after a decoder step. #[derive(Debug, Clone, Copy, PartialEq, Eq)] -enum LogitsSelection { +pub(crate) enum LogitsSelection { /// Copy the full `(batch, sequence, vocab)` logits tensor. Full, /// Copy only the final sequence position (shape becomes `(batch, 1, vocab)`). @@ -148,9 +141,30 @@ impl Decoder { self.execution_provider } - /// Creates an empty KV cache for this decoder and batch size. - pub fn empty_kv_cache(&self, batch_size: usize) -> Result { - KvCache::empty(&self.config, batch_size) + /// Creates an empty host-resident KV cache for this decoder and batch size. + pub fn empty_kv_cache(&self, batch_size: usize) -> Result { + KVCache::empty(&self.config, batch_size) + } + + /// Selects the KV backend for autoregressive generation. + /// + /// CUDA EP defaults to [`CudaKVCache`]. Set `FAST_LIGHTONOCR_CUDA_HOST_KV` + /// to force the host [`KVCache`] path for debugging. + fn create_kv_cache_backend(&self, batch_size: usize) -> Result { + #[cfg(feature = "cuda")] + if let ExecutionProvider::Cuda { .. } = self.execution_provider + && std::env::var_os("FAST_LIGHTONOCR_CUDA_HOST_KV").is_none() + { + return Ok(ActiveKVCache::Cuda(CudaKVCache::empty( + &self.config, + batch_size, + )?)); + } + + Ok(ActiveKVCache::Host(KVCache::empty( + &self.config, + batch_size, + )?)) } /// Returns the image placeholder token id used by the decoder config. @@ -181,82 +195,39 @@ impl Decoder { return Ok(GenerationOutput::new(Vec::new(), FinishReason::Length)); } - // Optional IoBinding device-KV path. Default CUDA EP uses the host KV - // loop below (Session::run): zero-length CUDA past tensors have been - // observed to segfault with ORT IoBinding on some stacks (e.g. Colab). - // Enable with FAST_LIGHTONOCR_CUDA_DEVICE_KV=1 for continued debugging. - #[cfg(feature = "cuda")] - if let ExecutionProvider::Cuda { device_id } = self.execution_provider - && std::env::var_os("FAST_LIGHTONOCR_CUDA_DEVICE_KV").is_some() - { - let device_id = - i32::try_from(device_id).map_err(|_| Error::OnnxRuntimeCompatibility { - reason: format!("CUDA device_id {device_id} is out of range for i32"), - })?; - - let mut generated = Vec::with_capacity(self.generation_config.max_new_tokens); - let mut kv_state = CudaKvState::empty(&self.config, decoder_input.batch_size())?; - attention_mask.reserve(self.generation_config.max_new_tokens); - - let hidden_size = decoder_input.hidden_size(); - let mut step_input = InputEmbeddings::new(vec![0.0; hidden_size], 1, 1, hidden_size)?; - let mut using_step_input = false; - let mut finish_reason = FinishReason::Length; - let cuda_mem = cuda_memory_info(device_id)?; - let cpu_mem = cpu_output_memory_info()?; - - for step in 0..self.generation_config.max_new_tokens { - let input = if using_step_input { - &step_input - } else { - &decoder_input - }; - let logits = self.decode_step_cuda( - input, - &attention_mask, - &mut kv_state, - &cuda_mem, - &cpu_mem, - LogitsSelection::FinalPosition, - )?; - - let next_token = - generation::next_token(&self.generation_config, &logits, &mut self.rng)?; - self.logits_scratch = logits.into_data(); - - generated.push(next_token); - on_token(next_token); - - if generation::is_eos(&self.generation_config, next_token) { - finish_reason = FinishReason::EndOfSequence; - break; - } - - if step + 1 == self.generation_config.max_new_tokens { - break; - } - - embedding_model.embed_into(&[next_token], &mut step_input)?; - using_step_input = true; - if step == 0 { - decoder_input = InputEmbeddings::default(); - attention_mask.fill_visible(); - } - attention_mask.push_visible(); - } - - return Ok(GenerationOutput::new(generated, finish_reason)); - } - let mut generated = Vec::with_capacity(self.generation_config.max_new_tokens); - - let mut kv_cache = self.empty_kv_cache(decoder_input.batch_size())?; + let mut kv = self.create_kv_cache_backend(decoder_input.batch_size())?; + debug_assert_eq!( + KVCacheBackend::batch_size(&kv), + decoder_input.batch_size(), + "KV backend batch size must match decoder input" + ); + debug_assert!( + KVCacheBackend::is_empty(&kv), + "generate must start from an empty KV backend" + ); + debug_assert_eq!( + KVCacheBackend::past_sequence_length(&kv), + 0, + "generate must start with past_sequence_length == 0" + ); attention_mask.reserve(self.generation_config.max_new_tokens); + #[cfg(feature = "cuda")] + let mut cuda_io = match (&kv, self.execution_provider) { + (ActiveKVCache::Cuda(_), ExecutionProvider::Cuda { device_id }) => { + let device_id = + i32::try_from(device_id).map_err(|_| Error::OnnxRuntimeCompatibility { + reason: format!("CUDA device_id {device_id} is out of range for i32"), + })?; + Some(CudaIoContext::new(device_id)?) + } + _ => None, + }; + let hidden_size = decoder_input.hidden_size(); let mut step_input = InputEmbeddings::new(vec![0.0; hidden_size], 1, 1, hidden_size)?; let mut using_step_input = false; - let mut finish_reason = FinishReason::Length; for step in 0..self.generation_config.max_new_tokens { @@ -265,10 +236,20 @@ impl Decoder { } else { &decoder_input }; + + #[cfg(feature = "cuda")] + let logits = self.decode_step( + input, + &attention_mask, + &mut kv, + LogitsSelection::FinalPosition, + cuda_io.as_mut(), + )?; + #[cfg(not(feature = "cuda"))] let logits = self.decode_step( input, &attention_mask, - &mut kv_cache, + &mut kv, LogitsSelection::FinalPosition, )?; @@ -312,27 +293,91 @@ impl Decoder { &mut self, input_embeddings: &InputEmbeddings, attention_mask: &AttentionMask, - kv_cache: &KvCache, + kv_cache: &KVCache, ) -> Result { - let mut cache = kv_cache.clone(); + let mut cache = ActiveKVCache::Host(kv_cache.clone()); + #[cfg(feature = "cuda")] let logits = self.decode_step( input_embeddings, attention_mask, &mut cache, LogitsSelection::Full, + None, )?; + #[cfg(not(feature = "cuda"))] + let logits = self.decode_step( + input_embeddings, + attention_mask, + &mut cache, + LogitsSelection::Full, + )?; + #[cfg(not(feature = "cuda"))] + let ActiveKVCache::Host(cache) = cache; + #[cfg(feature = "cuda")] + let cache = match cache { + ActiveKVCache::Host(cache) => cache, + ActiveKVCache::Cuda(_) => { + unreachable!("public decode always uses the host KV backend") + } + }; Ok(DecoderOutput::new(logits, cache)) } - /// Executes one decoder pass, writing the updated KV cache in place. - /// - /// Present tensors are copied into the existing per-layer buffers so - /// allocation capacity is reused across autoregressive steps. + /// Executes one decoder pass against the selected [`ActiveKVCache`] backend. + #[cfg(feature = "cuda")] + fn decode_step( + &mut self, + input_embeddings: &InputEmbeddings, + attention_mask: &AttentionMask, + kv: &mut ActiveKVCache, + logits_selection: LogitsSelection, + cuda_io: Option<&mut CudaIoContext>, + ) -> Result { + match kv { + ActiveKVCache::Host(cache) => { + self.decode_step_host(input_embeddings, attention_mask, cache, logits_selection) + } + ActiveKVCache::Cuda(cache) => { + let cuda_io = cuda_io.ok_or_else(|| Error::OnnxRuntimeCompatibility { + reason: "CUDA KV backend requires CudaIoContext".to_owned(), + })?; + cuda_backend::decode_step(cuda_backend::CudaDecodeStep { + session: &mut self.session, + config: &self.config, + has_cache_branch: self.has_cache_branch, + input_embeddings, + attention_mask, + kv_cache: cache, + cuda_io, + logits_selection, + logits_scratch: &mut self.logits_scratch, + }) + } + } + } + + /// Executes one decoder pass against the selected [`ActiveKVCache`] backend. + #[cfg(not(feature = "cuda"))] fn decode_step( &mut self, input_embeddings: &InputEmbeddings, attention_mask: &AttentionMask, - kv_cache: &mut KvCache, + kv: &mut ActiveKVCache, + logits_selection: LogitsSelection, + ) -> Result { + match kv { + ActiveKVCache::Host(cache) => { + self.decode_step_host(input_embeddings, attention_mask, cache, logits_selection) + } + } + } + + /// Host `Session::run` path: present tensors are copied into `KVCache` buffers. + fn decode_step_host( + &mut self, + input_embeddings: &InputEmbeddings, + attention_mask: &AttentionMask, + kv_cache: &mut KVCache, logits_selection: LogitsSelection, ) -> Result { let total_sequence_length = @@ -388,44 +433,15 @@ impl Decoder { .map_err(|source| Error::DecoderInference { source })? }; - let logits_output = - outputs - .remove(DECODER_LOGITS_NAME) - .ok_or_else(|| Error::InvalidDecoderOutput { - reason: format!("missing `{DECODER_LOGITS_NAME}` output"), - })?; - let (logits_shape, logits_data) = - logits_output - .try_extract_tensor::() - .map_err(|source| Error::InvalidDecoderOutput { - reason: format!( - "failed to extract `{DECODER_LOGITS_NAME}` as float32: {source}" - ), - })?; - let logits_shape = logits_shape.as_ref(); - validate_logits_shape( - logits_shape, + let logits = extract_logits( + &mut outputs, batch_size, sequence_length, self.config.vocab_size, + logits_selection, + &mut self.logits_scratch, )?; - let logits = match logits_selection { - LogitsSelection::Full => Logits::new( - logits_data.to_vec(), - usize::try_from(logits_shape[0]).expect("validated non-negative batch size"), - usize::try_from(logits_shape[1]).expect("validated non-negative sequence length"), - usize::try_from(logits_shape[2]).expect("validated non-negative vocabulary size"), - )?, - LogitsSelection::FinalPosition => materialize_final_position_logits( - logits_data, - batch_size, - sequence_length, - self.config.vocab_size, - &mut self.logits_scratch, - )?, - }; - update_cache_from_outputs( &mut outputs, kv_cache, @@ -438,204 +454,45 @@ impl Decoder { Ok(logits) } +} - /// CUDA IoBinding decode step: past/present stay on device; logits on host. - #[cfg(feature = "cuda")] - fn decode_step_cuda( - &mut self, - input_embeddings: &InputEmbeddings, - attention_mask: &AttentionMask, - kv_state: &mut CudaKvState, - cuda_mem: &MemoryInfo<'_>, - cpu_mem: &MemoryInfo<'_>, - logits_selection: LogitsSelection, - ) -> Result { - let (batch_size, sequence_length, hidden_size) = input_embeddings.shape(); - if batch_size != kv_state.batch_size { - return Err(Error::InvalidDecoderInput { - reason: format!( - "input embeddings batch size is {batch_size}, expected {}", - kv_state.batch_size - ), - }); - } - if hidden_size != self.config.hidden_size { - return Err(Error::InvalidDecoderInput { - reason: format!( - "input embeddings hidden size is {hidden_size}, expected {}", - self.config.hidden_size - ), - }); - } - - let total_sequence_length = kv_state - .past_sequence_length - .checked_add(sequence_length) - .ok_or_else(|| Error::InvalidDecoderInput { - reason: "total sequence length is too large".to_owned(), +pub(crate) fn extract_logits( + outputs: &mut ort::session::SessionOutputs<'_>, + batch_size: usize, + sequence_length: usize, + vocab_size: usize, + logits_selection: LogitsSelection, + logits_scratch: &mut Vec, +) -> Result { + let logits_output = + outputs + .remove(DECODER_LOGITS_NAME) + .ok_or_else(|| Error::InvalidDecoderOutput { + reason: format!("missing `{DECODER_LOGITS_NAME}` output"), })?; - let expected_attention_values = - batch_size - .checked_mul(total_sequence_length) - .ok_or_else(|| Error::InvalidDecoderInput { - reason: "attention mask shape is too large".to_owned(), - })?; - if attention_mask.len() != expected_attention_values { - return Err(Error::InvalidDecoderInput { - reason: format!( - "attention mask length is {}, expected {expected_attention_values}", - attention_mask.len() - ), - }); - } - - let use_cache_branch = !kv_state.is_empty(); - - let embeds = Tensor::from_array(( - [ - batch_size as i64, - sequence_length as i64, - hidden_size as i64, - ], - input_embeddings.as_slice().to_vec(), - )) - .map_err(|source| Error::DecoderTensorCreation { source })?; - let mask = Tensor::from_array(( - [batch_size as i64, total_sequence_length as i64], - attention_mask.as_slice().to_vec(), - )) - .map_err(|source| Error::DecoderTensorCreation { source })?; - - let mut binding = self - .session - .create_binding() - .map_err(|source| Error::DecoderInference { source })?; - - binding - .bind_input(DECODER_INPUT_EMBEDS_NAME, &embeds) - .map_err(|source| Error::DecoderInference { source })?; - binding - .bind_input(DECODER_ATTENTION_MASK_NAME, &mask) - .map_err(|source| Error::DecoderInference { source })?; - - let cache_branch_tensor = if self.has_cache_branch { - let tensor = Tensor::from_array(((), vec![use_cache_branch])) - .map_err(|source| Error::DecoderTensorCreation { source })?; - binding - .bind_input(DECODER_USE_CACHE_BRANCH_NAME, &tensor) - .map_err(|source| Error::DecoderInference { source })?; - Some(tensor) - } else { - None - }; - let _cache_branch_tensor = cache_branch_tensor; - - for layer_index in 0..self.config.num_hidden_layers { - binding - .bind_input(past_key_name(layer_index), &kv_state.past_keys[layer_index]) - .map_err(|source| Error::DecoderInference { source })?; - binding - .bind_input( - past_value_name(layer_index), - &kv_state.past_values[layer_index], - ) - .map_err(|source| Error::DecoderInference { source })?; - } - - binding - .bind_output_to_device(DECODER_LOGITS_NAME, cpu_mem) - .map_err(|source| Error::DecoderInference { source })?; - for layer_index in 0..self.config.num_hidden_layers { - binding - .bind_output_to_device(present_key_name(layer_index), cuda_mem) - .map_err(|source| Error::DecoderInference { source })?; - binding - .bind_output_to_device(present_value_name(layer_index), cuda_mem) - .map_err(|source| Error::DecoderInference { source })?; - } - - let mut outputs = self - .session - .run_binding(&binding) - .map_err(|source| Error::DecoderInference { source })?; - - let logits_output = - outputs - .remove(DECODER_LOGITS_NAME) - .ok_or_else(|| Error::InvalidDecoderOutput { - reason: format!("missing `{DECODER_LOGITS_NAME}` output"), - })?; - let (logits_shape, logits_data) = - logits_output - .try_extract_tensor::() - .map_err(|source| Error::InvalidDecoderOutput { - reason: format!( - "failed to extract `{DECODER_LOGITS_NAME}` as float32: {source}" - ), - })?; - let logits_shape = logits_shape.as_ref(); - validate_logits_shape( - logits_shape, + let (logits_shape, logits_data) = + logits_output + .try_extract_tensor::() + .map_err(|source| Error::InvalidDecoderOutput { + reason: format!("failed to extract `{DECODER_LOGITS_NAME}` as float32: {source}"), + })?; + let logits_shape = logits_shape.as_ref(); + validate_logits_shape(logits_shape, batch_size, sequence_length, vocab_size)?; + + match logits_selection { + LogitsSelection::Full => Logits::new( + logits_data.to_vec(), + usize::try_from(logits_shape[0]).expect("validated non-negative batch size"), + usize::try_from(logits_shape[1]).expect("validated non-negative sequence length"), + usize::try_from(logits_shape[2]).expect("validated non-negative vocabulary size"), + ), + LogitsSelection::FinalPosition => materialize_final_position_logits( + logits_data, batch_size, sequence_length, - self.config.vocab_size, - )?; - - let logits = match logits_selection { - LogitsSelection::Full => Logits::new( - logits_data.to_vec(), - usize::try_from(logits_shape[0]).expect("validated non-negative batch size"), - usize::try_from(logits_shape[1]).expect("validated non-negative sequence length"), - usize::try_from(logits_shape[2]).expect("validated non-negative vocabulary size"), - )?, - LogitsSelection::FinalPosition => materialize_final_position_logits( - logits_data, - batch_size, - sequence_length, - self.config.vocab_size, - &mut self.logits_scratch, - )?, - }; - - let mut next_keys = Vec::with_capacity(self.config.num_hidden_layers); - let mut next_values = Vec::with_capacity(self.config.num_hidden_layers); - for layer_index in 0..self.config.num_hidden_layers { - let key_name = present_key_name(layer_index); - let key = - outputs - .remove(key_name.as_str()) - .ok_or_else(|| Error::InvalidDecoderOutput { - reason: format!("missing `{key_name}` output"), - })?; - validate_device_cache_output( - &key, - &key_name, - batch_size, - total_sequence_length, - &self.config, - )?; - - let value_name = present_value_name(layer_index); - let value = - outputs - .remove(value_name.as_str()) - .ok_or_else(|| Error::InvalidDecoderOutput { - reason: format!("missing `{value_name}` output"), - })?; - validate_device_cache_output( - &value, - &value_name, - batch_size, - total_sequence_length, - &self.config, - )?; - - next_keys.push(key); - next_values.push(value); - } - - kv_state.promote_present(next_keys, next_values, total_sequence_length); - Ok(logits) + vocab_size, + logits_scratch, + ), } } @@ -871,7 +728,7 @@ fn validate_tensor_metadata( fn validate_decoder_inputs( input_embeddings: &InputEmbeddings, attention_mask: &AttentionMask, - kv_cache: &KvCache, + kv_cache: &KVCache, config: &DecoderConfig, ) -> Result { let (batch_size, sequence_length, hidden_size) = input_embeddings.shape(); @@ -983,7 +840,7 @@ fn validate_decoder_inputs( fn update_cache_from_outputs( outputs: &mut ort::session::SessionOutputs<'_>, - kv_cache: &mut KvCache, + kv_cache: &mut KVCache, batch_size: usize, total_sequence_length: usize, config: &DecoderConfig, @@ -1166,18 +1023,18 @@ fn validate_cache_output_shape( Ok(()) } -fn past_key_name(layer_index: usize) -> String { +pub(crate) fn past_key_name(layer_index: usize) -> String { format!("past_key_values.{layer_index}.key") } -fn past_value_name(layer_index: usize) -> String { +pub(crate) fn past_value_name(layer_index: usize) -> String { format!("past_key_values.{layer_index}.value") } -fn present_key_name(layer_index: usize) -> String { +pub(crate) fn present_key_name(layer_index: usize) -> String { format!("present.{layer_index}.key") } -fn present_value_name(layer_index: usize) -> String { +pub(crate) fn present_value_name(layer_index: usize) -> String { format!("present.{layer_index}.value") } diff --git a/src/model/decoder/kv_cache.rs b/src/model/decoder/kv_cache.rs index 956ce20..b0afb00 100644 --- a/src/model/decoder/kv_cache.rs +++ b/src/model/decoder/kv_cache.rs @@ -1,19 +1,24 @@ //! Decoder key/value cache representation. //! -//! The host [`KvCache`] is always available. With `--features cuda`, this module -//! also provides [`CudaKvState`] for device-resident past/present tensors used -//! by the decoder IoBinding path. +//! The host [`KVCache`] is always available and implements [`KVCacheBackend`]. +//! With `--features cuda`, [`ActiveKVCache`] can hold a CUDA-resident cache +//! from `cuda_backend`. +//! +//! See [`docs/KV.md`](../../../docs/KV.md) for the host vs CUDA design. use super::DecoderConfig; - #[cfg(feature = "cuda")] -use ort::memory::{AllocationDevice, AllocatorType, MemoryInfo, MemoryType}; -#[cfg(feature = "cuda")] -use ort::value::{DynValue, Tensor, TensorElementType, ValueType}; +use super::cuda_backend::CudaKVCache; -#[cfg(feature = "cuda")] use crate::{Error, Result}; +/// Pluggable past/present KV strategy used by autoregressive decode. +pub(crate) trait KVCacheBackend { + fn batch_size(&self) -> usize; + fn past_sequence_length(&self) -> usize; + fn is_empty(&self) -> bool; +} + /// Key/value tensors for one decoder layer. /// /// Each tensor stores contiguous `float32` data with shape @@ -49,14 +54,14 @@ impl LayerCache { } } -/// Opaque decoder key/value cache passed between decoder invocations. +/// Host-resident decoder key/value cache. /// -/// `KvCache` contains one [`LayerCache`] for each decoder layer. The cache is +/// `KVCache` contains one [`LayerCache`] for each decoder layer. The cache is /// initialized empty for the first decoder pass. During autoregressive /// generation the decoder overwrites each layer's buffers in place, reusing /// allocation capacity across steps. #[derive(Debug, Clone, Default, PartialEq)] -pub struct KvCache { +pub struct KVCache { layers: Vec, batch_size: usize, past_sequence_length: usize, @@ -64,11 +69,11 @@ pub struct KvCache { head_dim: usize, } -impl KvCache { +impl KVCache { /// Creates an empty KV cache for a decoder configuration and batch size. - pub fn empty(config: &DecoderConfig, batch_size: usize) -> crate::Result { + pub fn empty(config: &DecoderConfig, batch_size: usize) -> Result { if batch_size == 0 { - return Err(crate::Error::InvalidKvCache { + return Err(Error::InvalidKVCache { reason: "batch size must be greater than zero".to_owned(), }); } @@ -93,19 +98,19 @@ impl KvCache { past_sequence_length: usize, num_key_value_heads: usize, head_dim: usize, - ) -> crate::Result { + ) -> Result { if batch_size == 0 { - return Err(crate::Error::InvalidKvCache { + return Err(Error::InvalidKVCache { reason: "batch size must be greater than zero".to_owned(), }); } if num_key_value_heads == 0 { - return Err(crate::Error::InvalidKvCache { + return Err(Error::InvalidKVCache { reason: "KV head count must be greater than zero".to_owned(), }); } if head_dim == 0 { - return Err(crate::Error::InvalidKvCache { + return Err(Error::InvalidKVCache { reason: "KV head dimension must be greater than zero".to_owned(), }); } @@ -118,7 +123,7 @@ impl KvCache { )?; for (index, layer) in layers.iter().enumerate() { if layer.key.len() != expected { - return Err(crate::Error::InvalidKvCache { + return Err(Error::InvalidKVCache { reason: format!( "layer {index} key length {} does not match shape {:?}", layer.key.len(), @@ -132,7 +137,7 @@ impl KvCache { }); } if layer.value.len() != expected { - return Err(crate::Error::InvalidKvCache { + return Err(Error::InvalidKVCache { reason: format!( "layer {index} value length {} does not match shape {:?}", layer.value.len(), @@ -210,163 +215,90 @@ impl KvCache { } } +impl KVCacheBackend for KVCache { + fn batch_size(&self) -> usize { + self.batch_size + } + + fn past_sequence_length(&self) -> usize { + self.past_sequence_length + } + + fn is_empty(&self) -> bool { + self.past_sequence_length == 0 + } +} + +/// Selected KV backend for one autoregressive generate run. +pub(crate) enum ActiveKVCache { + Host(KVCache), + #[cfg(feature = "cuda")] + Cuda(CudaKVCache), +} + +impl KVCacheBackend for ActiveKVCache { + fn batch_size(&self) -> usize { + match self { + Self::Host(cache) => cache.batch_size, + #[cfg(feature = "cuda")] + Self::Cuda(cache) => cache.batch_size(), + } + } + + fn past_sequence_length(&self) -> usize { + match self { + Self::Host(cache) => cache.past_sequence_length, + #[cfg(feature = "cuda")] + Self::Cuda(cache) => cache.past_sequence_length(), + } + } + + fn is_empty(&self) -> bool { + match self { + Self::Host(cache) => cache.past_sequence_length == 0, + #[cfg(feature = "cuda")] + Self::Cuda(cache) => cache.is_empty(), + } + } +} + pub(crate) fn values_per_tensor( batch_size: usize, sequence_length: usize, num_key_value_heads: usize, head_dim: usize, -) -> crate::Result { +) -> Result { batch_size .checked_mul(num_key_value_heads) .and_then(|value| value.checked_mul(sequence_length)) .and_then(|value| value.checked_mul(head_dim)) - .ok_or_else(|| crate::Error::InvalidKvCache { + .ok_or_else(|| Error::InvalidKVCache { reason: "KV cache tensor shape is too large".to_owned(), }) } -/// Device-resident KV past tensors for CUDA IoBinding decode. -/// -/// Compiled only with `--features cuda`. Not a public API twin of [`KvCache`]; -/// it exists so the decoder can keep past/present on GPU across steps. -#[cfg(feature = "cuda")] -pub(crate) struct CudaKvState { - pub past_keys: Vec, - pub past_values: Vec, - pub past_sequence_length: usize, - pub batch_size: usize, -} +#[cfg(test)] +mod tests { + use super::*; -#[cfg(feature = "cuda")] -impl CudaKvState { - /// Creates empty `(batch, kv_heads, 0, head_dim)` past tensors on the **host**. - /// - /// Zero-length past tensors are allocated on CPU on purpose: CUDA allocations - /// with a zero sequence dimension are unreliable with ORT IoBinding and have - /// been observed to segfault. After the first decode step, [`Self::promote_present`] - /// replaces these with device-resident present tensors. - pub(crate) fn empty(config: &DecoderConfig, batch_size: usize) -> Result { - if batch_size == 0 { - return Err(Error::InvalidKvCache { - reason: "batch size must be greater than zero".to_owned(), - }); - } - - let shape = [ - batch_size as i64, - config.num_key_value_heads as i64, - 0_i64, - config.head_dim as i64, + #[test] + fn host_kv_cache_backend_surface() { + let layers = vec![ + LayerCache::new(Vec::new(), Vec::new()), + LayerCache::new(Vec::new(), Vec::new()), ]; - - let mut past_keys = Vec::with_capacity(config.num_hidden_layers); - let mut past_values = Vec::with_capacity(config.num_hidden_layers); - for _ in 0..config.num_hidden_layers { - // Empty f32 buffer with a zero-length sequence axis. - let key = Tensor::::from_array((shape, Vec::::new())).map_err(|source| { - Error::OnnxRuntimeCompatibility { - reason: format!("failed to create empty host past key: {source}"), - } - })?; - let value = - Tensor::::from_array((shape, Vec::::new())).map_err(|source| { - Error::OnnxRuntimeCompatibility { - reason: format!("failed to create empty host past value: {source}"), - } - })?; - past_keys.push(key.into_dyn()); - past_values.push(value.into_dyn()); - } - - Ok(Self { - past_keys, - past_values, - past_sequence_length: 0, - batch_size, - }) - } - - pub(crate) fn is_empty(&self) -> bool { - self.past_sequence_length == 0 + let cache = KVCache::new(layers, 1, 0, 2, 4).unwrap(); + assert!(cache.is_empty()); + assert_eq!(KVCacheBackend::batch_size(&cache), 1); + assert_eq!(KVCacheBackend::past_sequence_length(&cache), 0); } - /// Replaces past buffers with present outputs and advances sequence length. - pub(crate) fn promote_present( - &mut self, - past_keys: Vec, - past_values: Vec, - total_sequence_length: usize, - ) { - self.past_keys = past_keys; - self.past_values = past_values; - self.past_sequence_length = total_sequence_length; - } -} - -#[cfg(feature = "cuda")] -pub(crate) fn cuda_memory_info(device_id: i32) -> Result> { - MemoryInfo::new( - AllocationDevice::CUDA, - device_id, - AllocatorType::Device, - MemoryType::Default, - ) - .map_err(|source| Error::OnnxRuntimeCompatibility { - reason: format!("failed to create CUDA MemoryInfo: {source}"), - }) -} - -#[cfg(feature = "cuda")] -pub(crate) fn cpu_output_memory_info() -> Result> { - MemoryInfo::new( - AllocationDevice::CPU, - 0, - AllocatorType::Device, - MemoryType::CPUOutput, - ) - .map_err(|source| Error::OnnxRuntimeCompatibility { - reason: format!("failed to create CPU output MemoryInfo: {source}"), - }) -} - -#[cfg(feature = "cuda")] -pub(crate) fn validate_device_cache_output( - value: &DynValue, - name: &str, - batch_size: usize, - total_sequence_length: usize, - config: &DecoderConfig, -) -> Result<()> { - let ValueType::Tensor { ty, shape, .. } = value.dtype() else { - return Err(Error::InvalidDecoderOutput { - reason: format!("`{name}` is not a tensor"), - }); - }; - if *ty != TensorElementType::Float32 { - return Err(Error::InvalidDecoderOutput { - reason: format!("`{name}` has element type {ty:?}, expected Float32"), - }); - } - if shape.len() != 4 { - return Err(Error::InvalidDecoderOutput { - reason: format!("`{name}` has rank {}, expected 4", shape.len()), - }); - } - let expected = [ - batch_size as i64, - config.num_key_value_heads as i64, - total_sequence_length as i64, - config.head_dim as i64, - ]; - for (axis, expected_dim) in expected.into_iter().enumerate() { - if shape[axis] != expected_dim { - return Err(Error::InvalidDecoderOutput { - reason: format!( - "`{name}` dimension {axis} is {}, expected {expected_dim}", - shape[axis] - ), - }); - } + #[test] + fn active_host_backend() { + let layers = vec![LayerCache::new(Vec::new(), Vec::new())]; + let cache = KVCache::new(layers, 2, 0, 1, 4).unwrap(); + let active = ActiveKVCache::Host(cache); + assert_eq!(KVCacheBackend::batch_size(&active), 2); + assert!(KVCacheBackend::is_empty(&active)); } - Ok(()) } diff --git a/src/model/decoder/mod.rs b/src/model/decoder/mod.rs index f004ea0..da78689 100644 --- a/src/model/decoder/mod.rs +++ b/src/model/decoder/mod.rs @@ -1,11 +1,13 @@ //! Decoder, autoregressive generation, attention masks, logits, and KV-cache. //! //! This module owns the decoder-specific subset of `text_config` from -//! `config.json`, the [`KvCache`] representation, the ONNX Runtime wrapper, +//! `config.json`, the [`KVCache`] representation, the ONNX Runtime wrapper, //! and the autoregressive generation engine built on top of the decoder. mod attention; mod config; +#[cfg(feature = "cuda")] +mod cuda_backend; #[allow(clippy::module_inception)] mod decoder; mod generation; @@ -17,6 +19,6 @@ pub use attention::AttentionMask; pub use config::{DecoderConfig, GenerationConfig, LayerType}; pub use decoder::Decoder; pub use generation::{FinishReason, GenerationOutput}; -pub use kv_cache::{KvCache, LayerCache}; +pub use kv_cache::{KVCache, LayerCache}; pub use logits::Logits; pub use output::DecoderOutput; diff --git a/src/model/decoder/output.rs b/src/model/decoder/output.rs index ef1c398..0d4e448 100644 --- a/src/model/decoder/output.rs +++ b/src/model/decoder/output.rs @@ -1,6 +1,6 @@ //! Decoder output value. -use super::kv_cache::KvCache; +use super::kv_cache::KVCache; use super::logits::Logits; /// Output from a single decoder invocation. @@ -10,12 +10,12 @@ pub struct DecoderOutput { pub logits: Logits, /// Updated key/value cache returned by the decoder. - pub kv_cache: KvCache, + pub kv_cache: KVCache, } impl DecoderOutput { /// Creates decoder output from logits and an updated KV cache. - pub fn new(logits: Logits, kv_cache: KvCache) -> Self { + pub fn new(logits: Logits, kv_cache: KVCache) -> Self { Self { logits, kv_cache } } } diff --git a/src/model/mod.rs b/src/model/mod.rs index 5404bed..8a8f9b1 100644 --- a/src/model/mod.rs +++ b/src/model/mod.rs @@ -23,7 +23,7 @@ pub use embedding_model::{EmbeddingConfig, EmbeddingModel}; // Decoder and generation pub use decoder::{ - Decoder, DecoderConfig, FinishReason, GenerationConfig, GenerationOutput, KvCache, LayerCache, + Decoder, DecoderConfig, FinishReason, GenerationConfig, GenerationOutput, KVCache, LayerCache, }; // High-level OCR pipeline diff --git a/src/util/mod.rs b/src/util/mod.rs index 9016fd5..2b17cc3 100644 --- a/src/util/mod.rs +++ b/src/util/mod.rs @@ -284,7 +284,7 @@ pub enum Error { /// A decoder key/value cache was malformed. #[error("invalid KV cache: {reason}")] - InvalidKvCache { + InvalidKVCache { /// Explanation of the malformed cache. reason: String, }, diff --git a/tests/decoder.rs b/tests/decoder.rs index b5c1bc2..f1f3799 100644 --- a/tests/decoder.rs +++ b/tests/decoder.rs @@ -4,7 +4,7 @@ use std::process::Command; use fast_lightonocr::Error; use fast_lightonocr::Result; use fast_lightonocr::model::Logits; -use fast_lightonocr::model::decoder::{Decoder, DecoderConfig, KvCache, LayerCache, LayerType}; +use fast_lightonocr::model::decoder::{Decoder, DecoderConfig, KVCache, LayerCache, LayerType}; use fast_lightonocr::model::{AttentionMask, InputEmbeddings}; use fast_lightonocr::model::{DataType, ModelType}; use serde::Deserialize; @@ -59,7 +59,7 @@ fn reports_missing_decoder_model() { fn initializes_empty_kv_cache_from_decoder_config() { let config = DecoderConfig::from_file(fixture_path("lightonocr_config").join("config.json")).unwrap(); - let cache = KvCache::empty(&config, 1).unwrap(); + let cache = KVCache::empty(&config, 1).unwrap(); assert!(cache.is_empty()); assert_eq!(cache.layer_count(), 28); @@ -70,7 +70,7 @@ fn initializes_empty_kv_cache_from_decoder_config() { #[test] fn kv_cache_validates_layer_shapes() { - let error = KvCache::new( + let error = KVCache::new( vec![LayerCache::new(vec![0.0; 3], vec![0.0; 4])], 1, 1, @@ -80,7 +80,7 @@ fn kv_cache_validates_layer_shapes() { .expect_err("shape validation should fail"); match error { - Error::InvalidKvCache { reason } => { + Error::InvalidKVCache { reason } => { assert!(reason.contains("key length")); } other => panic!("expected invalid KV cache error, got {other:?}"),