43 const Tensor cu_seqlens_k,
const std::optional<Tensor> alibi_slopes,
44 const std::optional<Tensor> block_table,
const int64_t max_seqlen_q,
45 const int64_t max_seqlen_k,
const double dropout_p,
46 const std::optional<double> softmax_scale,
const bool causal,
47 const std::vector<int64_t> window_size,
const double softcap,
48 const bool deterministic,
const bool return_attn_probs,
Tensor out,
49 std::optional<Tensor> softmax_lse, std::optional<Tensor> s_dmask)
50 : q_shape_{q.shape()},
53 cu_seqlens_q_shape_{cu_seqlens_q.shape()},
54 cu_seqlens_k_shape_{cu_seqlens_k.shape()},
55 alibi_slopes_shape_{alibi_slopes.has_value()
56 ?
Tensor::Shape{alibi_slopes->shape()}
58 block_table_shape_{block_table.has_value()
59 ?
Tensor::Shape{block_table->shape()}
61 out_shape_{out.shape()},
62 softmax_lse_shape_{softmax_lse.has_value()
63 ?
Tensor::Shape{softmax_lse->shape()}
65 s_dmask_shape_{s_dmask.has_value() ?
Tensor::Shape{s_dmask->shape()}
67 q_strides_{q.strides()},
68 k_strides_{k.strides()},
69 v_strides_{v.strides()},
70 cu_seqlens_q_strides_{cu_seqlens_q.strides()},
71 cu_seqlens_k_strides_{cu_seqlens_k.strides()},
72 alibi_slopes_strides_{alibi_slopes.has_value()
73 ?
Tensor::Strides{alibi_slopes->strides()}
75 block_table_strides_{block_table.has_value()
76 ?
Tensor::Strides{block_table->strides()}
78 out_strides_{out.strides()},
79 softmax_lse_strides_{softmax_lse.has_value()
80 ?
Tensor::Strides{softmax_lse->strides()}
82 s_dmask_strides_{s_dmask.has_value()
83 ?
Tensor::Strides{s_dmask->strides()}
88 cu_seqlens_q_dtype_{cu_seqlens_q.dtype()},
89 cu_seqlens_k_dtype_{cu_seqlens_k.dtype()},
90 alibi_slopes_dtype_{alibi_slopes.has_value() ? alibi_slopes->dtype()
91 : DataType::kFloat32},
92 block_table_dtype_{block_table.has_value() ? block_table->dtype()
94 out_dtype_{out.dtype()},
95 softmax_lse_dtype_{softmax_lse.has_value() ? softmax_lse->dtype()
96 : DataType::kFloat32},
97 s_dmask_dtype_{s_dmask.has_value() ? s_dmask->dtype() : q.dtype()},
98 has_auxiliary_outputs_{softmax_lse.has_value() && s_dmask.has_value()},
99 device_index_{q.device().index()} {
100 assert(q.ndim() == 3 &&
101 ((!block_table.has_value() && k.ndim() == 3 && v.ndim() == 3) ||
102 (block_table.has_value() && k.ndim() == 4 && v.ndim() == 4)) &&
103 "`FlashAttnVarlenFunc` requires packed 3D Q and either packed 3D "
104 "or paged 4D K and V tensors");
105 assert(k.shape() == v.shape() &&
106 "`FlashAttnVarlenFunc` requires K and V to have the same shape");
107 assert(q.size(1) > 0 && k.size(-2) > 0 && q.size(2) == k.size(-1) &&
108 q.size(1) % k.size(-2) == 0 &&
109 "`FlashAttnVarlenFunc` requires compatible Q and KV heads");
110 assert(q.size(2) > 0 && q.size(2) <= 256 && q.size(2) % 8 == 0 &&
111 "`FlashAttnVarlenFunc` requires a head dimension divisible by 8 "
112 "and no greater than 256");
113 assert(out.shape() == q.shape() &&
114 "`FlashAttnVarlenFunc` output must have the same shape as Q");
115 assert(softmax_lse.has_value() == s_dmask.has_value() &&
116 "`FlashAttnVarlenFunc` auxiliary outputs must be provided "
118 assert(return_attn_probs == has_auxiliary_outputs_ &&
119 "`FlashAttnVarlenFunc` requires auxiliary outputs exactly when "
120 "`return_attn_probs` is true");
121 if (has_auxiliary_outputs_) {
122 assert((softmax_lse->shape() == Tensor::Shape{q.size(1), q.size(0)} &&
123 softmax_lse_dtype_ == DataType::kFloat32 &&
124 "`FlashAttnVarlenFunc` softmax LSE output must have shape "
125 "(num_heads, total_q) and dtype float32"));
126 assert(s_dmask->shape() == Tensor::Shape{0} &&
127 s_dmask_dtype_ == q_dtype_ &&
128 "`FlashAttnVarlenFunc` inference attention mask output must be "
129 "empty and match the Q dtype");
132 (q_dtype_ == DataType::kFloat16 || q_dtype_ == DataType::kBFloat16) &&
133 q_dtype_ == k_dtype_ && q_dtype_ == v_dtype_ &&
134 q_dtype_ == out_dtype_ &&
135 "`FlashAttnVarlenFunc` requires matching float16 or bfloat16 Q, "
136 "K, V, and output dtypes");
137 assert(q.stride(-1) == 1 && k.stride(-1) == 1 && v.stride(-1) == 1 &&
138 out.stride(-1) == 1 &&
139 "`FlashAttnVarlenFunc` requires contiguous head dimensions");
140 assert(cu_seqlens_q.ndim() == 1 && cu_seqlens_k.ndim() == 1 &&
141 cu_seqlens_q.shape() == cu_seqlens_k.shape() &&
142 cu_seqlens_q.numel() >= 2 &&
143 "`FlashAttnVarlenFunc` cumulative sequence tensors must be "
144 "matching non-empty vectors");
145 assert(cu_seqlens_q_dtype_ == DataType::kInt32 &&
146 cu_seqlens_k_dtype_ == DataType::kInt32 &&
147 cu_seqlens_q.IsContiguous() && cu_seqlens_k.IsContiguous() &&
148 "`FlashAttnVarlenFunc` cumulative sequence tensors must be "
149 "contiguous int32 tensors");
150 assert(max_seqlen_q > 0 && max_seqlen_k > 0 &&
151 "`FlashAttnVarlenFunc` maximum sequence lengths must be positive");
152 assert(window_size.size() == 2 && window_size[0] >= -1 &&
153 window_size[1] >= -1 &&
154 "`FlashAttnVarlenFunc` `window_size` must contain two values >= -1");
156 assert(dropout_p == 0.0 &&
157 "`FlashAttnVarlenFunc` initially supports inference only");
158 assert(softcap == 0.0 &&
159 "`FlashAttnVarlenFunc` does not yet support softcap");
160 assert(!deterministic &&
161 "`FlashAttnVarlenFunc` does not yet support deterministic mode");
162 if (block_table.has_value()) {
163 assert(block_table->ndim() == 2 &&
164 block_table->size(0) + 1 == cu_seqlens_q.size(0) &&
165 block_table_dtype_ == DataType::kInt32 &&
166 block_table->IsContiguous() && k.size(1) % 256 == 0 &&
167 "`FlashAttnVarlenFunc` requires a contiguous int32 block table "
168 "and page size divisible by 256");
170 if (alibi_slopes.has_value()) {
172 (alibi_slopes->ndim() == 1 || alibi_slopes->ndim() == 2) &&
173 alibi_slopes_dtype_ == DataType::kFloat32 &&
174 alibi_slopes->IsContiguous() &&
175 ((alibi_slopes->ndim() == 1 && alibi_slopes->size(0) == q.size(1)) ||
176 (alibi_slopes->ndim() == 2 &&
177 alibi_slopes->size(0) + 1 == cu_seqlens_q.size(0) &&
178 alibi_slopes->size(1) == q.size(1))) &&
179 "`FlashAttnVarlenFunc` received incompatible ALiBi slopes");
182 const auto same_device_as_q = [&](
const Tensor tensor) {
183 return tensor.device().type() == q.device().type() &&
184 tensor.device().index() == q.device().index();
186 assert(same_device_as_q(k) && same_device_as_q(v) &&
187 same_device_as_q(cu_seqlens_q) && same_device_as_q(cu_seqlens_k) &&
188 same_device_as_q(out) &&
189 (!alibi_slopes.has_value() || same_device_as_q(*alibi_slopes)) &&
190 (!block_table.has_value() || same_device_as_q(*block_table)) &&
191 (!softmax_lse.has_value() || same_device_as_q(*softmax_lse)) &&
192 (!s_dmask.has_value() || same_device_as_q(*s_dmask)) &&
193 "`FlashAttnVarlenFunc` tensors must be on the same device");