InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
paged_caching_infinilm.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_PAGED_CACHING_INFINILM_H_
2#define INFINI_OPS_BASE_PAGED_CACHING_INFINILM_H_
3
4#include <cassert>
5#include <cstddef>
6
7#include "data_type.h"
8#include "operator.h"
9#include "tensor.h"
10
11namespace infini::ops {
12
15class [[deprecated(
16 "Migrate to an open-source-aligned operator when available.")]]
17PagedCachingInfinilm : public Operator<PagedCachingInfinilm> {
18 public:
20 const Tensor slot_mapping, Tensor k_cache,
21 Tensor v_cache)
22 : dtype_{k.dtype()},
23 num_tokens_{slot_mapping.size(0)},
24 num_kv_heads_{k.size(1)},
25 head_size_{k.size(2)},
26 block_size_{k_cache.size(2)},
27 k_src_stride_{k.stride(0)},
28 v_src_stride_{v.stride(0)},
29 k_cache_block_stride_{k_cache.stride(0)},
30 v_cache_block_stride_{v_cache.stride(0)},
31 k_cache_head_stride_{k_cache.stride(1)},
32 v_cache_head_stride_{v_cache.stride(1)},
33 k_cache_slot_stride_{k_cache.stride(2)},
34 v_cache_slot_stride_{v_cache.stride(2)} {
35 assert(k.ndim() == 3 && v.ndim() == 3 &&
36 "`PagedCachingInfinilm` requires `k` and `v` to be 3D");
37 assert(k_cache.ndim() == 4 && v_cache.ndim() == 4 &&
38 "`PagedCachingInfinilm` requires 4D cache tensors");
39 assert(slot_mapping.ndim() == 1 &&
40 "`PagedCachingInfinilm` requires 1D slot mapping");
41 assert((dtype_ == DataType::kFloat16 || dtype_ == DataType::kBFloat16 ||
42 dtype_ == DataType::kFloat32) &&
43 "`PagedCachingInfinilm` supports float16, bfloat16, and float32");
44 assert(v.dtype() == dtype_ && k_cache.dtype() == dtype_ &&
45 v_cache.dtype() == dtype_);
46 assert(slot_mapping.dtype() == DataType::kInt64 &&
47 "`PagedCachingInfinilm` requires int64 slot mapping");
48 assert(k.shape() == v.shape());
49 assert(k_cache.shape() == v_cache.shape());
50 assert(k_cache.size(1) == num_kv_heads_ && k_cache.size(3) == head_size_);
51 assert(k.stride(2) == 1 && v.stride(2) == 1);
52 assert(k_cache.stride(3) == 1 && v_cache.stride(3) == 1);
53 }
54
55 virtual void operator()(const Tensor k, const Tensor v,
56 const Tensor slot_mapping, Tensor k_cache,
57 Tensor v_cache) const = 0;
58
59 protected:
60 DataType dtype_;
61
62 std::size_t num_tokens_{0};
63
64 std::size_t num_kv_heads_{0};
65
66 std::size_t head_size_{0};
67
68 std::size_t block_size_{0};
69
70 Tensor::Stride k_src_stride_{0};
71
72 Tensor::Stride v_src_stride_{0};
73
74 Tensor::Stride k_cache_block_stride_{0};
75
76 Tensor::Stride v_cache_block_stride_{0};
77
78 Tensor::Stride k_cache_head_stride_{0};
79
80 Tensor::Stride v_cache_head_stride_{0};
81
82 Tensor::Stride k_cache_slot_stride_{0};
83
84 Tensor::Stride v_cache_slot_stride_{0};
85};
86
87} // namespace infini::ops
88
89#endif
Definition generated/include/operator.h:282
Definition paged_caching_infinilm.h:17
PagedCachingInfinilm(const Tensor k, const Tensor v, const Tensor slot_mapping, Tensor k_cache, Tensor v_cache)
Definition paged_caching_infinilm.h:19
virtual void operator()(const Tensor k, const Tensor v, const Tensor slot_mapping, Tensor k_cache, Tensor v_cache) const =0
DataType dtype_
Definition paged_caching_infinilm.h:60
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8