InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
scaled_dot_product_attention.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_SCALED_DOT_PRODUCT_ATTENTION_H_
2#define INFINI_OPS_BASE_SCALED_DOT_PRODUCT_ATTENTION_H_
3
4#include <optional>
5
6#include "operator.h"
7
8namespace infini::ops {
9
10class ScaledDotProductAttention : public Operator<ScaledDotProductAttention> {
11 public:
12 ScaledDotProductAttention(const Tensor query, const Tensor key,
13 const Tensor value,
14 const std::optional<Tensor> attn_mask,
15 double dropout_p, bool is_causal,
16 const std::optional<double> scale, bool enable_gqa,
17 Tensor out)
18 : query_shape_{query.shape()},
19 key_shape_{key.shape()},
20 value_shape_{value.shape()},
21 attn_mask_shape_{attn_mask.has_value()
22 ? Tensor::Shape{attn_mask->shape()}
23 : Tensor::Shape{}},
24 out_shape_{out.shape()},
25 query_strides_{query.strides()},
26 key_strides_{key.strides()},
27 value_strides_{value.strides()},
28 attn_mask_strides_{attn_mask.has_value()
29 ? Tensor::Strides{attn_mask->strides()}
30 : Tensor::Strides{}},
31 out_strides_{out.strides()},
32 query_type_{query.dtype()},
33 attn_mask_type_{attn_mask.has_value() ? attn_mask->dtype()
34 : query.dtype()},
35 dropout_p_{dropout_p},
36 is_causal_{is_causal},
37 scale_{scale},
38 enable_gqa_{enable_gqa},
39 device_index_{query.device().index()} {
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 "
45 "dtypes");
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");
52 assert(dropout_p_ >= 0.0 && dropout_p_ <= 1.0 &&
53 "`ScaledDotProductAttention` requires `dropout_p` in [0, 1]");
54 }
55
56 ScaledDotProductAttention(const Tensor query, const Tensor key,
57 const Tensor value, Tensor out)
58 : ScaledDotProductAttention{query, key, value, std::nullopt, 0.0,
59 false, std::nullopt, false, out} {}
60
61 virtual void operator()(const Tensor query, const Tensor key,
62 const Tensor value,
63 const std::optional<Tensor> attn_mask,
64 double dropout_p, bool is_causal,
65 const std::optional<double> scale, bool enable_gqa,
66 Tensor out) const = 0;
67
68 void operator()(const Tensor query, const Tensor key, const Tensor value,
69 Tensor out) const {
70 (*this)(query, key, value, std::nullopt, 0.0, false, std::nullopt, false,
71 out);
72 }
73
74 template <typename TensorLike>
75 static auto MakeReturnValue(
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) {
81 (void)key;
82 (void)attn_mask;
83 (void)dropout_p;
84 (void)is_causal;
85 (void)scale;
86 (void)enable_gqa;
87
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());
91 }
92
93 protected:
94 Tensor::Shape query_shape_;
95
96 Tensor::Shape key_shape_;
97
98 Tensor::Shape value_shape_;
99
100 Tensor::Shape attn_mask_shape_;
101
102 Tensor::Shape out_shape_;
103
104 Tensor::Strides query_strides_;
105
106 Tensor::Strides key_strides_;
107
108 Tensor::Strides value_strides_;
109
110 Tensor::Strides attn_mask_strides_;
111
112 Tensor::Strides out_strides_;
113
114 DataType query_type_;
115
117
118 double dropout_p_{0.0};
119
120 bool is_causal_{false};
121
122 std::optional<double> scale_;
123
124 bool enable_gqa_{false};
125
127};
128
129} // namespace infini::ops
130
131#endif
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