InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
reshape_and_cache_flash.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_RESHAPE_AND_CACHE_FLASH_H_
2#define INFINI_OPS_BASE_RESHAPE_AND_CACHE_FLASH_H_
3
4#include <cassert>
5#include <cstddef>
6#include <cstdint>
7#include <string>
8
9#include "data_type.h"
10#include "operator.h"
11
12namespace infini::ops {
13
14// Aligned with vLLM `_custom_ops.reshape_and_cache_flash`.
15class ReshapeAndCacheFlash : public Operator<ReshapeAndCacheFlash> {
16 public:
17 ReshapeAndCacheFlash(const Tensor key, const Tensor value,
18 const Tensor slot_mapping, const Tensor k_scale,
19 const Tensor v_scale, const std::string kv_cache_dtype,
20 Tensor key_cache, Tensor value_cache)
21 : dtype_{key.dtype()},
22 num_tokens_{slot_mapping.size(0)},
23 num_heads_{key.size(1)},
24 head_size_{key.size(2)},
25 block_size_{key_cache.size(1)},
26 key_token_stride_{key.stride(0)},
27 value_token_stride_{value.stride(0)},
28 key_head_stride_{key.stride(1)},
29 value_head_stride_{value.stride(1)},
30 key_cache_block_stride_{key_cache.stride(0)},
31 value_cache_block_stride_{value_cache.stride(0)},
32 key_cache_page_stride_{key_cache.stride(1)},
33 value_cache_page_stride_{value_cache.stride(1)},
34 key_cache_head_stride_{key_cache.stride(2)},
35 value_cache_head_stride_{value_cache.stride(2)} {
36 assert(key.ndim() == 3 && value.ndim() == 3 &&
37 "`ReshapeAndCacheFlash` requires 3D `key` and `value`");
38 assert(key.shape() == value.shape() && key.dtype() == value.dtype() &&
39 "`ReshapeAndCacheFlash` requires `key` and `value` to have matching "
40 "shapes and dtypes");
41 assert((dtype_ == DataType::kFloat16 || dtype_ == DataType::kBFloat16 ||
42 dtype_ == DataType::kFloat32) &&
43 "`ReshapeAndCacheFlash` supports float16, bfloat16, and float32");
44 assert(slot_mapping.ndim() == 1 &&
45 slot_mapping.dtype() == DataType::kInt64 &&
46 slot_mapping.stride(0) == 1 &&
47 "`ReshapeAndCacheFlash` requires contiguous int64 `slot_mapping`");
48 assert(num_tokens_ <= key.size(0) &&
49 "`ReshapeAndCacheFlash` requires enough `key` and `value` rows");
50 assert(num_heads_ > 0 && head_size_ > 0 &&
51 "`ReshapeAndCacheFlash` requires non-empty head dimensions");
52 assert(key_cache.ndim() == 4 && value_cache.ndim() == 4 &&
53 key_cache.shape() == value_cache.shape() &&
54 "`ReshapeAndCacheFlash` requires matching 4D caches");
55 assert(block_size_ > 0 && key_cache.size(2) == num_heads_ &&
56 key_cache.size(3) == head_size_ &&
57 "`ReshapeAndCacheFlash` cache shape must be "
58 "[`num_blocks`, `block_size`, `num_heads`, `head_size`]");
59 assert(key.stride(2) == 1 && value.stride(2) == 1 &&
60 key_cache.stride(3) == 1 && value_cache.stride(3) == 1 &&
61 "`ReshapeAndCacheFlash` requires contiguous head dimensions");
62 assert(key.device() == value.device() &&
63 key.device() == slot_mapping.device() &&
64 key.device() == k_scale.device() &&
65 key.device() == v_scale.device() &&
66 key.device() == key_cache.device() &&
67 key.device() == value_cache.device() &&
68 "`ReshapeAndCacheFlash` tensors must be on the same device");
69 assert(k_scale.shape() == v_scale.shape() &&
70 (k_scale.numel() == 1 || k_scale.numel() == num_heads_) &&
71 k_scale.dtype() == DataType::kFloat32 &&
72 v_scale.dtype() == DataType::kFloat32 &&
73 "`ReshapeAndCacheFlash` scales must be float32 scalar or per-head "
74 "tensors");
75 assert(kv_cache_dtype == "auto" && key_cache.dtype() == dtype_ &&
76 value_cache.dtype() == dtype_ &&
77 "`ReshapeAndCacheFlash` currently supports `auto` cache dtype");
78 assert(!key_cache.HasBroadcastDim() && !value_cache.HasBroadcastDim() &&
79 "`ReshapeAndCacheFlash` caches must not have broadcast dimensions");
80 }
81
82 virtual void operator()(const Tensor key, const Tensor value,
83 const Tensor slot_mapping, const Tensor k_scale,
84 const Tensor v_scale,
85 const std::string kv_cache_dtype, Tensor key_cache,
86 Tensor value_cache) const = 0;
87
88 protected:
89 DataType dtype_;
90
91 std::size_t num_tokens_{0};
92
93 std::size_t num_heads_{0};
94
95 std::size_t head_size_{0};
96
97 std::size_t block_size_{0};
98
99 Tensor::Stride key_token_stride_{0};
100
101 Tensor::Stride value_token_stride_{0};
102
103 Tensor::Stride key_head_stride_{0};
104
105 Tensor::Stride value_head_stride_{0};
106
107 Tensor::Stride key_cache_block_stride_{0};
108
109 Tensor::Stride value_cache_block_stride_{0};
110
111 Tensor::Stride key_cache_page_stride_{0};
112
113 Tensor::Stride value_cache_page_stride_{0};
114
115 Tensor::Stride key_cache_head_stride_{0};
116
117 Tensor::Stride value_cache_head_stride_{0};
118};
119
120} // namespace infini::ops
121
122#endif
Definition generated/include/operator.h:282
Definition reshape_and_cache_flash.h:15
Tensor::Stride key_cache_head_stride_
Definition reshape_and_cache_flash.h:115
ReshapeAndCacheFlash(const Tensor key, const Tensor value, const Tensor slot_mapping, const Tensor k_scale, const Tensor v_scale, const std::string kv_cache_dtype, Tensor key_cache, Tensor value_cache)
Definition reshape_and_cache_flash.h:17
std::size_t head_size_
Definition reshape_and_cache_flash.h:95
std::size_t num_tokens_
Definition reshape_and_cache_flash.h:91
virtual void operator()(const Tensor key, const Tensor value, const Tensor slot_mapping, const Tensor k_scale, const Tensor v_scale, const std::string kv_cache_dtype, Tensor key_cache, Tensor value_cache) const =0
Tensor::Stride value_cache_page_stride_
Definition reshape_and_cache_flash.h:113
Tensor::Stride value_head_stride_
Definition reshape_and_cache_flash.h:105
Tensor::Stride key_cache_block_stride_
Definition reshape_and_cache_flash.h:107
Tensor::Stride key_cache_page_stride_
Definition reshape_and_cache_flash.h:111
std::size_t num_heads_
Definition reshape_and_cache_flash.h:93
Tensor::Stride value_cache_head_stride_
Definition reshape_and_cache_flash.h:117
Tensor::Stride key_token_stride_
Definition reshape_and_cache_flash.h:99
Tensor::Stride value_token_stride_
Definition reshape_and_cache_flash.h:101
DataType dtype_
Definition reshape_and_cache_flash.h:89
Tensor::Stride key_head_stride_
Definition reshape_and_cache_flash.h:103
std::size_t block_size_
Definition reshape_and_cache_flash.h:97
Tensor::Stride value_cache_block_stride_
Definition reshape_and_cache_flash.h:109
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8