1#ifndef INFINI_OPS_BASE_SCALED_DOT_PRODUCT_ATTENTION_H_
2#define INFINI_OPS_BASE_SCALED_DOT_PRODUCT_ATTENTION_H_
14 const std::optional<Tensor> attn_mask,
15 double dropout_p,
bool is_causal,
16 const std::optional<double> scale,
bool enable_gqa,
22 ?
Tensor::Shape{attn_mask->shape()}
29 ?
Tensor::Strides{attn_mask->strides()}
40 assert(query.ndim() >= 2 && key.ndim() >= 2 && value.ndim() >= 2 &&
41 "`ScaledDotProductAttention` requires rank-2 or higher inputs");
42 assert(query.dtype() == key.dtype() && query.dtype() == value.dtype() &&
43 query.dtype() == out.dtype() &&
44 "`ScaledDotProductAttention` requires matching input/output "
46 assert(query.size(-1) == key.size(-1) && key.size(-2) == value.size(-2) &&
47 "`ScaledDotProductAttention` input dimensions are incompatible");
48 Tensor::Shape expected_out_shape{query.shape()};
49 expected_out_shape.back() = value.size(-1);
50 assert(out.shape() == expected_out_shape &&
51 "`ScaledDotProductAttention` output shape is incorrect");
53 "`ScaledDotProductAttention` requires `dropout_p` in [0, 1]");
59 false, std::nullopt, false, out} {}
63 const std::optional<Tensor> attn_mask,
64 double dropout_p,
bool is_causal,
65 const std::optional<double> scale,
bool enable_gqa,
70 (*this)(query, key, value, std::nullopt, 0.0,
false, std::nullopt,
false,
74 template <
typename TensorLike>
76 const TensorLike& query,
const TensorLike& key,
const TensorLike& value,
77 const std::optional<TensorLike> attn_mask = std::nullopt,
78 double dropout_p = 0.0,
bool is_causal =
false,
79 const std::optional<double> scale = std::nullopt,
80 bool enable_gqa =
false) {
88 typename TensorLike::Shape out_shape{query.shape()};
89 out_shape.back() = value.size(-1);
90 return TensorLike::Empty(out_shape, query.dtype(), query.device());
Definition generated/include/operator.h:282
Definition scaled_dot_product_attention.h:10
ScaledDotProductAttention(const Tensor query, const Tensor key, const Tensor value, Tensor out)
Definition scaled_dot_product_attention.h:56
virtual void operator()(const Tensor query, const Tensor key, const Tensor value, const std::optional< Tensor > attn_mask, double dropout_p, bool is_causal, const std::optional< double > scale, bool enable_gqa, Tensor out) const =0
std::optional< double > scale_
Definition scaled_dot_product_attention.h:122
ScaledDotProductAttention(const Tensor query, const Tensor key, const Tensor value, const std::optional< Tensor > attn_mask, double dropout_p, bool is_causal, const std::optional< double > scale, bool enable_gqa, Tensor out)
Definition scaled_dot_product_attention.h:12
static auto MakeReturnValue(const TensorLike &query, const TensorLike &key, const TensorLike &value, const std::optional< TensorLike > attn_mask=std::nullopt, double dropout_p=0.0, bool is_causal=false, const std::optional< double > scale=std::nullopt, bool enable_gqa=false)
Definition scaled_dot_product_attention.h:75
Tensor::Strides key_strides_
Definition scaled_dot_product_attention.h:106
Tensor::Strides query_strides_
Definition scaled_dot_product_attention.h:104
Tensor::Shape query_shape_
Definition scaled_dot_product_attention.h:94
Tensor::Shape attn_mask_shape_
Definition scaled_dot_product_attention.h:100
bool is_causal_
Definition scaled_dot_product_attention.h:120
Tensor::Shape key_shape_
Definition scaled_dot_product_attention.h:96
Tensor::Shape value_shape_
Definition scaled_dot_product_attention.h:98
DataType query_type_
Definition scaled_dot_product_attention.h:114
Tensor::Strides attn_mask_strides_
Definition scaled_dot_product_attention.h:110
int device_index_
Definition scaled_dot_product_attention.h:126
bool enable_gqa_
Definition scaled_dot_product_attention.h:124
void operator()(const Tensor query, const Tensor key, const Tensor value, Tensor out) const
Definition scaled_dot_product_attention.h:68
Tensor::Strides value_strides_
Definition scaled_dot_product_attention.h:108
Tensor::Strides out_strides_
Definition scaled_dot_product_attention.h:112
DataType attn_mask_type_
Definition scaled_dot_product_attention.h:116
double dropout_p_
Definition scaled_dot_product_attention.h:118
Tensor::Shape out_shape_
Definition scaled_dot_product_attention.h:102
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8