InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
kv_caching_infinilm.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_KV_CACHING_INFINILM_H_
2#define INFINI_OPS_BASE_KV_CACHING_INFINILM_H_
3
4#include <cassert>
5
6#include "operator.h"
7
8namespace infini::ops {
9
12class [[deprecated(
13 "Migrate to an open-source-aligned operator when available.")]]
14KvCachingInfinilm : public Operator<KvCachingInfinilm> {
15 public:
17 const Tensor past_kv_lengths, Tensor k_cache,
18 Tensor v_cache)
19 : k_cache_shape_{k_cache.shape()},
20 k_cache_strides_{k_cache.strides()},
21 v_cache_shape_{v_cache.shape()},
22 v_cache_strides_{v_cache.strides()},
23 k_shape_{k.shape()},
24 k_strides_{k.strides()},
25 v_shape_{v.shape()},
26 v_strides_{v.strides()},
27 past_kv_lengths_shape_{past_kv_lengths.shape()},
28 data_type_{k_cache.dtype()},
29 past_kv_lengths_type_{past_kv_lengths.dtype()},
30 batch_size_{k_cache.size(0)},
31 num_kv_heads_{k_cache.size(1)},
32 max_seq_len_{k_cache.size(2)},
33 seq_len_{k.size(2)},
34 hidden_size_{k_cache.size(3)},
35 output_size_{k.numel()},
36 device_index_{k_cache.device().index()} {
37 assert(k_cache.ndim() == 4 && v_cache.ndim() == 4 && k.ndim() == 4 &&
38 v.ndim() == 4 && "`KvCachingInfinilm` tensors must be 4D");
39 assert(k_cache_shape_ == v_cache_shape_ &&
40 "`KvCachingInfinilm` cache shapes must match");
41 assert(k_shape_ == v_shape_ &&
42 "`KvCachingInfinilm` source shapes must match");
43 assert(k.size(0) == batch_size_ && k.size(1) == num_kv_heads_ &&
44 k.size(3) == hidden_size_ &&
45 "`KvCachingInfinilm` source shape must match cache "
46 "batch/head/hidden dims");
47 assert(seq_len_ <= max_seq_len_ &&
48 "`KvCachingInfinilm` source sequence length exceeds cache length");
49 assert(k_cache.dtype() == v_cache.dtype() && k_cache.dtype() == k.dtype() &&
50 k_cache.dtype() == v.dtype() &&
51 "`KvCachingInfinilm` K/V tensors must have the same dtype");
52 assert(
53 (data_type_ == DataType::kFloat16 ||
54 data_type_ == DataType::kBFloat16 ||
55 data_type_ == DataType::kFloat32) &&
56 "`KvCachingInfinilm` K/V dtype must be float16, bfloat16, or float32");
57 assert((past_kv_lengths_type_ == DataType::kInt32 ||
58 past_kv_lengths_type_ == DataType::kInt64) &&
59 "`KvCachingInfinilm` past_kv_lengths dtype must be int32 or int64");
60 assert(past_kv_lengths.ndim() == 1 &&
61 past_kv_lengths.size(0) == batch_size_ &&
62 "`KvCachingInfinilm` past_kv_lengths shape must be (batch_size,)");
63 assert(!k_cache.HasBroadcastDim() && !v_cache.HasBroadcastDim() &&
64 "`KvCachingInfinilm` caches must not have broadcasted dimensions");
65 }
66
67 virtual void operator()(const Tensor k, const Tensor v,
68 const Tensor past_kv_lengths, Tensor k_cache,
69 Tensor v_cache) const = 0;
70
71 protected:
72 Tensor::Shape k_cache_shape_;
73
74 Tensor::Strides k_cache_strides_;
75
76 Tensor::Shape v_cache_shape_;
77
78 Tensor::Strides v_cache_strides_;
79
80 Tensor::Shape k_shape_;
81
82 Tensor::Strides k_strides_;
83
84 Tensor::Shape v_shape_;
85
86 Tensor::Strides v_strides_;
87
89
90 DataType data_type_;
91
93
94 Tensor::Size batch_size_{0};
95
96 Tensor::Size num_kv_heads_{0};
97
98 Tensor::Size max_seq_len_{0};
99
100 Tensor::Size seq_len_{0};
101
102 Tensor::Size hidden_size_{0};
103
104 Tensor::Size output_size_{0};
105
106 int device_index_{0};
107};
108
109} // namespace infini::ops
110
111#endif
Definition kv_caching_infinilm.h:14
Tensor::Strides v_cache_strides_
Definition kv_caching_infinilm.h:78
DataType past_kv_lengths_type_
Definition kv_caching_infinilm.h:92
Tensor::Strides v_strides_
Definition kv_caching_infinilm.h:86
virtual void operator()(const Tensor k, const Tensor v, const Tensor past_kv_lengths, Tensor k_cache, Tensor v_cache) const =0
Tensor::Shape k_shape_
Definition kv_caching_infinilm.h:80
Tensor::Strides k_cache_strides_
Definition kv_caching_infinilm.h:74
Tensor::Shape past_kv_lengths_shape_
Definition kv_caching_infinilm.h:88
Tensor::Shape v_cache_shape_
Definition kv_caching_infinilm.h:76
Tensor::Strides k_strides_
Definition kv_caching_infinilm.h:82
Tensor::Shape v_shape_
Definition kv_caching_infinilm.h:84
KvCachingInfinilm(const Tensor k, const Tensor v, const Tensor past_kv_lengths, Tensor k_cache, Tensor v_cache)
Definition kv_caching_infinilm.h:16
DataType data_type_
Definition kv_caching_infinilm.h:90
Tensor::Shape k_cache_shape_
Definition kv_caching_infinilm.h:72
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8