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}