1#ifndef INFINI_OPS_BASE_RESHAPE_AND_CACHE_FLASH_H_
2#define INFINI_OPS_BASE_RESHAPE_AND_CACHE_FLASH_H_
19 const Tensor v_scale,
const std::string kv_cache_dtype,
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 "
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`");
49 "`ReshapeAndCacheFlash` requires enough `key` and `value` rows");
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");
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 "
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");
85 const std::string kv_cache_dtype,
Tensor key_cache,
86 Tensor value_cache)
const = 0;
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