Contiguous KV decode on CPU#
This example builds an immutable Attention model whose graph declares K/V feedback, then runs twenty single-token decode steps. Zero Q/K tensors make attention uniform, so every result is checked against the mean of the values seen so far. The example compares one KV head, eligible for contiguous tail reuse, with two KV heads, which require dense concatenation.
The initial capacity is four tokens, so tokens 5, 9 and 17 exercise growth. CSV output reports latency, kernel allocation requests and capacity bytes, prefix/append copy bytes, buffer reuse counts, and I/O-arena live/peak bytes. Each row also verifies that the state and returned K/V outputs share pointers: state forwarding does not duplicate payloads. The per-token kernel counters exclude feed construction, Attention arithmetic, and score/output workspace.
For one KV head of width four, nongrowing reuse steps report zero allocations,
zero prefix-copy bytes and 32 append-copy bytes. The two-head fallback reports
two allocations and 64 * (token - 1) prefix-copy bytes per step.
Build against an installed native onnx-light tree:
cmake -S examples/contiguous_kv_decode -B build-contiguous-kv-decode \
-DCMAKE_PREFIX_PATH=/path/to/onnx-light-install
cmake --build build-contiguous-kv-decode
./build-contiguous-kv-decode/contiguous_kv_decode
Release each returned output before the next call to allow exclusive reuse. Holding an output or state view is supported, but forces a fresh allocation when that buffer would otherwise be reused. See Persistent input/output feedback for the ownership, valid-length and cancellation contracts.
1// Copyright (c) ONNX Project Contributors
2//
3// SPDX-License-Identifier: Apache-2.0
4
5#include "onnx_core/compute/raw_buffer_allocator.h"
6#include "onnx_core/runtime/persistent_value_state.h"
7#include "onnx_extensions/kernels/kernel_dispatch_table.h"
8#include <chrono>
9#include <cmath>
10#include <iostream>
11
12using namespace ONNX_LIGHT_NAMESPACE;
13using namespace ONNX_LIGHT_NAMESPACE::core::runtime;
14
15namespace {
16
17constexpr int64_t kWidth = 4;
18
19ModelProto DecodeModel(int64_t heads) {
20 ModelProto model;
21 model.set_ir_version(10);
22 model.add_opset_import()->set_version(24);
23 auto *graph = model.mutable_graph();
24 graph->set_name("contiguous_kv_decode");
25 auto declare = [heads](ValueInfoProto *value, const char *name, bool variable_length) {
26 value->set_name(name);
27 auto *tensor = value->mutable_type()->mutable_tensor_type();
28 tensor->set_elem_type(TensorProto::FLOAT);
29 auto *shape = tensor->mutable_shape();
30 shape->add_dim()->set_dim_value(1);
31 shape->add_dim()->set_dim_value(heads);
32 if (variable_length)
33 shape->add_dim();
34 else
35 shape->add_dim()->set_dim_value(1);
36 shape->add_dim()->set_dim_value(kWidth);
37 };
38 for (const char *name : {"q", "k", "v"})
39 declare(graph->add_input(), name, false);
40 for (const char *name : {"past_key", "past_value"})
41 declare(graph->add_input(), name, true);
42 declare(graph->add_output(), "y", false);
43 for (const char *name : {"present_key", "present_value"})
44 declare(graph->add_output(), name, true);
45 auto *node = graph->add_node();
46 node->set_op_type("Attention");
47 for (const char *name : {"q", "k", "v", "", "past_key", "past_value"})
48 node->add_input(name);
49 for (const char *name : {"y", "present_key", "present_value"})
50 node->add_output(name);
51 for (const auto &[input, output] :
52 {std::pair{"past_key", "present_key"}, std::pair{"past_value", "present_value"}}) {
53 auto *binding = graph->add_persistent_bindings();
54 binding->set_input_name(input);
55 binding->set_output_name(output);
56 }
57 return model;
58}
59
60RuntimeValue Filled(int64_t heads, int64_t length, float value) {
61 return RuntimeValue(
62 Tensor::FromFloat("", {1, heads, length, kWidth},
63 std::vector<float>(static_cast<size_t>(heads * length * kWidth), value)));
64}
65
66} // namespace
67
68int main() {
69 onnx_kernels::RegisterKernelFunctions();
70 std::cout << "kv_heads,token,elapsed_ns,kv_allocations,kv_allocated_bytes,"
71 "prefix_copied_bytes,append_copied_bytes,reused_buffers,"
72 "io_live_bytes,io_peak_bytes,state_alias_verified\n";
73 for (int64_t heads : {1, 2}) {
74 const ModelProto model = DecodeModel(heads);
75 SimpleRawBufferAllocator execution(64);
76 auto io = IOArena::Create(32);
77 RuntimeContext context(KernelContext(DefaultOpset(24)),
78 RuntimeContextOptions{.allocator = &execution,
79 .io_allocator = io.get(),
80 .events_enabled = true});
81 PersistentValueState state(
82 model, {{"past_key", Filled(heads, 0, 0)}, {"past_value", Filled(heads, 0, 0)}},
83 RuntimeSessionOptions{.persistent_tensor_initial_capacity = 4});
84 for (int step = 0; step < 20; ++step) {
85 context.ClearEvents();
86 RuntimeValueMap feeds{{"q", Filled(heads, 1, 0)},
87 {"k", Filled(heads, 1, 0)},
88 {"v", Filled(heads, 1, static_cast<float>(step + 1))}};
89 const auto start = std::chrono::steady_clock::now();
90 const auto output = state.Run(context, feeds);
91 const auto elapsed = std::chrono::duration_cast<std::chrono::nanoseconds>(
92 std::chrono::steady_clock::now() - start);
93 const Tensor &y = output.at("y").tensor;
94 for (int64_t i = 0; i < y.element_count(); ++i) {
95 if (std::fabs(y.AsFloat()[i] - static_cast<float>(step + 2) / 2) > 1e-5f) {
96 std::cerr << "Attention result differs from the uniform-attention mean.\n";
97 return 1;
98 }
99 }
100 const auto retained = state.Values();
101 if (retained.at("past_key").tensor.bytes() != output.at("present_key").tensor.bytes() ||
102 retained.at("past_value").tensor.bytes() != output.at("present_value").tensor.bytes()) {
103 std::cerr << "State forwarding unexpectedly copied a KV payload.\n";
104 return 1;
105 }
106 uint64_t allocations = 0, allocated_bytes = 0, prefix_copied = 0, append_copied = 0,
107 reused = 0;
108 for (const auto &event : context.events()) {
109 if (event.action != RuntimeEventAction::kPersistentStorage)
110 continue;
111 allocations += event.storage_allocations;
112 allocated_bytes += event.storage_allocated_bytes;
113 prefix_copied += event.storage_prefix_copied_bytes;
114 append_copied += event.storage_append_copied_bytes;
115 reused += event.storage_reuse_count;
116 }
117 std::cout << heads << "," << step + 1 << "," << elapsed.count() << "," << allocations << ","
118 << allocated_bytes << "," << prefix_copied << "," << append_copied << "," << reused
119 << "," << io->TotalAllocatedSize() << "," << io->PeakAllocatedSize() << ",1\n";
120 }
121 }
122}