Continuum C++ API
Unified runtime for token + tensor execution
Loading...
Searching...
No Matches
memo_table.hpp
Go to the documentation of this file.
1#pragma once
2
5
6#include <cstdint>
7#include <functional>
8#include <mutex>
9#include <optional>
10#include <string>
11#include <unordered_map>
12#include <vector>
13
14namespace continuum::runtime {
15
16struct MemoKey {
17 std::string node_kind_str;
18 std::string payload_hash;
19 std::vector<std::uint8_t> inputs_hash;
20 std::string cache_namespace;
21
22 bool operator==(const MemoKey& o) const {
23 return node_kind_str == o.node_kind_str &&
27 }
28};
29
30struct MemoEntry {
31 std::vector<std::uint8_t> output_bytes;
32 std::uint64_t version = 0;
33 std::int64_t access_count = 0;
34 std::int64_t last_access_ns = 0;
35};
36
38 std::size_t operator()(const MemoKey& k) const {
39 std::size_t h = std::hash<std::string>{}(k.node_kind_str);
40 h ^= std::hash<std::string>{}(k.payload_hash) + 0x9e3779b97f4a7c15ULL + (h << 6) + (h >> 2);
41 for (auto b : k.inputs_hash) {
42 h ^= static_cast<std::size_t>(b) + 0x9e3779b97f4a7c15ULL + (h << 6) + (h >> 2);
43 }
44 h ^= std::hash<std::string>{}(k.cache_namespace) + 0x9e3779b97f4a7c15ULL + (h << 6) + (h >> 2);
45 return h;
46 }
47};
48
55class MemoTable {
56 public:
57 explicit MemoTable(std::size_t max_entries = 4096,
58 std::size_t version = 0);
59
60 std::optional<MemoEntry> lookup(const MemoKey& key) const;
61 void insert(MemoKey key, MemoEntry entry);
62
63 void invalidate_version(std::size_t version);
64 void invalidate_node(const std::string& node_kind_str);
65 void clear();
66
67 std::size_t size() const;
69 std::size_t max_entries() const { return max_entries_; }
71 std::size_t estimated_bytes() const;
72 std::size_t version() const;
73 void set_version(std::size_t v);
74
75 MemoKey make_key(const ir::Node& node, const std::vector<continuum::Value>& inputs,
76 const std::string& cache_namespace = {}) const;
77
78 static std::vector<std::uint8_t> serialize_value(const continuum::Value& v);
79 static std::optional<continuum::Value> deserialize_value(const std::vector<std::uint8_t>& bytes);
80
81 private:
82 mutable std::mutex mu_;
83 std::unordered_map<MemoKey, MemoEntry, MemoKeyHash> table_;
84 std::size_t max_entries_;
85 std::size_t version_;
86 mutable std::uint64_t clock_ = 0;
87};
88
89} // namespace continuum::runtime
Definition memo_table.hpp:55
MemoKey make_key(const ir::Node &node, const std::vector< continuum::Value > &inputs, const std::string &cache_namespace={}) const
void insert(MemoKey key, MemoEntry entry)
std::size_t estimated_bytes() const
Approximate resident bytes: keys, cached outputs, and per-entry overhead.
void set_version(std::size_t v)
void invalidate_node(const std::string &node_kind_str)
MemoTable(std::size_t max_entries=4096, std::size_t version=0)
static std::vector< std::uint8_t > serialize_value(const continuum::Value &v)
static std::optional< continuum::Value > deserialize_value(const std::vector< std::uint8_t > &bytes)
std::size_t version() const
std::size_t max_entries() const
Capacity in entries passed at construction.
Definition memo_table.hpp:69
void invalidate_version(std::size_t version)
std::optional< MemoEntry > lookup(const MemoKey &key) const
std::size_t size() const
Definition checkpoint.hpp:12
std::variant< TensorValue, MlxTensorValue, TokensValue, SchemaValue, std::string, double, int64_t > Value
Definition value.hpp:30
Definition node.hpp:51
Definition memo_table.hpp:30
std::uint64_t version
Definition memo_table.hpp:32
std::int64_t access_count
Definition memo_table.hpp:33
std::vector< std::uint8_t > output_bytes
Definition memo_table.hpp:31
std::int64_t last_access_ns
Definition memo_table.hpp:34
Definition memo_table.hpp:37
std::size_t operator()(const MemoKey &k) const
Definition memo_table.hpp:38
Definition memo_table.hpp:16
std::string payload_hash
Definition memo_table.hpp:18
std::string cache_namespace
Definition memo_table.hpp:20
std::vector< std::uint8_t > inputs_hash
Definition memo_table.hpp:19
std::string node_kind_str
Definition memo_table.hpp:17
bool operator==(const MemoKey &o) const
Definition memo_table.hpp:22