84 const std::optional<Tensor> k,
85 const std::optional<Tensor> v,
86 const std::optional<Tensor> rotary_cos,
87 const std::optional<Tensor> rotary_sin,
88 const std::optional<Tensor> cache_seqlens,
89 const std::optional<Tensor> cache_batch_idx,
90 const std::optional<Tensor> cache_leftpad,
91 const std::optional<Tensor> block_table,
92 const std::optional<Tensor> alibi_slopes,
93 const std::optional<double> softmax_scale,
95 const std::vector<int64_t> window_size,
96 const double softcap,
const bool rotary_interleaved,
97 const int64_t num_splits,
const bool return_softmax_lse,
98 Tensor out, std::optional<Tensor> softmax_lse)
99 : q_shape_{q.shape()},
100 k_cache_shape_{k_cache.shape()},
101 v_cache_shape_{v_cache.shape()},
102 k_shape_{k.has_value() ?
Tensor::Shape{k->shape()} :
Tensor::Shape{}},
103 v_shape_{v.has_value() ?
Tensor::Shape{v->shape()} :
Tensor::Shape{}},
104 rotary_cos_shape_{rotary_cos.has_value()
105 ?
Tensor::Shape{rotary_cos->shape()}
107 rotary_sin_shape_{rotary_sin.has_value()
108 ?
Tensor::Shape{rotary_sin->shape()}
110 cache_seqlens_shape_{cache_seqlens.has_value()
111 ?
Tensor::Shape{cache_seqlens->shape()}
113 cache_batch_idx_shape_{cache_batch_idx.has_value()
114 ?
Tensor::Shape{cache_batch_idx->shape()}
116 cache_leftpad_shape_{cache_leftpad.has_value()
117 ?
Tensor::Shape{cache_leftpad->shape()}
119 block_table_shape_{block_table.has_value()
120 ?
Tensor::Shape{block_table->shape()}
122 alibi_slopes_shape_{alibi_slopes.has_value()
123 ?
Tensor::Shape{alibi_slopes->shape()}
125 out_shape_{out.shape()},
126 softmax_lse_shape_{softmax_lse.has_value()
127 ?
Tensor::Shape{softmax_lse->shape()}
129 q_strides_{q.strides()},
130 k_cache_strides_{k_cache.strides()},
131 v_cache_strides_{v_cache.strides()},
132 k_strides_{k.has_value() ?
Tensor::Strides{k->strides()}
134 v_strides_{v.has_value() ?
Tensor::Strides{v->strides()}
136 rotary_cos_strides_{rotary_cos.has_value()
137 ?
Tensor::Strides{rotary_cos->strides()}
139 rotary_sin_strides_{rotary_sin.has_value()
140 ?
Tensor::Strides{rotary_sin->strides()}
142 cache_seqlens_strides_{cache_seqlens.has_value()
143 ?
Tensor::Strides{cache_seqlens->strides()}
145 cache_batch_idx_strides_{
146 cache_batch_idx.has_value()
147 ?
Tensor::Strides{cache_batch_idx->strides()}
149 cache_leftpad_strides_{cache_leftpad.has_value()
150 ?
Tensor::Strides{cache_leftpad->strides()}
152 block_table_strides_{block_table.has_value()
153 ?
Tensor::Strides{block_table->strides()}
155 alibi_slopes_strides_{alibi_slopes.has_value()
156 ?
Tensor::Strides{alibi_slopes->strides()}
158 out_strides_{out.strides()},
159 softmax_lse_strides_{softmax_lse.has_value()
160 ?
Tensor::Strides{softmax_lse->strides()}
163 k_cache_dtype_{k_cache.dtype()},
164 v_cache_dtype_{v_cache.dtype()},
165 k_dtype_{k.has_value() ? k->dtype() : q.dtype()},
166 v_dtype_{v.has_value() ? v->dtype() : q.dtype()},
167 rotary_cos_dtype_{rotary_cos.has_value() ? rotary_cos->dtype()
169 rotary_sin_dtype_{rotary_sin.has_value() ? rotary_sin->dtype()
171 cache_seqlens_dtype_{cache_seqlens.has_value() ? cache_seqlens->dtype()
173 cache_batch_idx_dtype_{cache_batch_idx.has_value()
174 ? cache_batch_idx->dtype()
176 cache_leftpad_dtype_{cache_leftpad.has_value() ? cache_leftpad->dtype()
178 block_table_dtype_{block_table.has_value() ? block_table->dtype()
180 alibi_slopes_dtype_{alibi_slopes.has_value() ? alibi_slopes->dtype()
181 : DataType::kFloat32},
182 out_dtype_{out.dtype()},
183 softmax_lse_dtype_{softmax_lse.has_value() ? softmax_lse->dtype()
184 : DataType::kFloat32},
185 has_k_{k.has_value()},
186 has_v_{v.has_value()},
187 has_rotary_cos_{rotary_cos.has_value()},
188 has_rotary_sin_{rotary_sin.has_value()},
189 has_cache_batch_idx_{cache_batch_idx.has_value()},
190 has_block_table_{block_table.has_value()},
191 has_alibi_slopes_{alibi_slopes.has_value()},
192 has_softmax_lse_{softmax_lse.has_value()},
193 batch_size_{q.ndim() > 0 ? q.size(0) : 0},
194 head_size_{q.ndim() == 4 ? q.size(3) : 0},
195 device_index_{q.device().index()} {
196 assert(q.ndim() == 4 && k_cache.ndim() == 4 && v_cache.ndim() == 4 &&
197 "`FlashAttnWithKvcache` requires 4D `q`, `k_cache`, and "
199 assert(k_cache.shape() == v_cache.shape() &&
200 "`FlashAttnWithKvcache` requires matching K/V cache shapes");
202 (q_dtype_ == DataType::kFloat16 || q_dtype_ == DataType::kBFloat16) &&
203 q_dtype_ == k_cache_dtype_ && q_dtype_ == v_cache_dtype_ &&
204 q_dtype_ == out_dtype_ &&
205 "`FlashAttnWithKvcache` requires matching float16 or bfloat16 "
206 "Q, cache, and output dtypes");
207 assert(q.size(0) > 0 && q.size(1) > 0 && q.size(2) > 0 &&
208 k_cache.size(0) > 0 && k_cache.size(1) > 0 && k_cache.size(2) > 0 &&
209 "`FlashAttnWithKvcache` requires non-empty Q and KV cache "
211 assert(q.size(2) % k_cache.size(2) == 0 && q.size(3) == k_cache.size(3) &&
212 "`FlashAttnWithKvcache` requires compatible Q and KV heads");
213 assert(head_size_ > 0 && head_size_ <= 256 &&
214 "`FlashAttnWithKvcache` requires a head dimension no greater than "
216 assert(out.shape() == q.shape() &&
217 "`FlashAttnWithKvcache` output must have the same shape as Q");
218 assert(return_softmax_lse == has_softmax_lse_ &&
219 "`FlashAttnWithKvcache` requires `softmax_lse` exactly when "
220 "`return_softmax_lse` is true");
221 if (has_softmax_lse_) {
222 assert((softmax_lse->shape() ==
223 Tensor::Shape{q.size(0), q.size(2), q.size(1)}) &&
224 softmax_lse_dtype_ == DataType::kFloat32 &&
225 "`FlashAttnWithKvcache` softmax LSE output must have shape "
226 "(batch_size, num_heads, seqlen) and dtype float32");
228 assert(q.stride(-1) == 1 && k_cache.stride(-1) == 1 &&
229 v_cache.stride(-1) == 1 && out.stride(-1) == 1 &&
230 "`FlashAttnWithKvcache` requires contiguous head dimensions");
231 assert(has_k_ == has_v_ &&
232 "`FlashAttnWithKvcache` requires `k` and `v` together");
234 assert(k->ndim() == 4 && k->shape() == v->shape() &&
235 k->size(0) == batch_size_ && k->size(2) == k_cache.size(2) &&
236 k->size(3) == head_size_ &&
237 "`FlashAttnWithKvcache` received incompatible new K/V shapes");
238 assert(k_dtype_ == q_dtype_ && v_dtype_ == q_dtype_ &&
239 "`FlashAttnWithKvcache` requires matching new K/V dtypes");
240 assert(k->stride(-1) == 1 && v->stride(-1) == 1 &&
241 "`FlashAttnWithKvcache` requires contiguous new K/V head "
244 assert(has_rotary_cos_ == has_rotary_sin_ &&
245 "`FlashAttnWithKvcache` requires rotary cosine and sine together");
246 if (has_rotary_cos_) {
247 assert(has_k_ && rotary_cos->ndim() == 2 &&
248 rotary_cos->shape() == rotary_sin->shape() &&
249 rotary_cos_dtype_ == q_dtype_ && rotary_sin_dtype_ == q_dtype_ &&
250 rotary_cos->size(1) > 0 && rotary_cos->size(1) * 2 <= head_size_ &&
251 (rotary_cos->size(1) * 2) % 16 == 0 &&
252 "`FlashAttnWithKvcache` received incompatible rotary tables");
254 ValidateIndexVector(cache_seqlens);
255 ValidateIndexVector(cache_batch_idx);
256 ValidateIndexVector(cache_leftpad);
257 if (has_block_table_) {
258 assert(block_table->ndim() == 2 && block_table->size(0) == batch_size_ &&
259 block_table_dtype_ == DataType::kInt32 &&
260 block_table->IsContiguous() && k_cache.size(1) % 256 == 0 &&
261 "`FlashAttnWithKvcache` requires a contiguous int32 block table "
262 "and page size divisible by 256");
264 assert((has_cache_batch_idx_ || k_cache.size(0) >= batch_size_) &&
265 "`FlashAttnWithKvcache` cache batch is too small");
267 if (has_alibi_slopes_) {
268 assert((alibi_slopes->ndim() == 1 || alibi_slopes->ndim() == 2) &&
269 alibi_slopes_dtype_ == DataType::kFloat32 &&
270 alibi_slopes->IsContiguous() &&
271 "`FlashAttnWithKvcache` requires contiguous float32 ALiBi "
274 ((alibi_slopes->ndim() == 1 && alibi_slopes->size(0) == q.size(2)) ||
275 (alibi_slopes->ndim() == 2 && alibi_slopes->size(0) == batch_size_ &&
276 alibi_slopes->size(1) == q.size(2))) &&
277 "`FlashAttnWithKvcache` received incompatible ALiBi slopes");
279 assert(window_size.size() == 2 && window_size[0] >= -1 &&
280 window_size[1] >= -1 &&
281 "`FlashAttnWithKvcache` `window_size` must contain two values >= "
283 assert(softcap >= 0.0 &&
284 "`FlashAttnWithKvcache` requires non-negative `softcap`");
285 assert(num_splits >= 0 &&
286 "`FlashAttnWithKvcache` requires non-negative `num_splits`");
287 const auto same_device_as_q = [&](
const Tensor tensor) {
288 return tensor.device().type() == q.device().type() &&
289 tensor.device().index() == q.device().index();
292 same_device_as_q(k_cache) && same_device_as_q(v_cache) &&
293 same_device_as_q(out) &&
294 (!softmax_lse.has_value() || same_device_as_q(*softmax_lse)) &&
295 (!k.has_value() || same_device_as_q(*k)) &&
296 (!v.has_value() || same_device_as_q(*v)) &&
297 (!rotary_cos.has_value() || same_device_as_q(*rotary_cos)) &&
298 (!rotary_sin.has_value() || same_device_as_q(*rotary_sin)) &&
299 (!cache_seqlens.has_value() || same_device_as_q(*cache_seqlens)) &&
300 (!cache_batch_idx.has_value() || same_device_as_q(*cache_batch_idx)) &&
301 (!cache_leftpad.has_value() || same_device_as_q(*cache_leftpad)) &&
302 (!block_table.has_value() || same_device_as_q(*block_table)) &&
303 (!alibi_slopes.has_value() || same_device_as_q(*alibi_slopes)) &&
304 "`FlashAttnWithKvcache` tensors must be on the same device");
308 (void)rotary_interleaved;