InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
reshape_and_cache.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_RESHAPE_AND_CACHE_H_
2#define INFINI_OPS_BASE_RESHAPE_AND_CACHE_H_
3
4#include <string>
5
6#include "data_type.h"
7#include "operator.h"
8
9namespace infini::ops {
10
11class ReshapeAndCache : public Operator<ReshapeAndCache> {
12 public:
13 ReshapeAndCache(const Tensor key, const Tensor value, Tensor key_cache,
14 Tensor value_cache, const Tensor slot_mapping,
15 const Tensor k_scale, const Tensor v_scale,
16 const std::string kv_cache_dtype)
17 : key_shape_{key.shape()},
18 value_shape_{value.shape()},
19 key_cache_shape_{key_cache.shape()},
20 value_cache_shape_{value_cache.shape()},
21 slot_mapping_shape_{slot_mapping.shape()},
22 k_scale_shape_{k_scale.shape()},
23 v_scale_shape_{v_scale.shape()},
24 key_strides_{key.strides()},
25 value_strides_{value.strides()},
26 key_cache_strides_{key_cache.strides()},
27 value_cache_strides_{value_cache.strides()},
28 slot_mapping_strides_{slot_mapping.strides()},
29 k_scale_strides_{k_scale.strides()},
30 v_scale_strides_{v_scale.strides()},
31 key_type_{key.dtype()},
32 key_cache_type_{key_cache.dtype()},
33 value_cache_type_{value_cache.dtype()},
34 k_scale_type_{k_scale.dtype()},
35 v_scale_type_{v_scale.dtype()},
36 kv_cache_dtype_{kv_cache_dtype},
37 num_tokens_{slot_mapping.size(0)},
38 num_heads_{key.size(1)},
39 head_size_{key.size(2)},
40 block_size_{key_cache.size(3)},
41 x_{key_cache.size(4)},
42 device_index_{key.device().index()} {
43 assert(key.ndim() == 3 && value.ndim() == 3 &&
44 "`ReshapeAndCache` requires 3D `key` and `value`");
45 assert(key.shape() == value.shape() &&
46 "`ReshapeAndCache` requires `key` and `value` to have the same "
47 "shape");
48 assert(key.dtype() == value.dtype() &&
49 "`ReshapeAndCache` requires `key` and `value` to have the same "
50 "dtype");
51 assert((key_type_ == DataType::kFloat16 ||
52 key_type_ == DataType::kBFloat16 ||
53 key_type_ == DataType::kFloat32) &&
54 "`ReshapeAndCache` supports float16, bfloat16, and float32 inputs");
55 assert(key_cache.ndim() == 5 && value_cache.ndim() == 4 &&
56 "`ReshapeAndCache` requires 5D `key_cache` and 4D `value_cache`");
57 assert(slot_mapping.ndim() == 1 &&
58 slot_mapping.dtype() == DataType::kInt64 &&
59 slot_mapping.stride(0) == 1 &&
60 "`ReshapeAndCache` requires contiguous int64 `slot_mapping`");
61 assert(num_tokens_ <= key.size(0) &&
62 "`ReshapeAndCache` requires enough key/value rows for all slots");
63 assert(x_ > 0 && head_size_ % x_ == 0 &&
64 "`ReshapeAndCache` requires `head_size` divisible by cache vector "
65 "width");
66 assert(key_cache.size(0) == value_cache.size(0) &&
67 key_cache.size(1) == num_heads_ &&
68 value_cache.size(1) == num_heads_ &&
69 key_cache.size(2) == head_size_ / x_ &&
70 value_cache.size(2) == head_size_ &&
71 value_cache.size(3) == block_size_ &&
72 "`ReshapeAndCache` cache shapes do not match key/value geometry");
73 assert(key.stride(2) == 1 && key.stride(1) == head_size_ &&
74 value.stride(2) == 1 && value.stride(1) == head_size_ &&
75 "`ReshapeAndCache` requires contiguous head dimensions");
76 assert(key_cache.stride(4) == 1 && key_cache.stride(3) == x_ &&
77 key_cache.stride(2) == block_size_ * x_ &&
78 key_cache.stride(1) == head_size_ * block_size_ &&
79 key_cache.stride(0) == num_heads_ * head_size_ * block_size_ &&
80 "`ReshapeAndCache` requires contiguous vLLM key-cache layout");
81 assert(value_cache.stride(3) == 1 && value_cache.stride(2) == block_size_ &&
82 value_cache.stride(1) == head_size_ * block_size_ &&
83 value_cache.stride(0) == num_heads_ * head_size_ * block_size_ &&
84 "`ReshapeAndCache` requires contiguous vLLM value-cache layout");
85 assert((kv_cache_dtype_ == "auto" || kv_cache_dtype_ == "fp8" ||
86 kv_cache_dtype_ == "fp8_e4m3" || kv_cache_dtype_ == "fp8_e5m2") &&
87 "`ReshapeAndCache` received unsupported `kv_cache_dtype`");
88
89 const bool quantized = kv_cache_dtype_ != "auto";
90 assert((quantized ? key_cache_type_ == DataType::kUInt8
93 "`ReshapeAndCache` cache storage dtype does not match "
94 "`kv_cache_dtype`");
95 assert(k_scale.numel() == 1 && v_scale.numel() == 1 &&
96 k_scale_type_ == DataType::kFloat32 &&
97 v_scale_type_ == DataType::kFloat32 &&
98 "`ReshapeAndCache` requires scalar float32 scales");
99 }
100
101 virtual void operator()(const Tensor key, const Tensor value,
102 Tensor key_cache, Tensor value_cache,
103 const Tensor slot_mapping, const Tensor k_scale,
104 const Tensor v_scale,
105 const std::string kv_cache_dtype) const = 0;
106
107 protected:
108 Tensor::Shape key_shape_;
109
110 Tensor::Shape value_shape_;
111
112 Tensor::Shape key_cache_shape_;
113
114 Tensor::Shape value_cache_shape_;
115
116 Tensor::Shape slot_mapping_shape_;
117
118 Tensor::Shape k_scale_shape_;
119
120 Tensor::Shape v_scale_shape_;
121
122 Tensor::Strides key_strides_;
123
124 Tensor::Strides value_strides_;
125
126 Tensor::Strides key_cache_strides_;
127
128 Tensor::Strides value_cache_strides_;
129
130 Tensor::Strides slot_mapping_strides_;
131
132 Tensor::Strides k_scale_strides_;
133
134 Tensor::Strides v_scale_strides_;
135
136 DataType key_type_;
137
139
141
143
145
146 std::string kv_cache_dtype_;
147
148 Tensor::Size num_tokens_{0};
149
150 Tensor::Size num_heads_{0};
151
152 Tensor::Size head_size_{0};
153
154 Tensor::Size block_size_{0};
155
156 Tensor::Size x_{0};
157
159};
160
161} // namespace infini::ops
162
163#endif
Definition generated/include/operator.h:282
Definition reshape_and_cache.h:11
Tensor::Shape value_shape_
Definition reshape_and_cache.h:110
Tensor::Size x_
Definition reshape_and_cache.h:156
Tensor::Strides key_cache_strides_
Definition reshape_and_cache.h:126
Tensor::Size head_size_
Definition reshape_and_cache.h:152
Tensor::Shape key_cache_shape_
Definition reshape_and_cache.h:112
DataType value_cache_type_
Definition reshape_and_cache.h:140
DataType k_scale_type_
Definition reshape_and_cache.h:142
Tensor::Shape key_shape_
Definition reshape_and_cache.h:108
virtual void operator()(const Tensor key, const Tensor value, Tensor key_cache, Tensor value_cache, const Tensor slot_mapping, const Tensor k_scale, const Tensor v_scale, const std::string kv_cache_dtype) const =0
std::string kv_cache_dtype_
Definition reshape_and_cache.h:146
Tensor::Shape v_scale_shape_
Definition reshape_and_cache.h:120
Tensor::Size num_heads_
Definition reshape_and_cache.h:150
int device_index_
Definition reshape_and_cache.h:158
DataType v_scale_type_
Definition reshape_and_cache.h:144
Tensor::Strides k_scale_strides_
Definition reshape_and_cache.h:132
Tensor::Shape slot_mapping_shape_
Definition reshape_and_cache.h:116
Tensor::Strides slot_mapping_strides_
Definition reshape_and_cache.h:130
Tensor::Strides value_strides_
Definition reshape_and_cache.h:124
Tensor::Strides value_cache_strides_
Definition reshape_and_cache.h:128
Tensor::Size block_size_
Definition reshape_and_cache.h:154
Tensor::Shape k_scale_shape_
Definition reshape_and_cache.h:118
Tensor::Strides key_strides_
Definition reshape_and_cache.h:122
DataType key_type_
Definition reshape_and_cache.h:136
Tensor::Shape value_cache_shape_
Definition reshape_and_cache.h:114
ReshapeAndCache(const Tensor key, const Tensor value, Tensor key_cache, Tensor value_cache, const Tensor slot_mapping, const Tensor k_scale, const Tensor v_scale, const std::string kv_cache_dtype)
Definition reshape_and_cache.h:13
Tensor::Strides v_scale_strides_
Definition reshape_and_cache.h:134
Tensor::Size num_tokens_
Definition reshape_and_cache.h:148
DataType key_cache_type_
Definition reshape_and_cache.h:138
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8