Skip to content

Commit 8c0c777

Browse files
committed
change std::optional<at::Tensor> position_ids to const std::optional<at::Tensor>& position_ids
1 parent 4459b41 commit 8c0c777

2 files changed

Lines changed: 26 additions & 25 deletions

File tree

include/flag_gems/operators.h

Lines changed: 8 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -20,12 +20,13 @@ void rotary_embedding_inplace(at::Tensor &q,
2020
at::Tensor &k,
2121
const at::Tensor &cos,
2222
const at::Tensor &sin,
23-
::std::optional<at::Tensor> position_ids = ::std::nullopt,
23+
const std::optional<at::Tensor> &position_ids = std::nullopt,
2424
bool rotary_interleaved = false);
25-
std::tuple<at::Tensor, at::Tensor> rotary_embedding(const at::Tensor &q,
26-
const at::Tensor &k,
27-
const at::Tensor &cos,
28-
const at::Tensor &sin,
29-
::std::optional<at::Tensor> position_ids = ::std::nullopt,
30-
bool rotary_interleaved = false);
25+
std::tuple<at::Tensor, at::Tensor> rotary_embedding(
26+
const at::Tensor &q,
27+
const at::Tensor &k,
28+
const at::Tensor &cos,
29+
const at::Tensor &sin,
30+
const std::optional<at::Tensor> &position_ids = std::nullopt,
31+
bool rotary_interleaved = false);
3132
} // namespace flag_gems

lib/rotary_embedding.cpp

Lines changed: 18 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -98,21 +98,20 @@ void rotary_embedding_inplace(
9898
at::Tensor& k, // [batch_size, seq_len, k_heads, head_dim] or [num_tokens, k_heads, head_dim]
9999
const at::Tensor& cos, // [max_seq_len, head_dim // 2]
100100
const at::Tensor& sin, // [max_seq_len, head_dim // 2]
101-
std::optional<at::Tensor> position_ids, // None or [..., seq_len]
102-
bool rotary_interleaved) { // default false
101+
const std::optional<at::Tensor>& position_ids, // None or [..., seq_len]
102+
bool rotary_interleaved) { // default false
103103

104104
check_rotary_embedding_inputs(q, k, cos, sin, position_ids);
105105

106106
auto q_sizes = q.sizes();
107107
auto k_sizes = k.sizes();
108-
std::optional<int64_t> seq_len;
109108

109+
std::optional<int64_t> seq_len = std::nullopt;
110+
std::optional<at::Tensor> flat_position_ids = std::nullopt;
110111
if (!position_ids.has_value()) {
111112
seq_len = q_sizes[1];
112-
113-
} else { // default case
114-
position_ids = position_ids.value().view({-1}); // flatten the position_ids tensor
115-
seq_len = std::nullopt;
113+
} else { // default case
114+
flat_position_ids = position_ids.value().view({-1}); // flatten the position_ids tensor
116115
}
117116

118117
q = q.view({-1, q.size(-2), q.size(-1)}); // [num_tokens, q_heads, head_dim]
@@ -167,14 +166,15 @@ def apply_rotary_pos_emb_inplace_kernel(
167166
k,
168167
cos,
169168
sin,
170-
position_ids, // std::optional<at::Tensor>
169+
flat_position_ids, // std::optional<at::Tensor>
171170
q.stride(0),
172171
q.stride(1),
173172
q.stride(2),
174173
k.stride(0),
175174
k.stride(1),
176175
k.stride(2),
177-
position_ids.has_value() ? position_ids.value().stride(0) : 0, // 0 if position_ids is not defined
176+
flat_position_ids.has_value() ? flat_position_ids.value().stride(0)
177+
: 0, // 0 if flat_position_ids is not defined
178178
cos.stride(0),
179179
sin.stride(0),
180180
seq_len, // std::optional<long int>
@@ -197,21 +197,20 @@ std::tuple<at::Tensor, at::Tensor> rotary_embedding(const at::Tensor& q,
197197
const at::Tensor& k,
198198
const at::Tensor& cos,
199199
const at::Tensor& sin,
200-
std::optional<at::Tensor> position_ids,
200+
const std::optional<at::Tensor>& position_ids,
201201
bool rotary_interleaved) {
202202
// Check inputs
203203
check_rotary_embedding_inputs(q, k, cos, sin, position_ids);
204204

205205
auto q_sizes = q.sizes();
206206
auto k_sizes = k.sizes();
207-
std::optional<int64_t> seq_len;
207+
std::optional<int64_t> seq_len = std::nullopt;
208+
std::optional<at::Tensor> flat_position_ids = std::nullopt;
208209

209210
if (!position_ids.has_value()) {
210211
seq_len = q_sizes[1];
211-
212-
} else { // default case
213-
position_ids = position_ids.value().view({-1}); // flatten the position_ids tensor
214-
seq_len = std::nullopt;
212+
} else { // default case
213+
flat_position_ids = position_ids.value().view({-1}); // flatten the position_ids tensor
215214
}
216215

217216
auto q_view = q.view({-1, q.size(-2), q.size(-1)}); // [num_tokens, q_heads, head_dim]
@@ -248,7 +247,7 @@ std::tuple<at::Tensor, at::Tensor> rotary_embedding(const at::Tensor& q,
248247
k_view,
249248
cos,
250249
sin,
251-
position_ids, // std::optional<at::Tensor>
250+
flat_position_ids, // std::optional<at::Tensor>
252251
q_view.stride(0),
253252
q_view.stride(1),
254253
q_view.stride(2),
@@ -261,7 +260,8 @@ std::tuple<at::Tensor, at::Tensor> rotary_embedding(const at::Tensor& q,
261260
k_embed_stride[0],
262261
k_embed_stride[1],
263262
k_embed_stride[2],
264-
position_ids.has_value() ? position_ids.value().stride(0) : 0, // 0 if position_ids is not defined
263+
flat_position_ids.has_value() ? flat_position_ids.value().stride(0)
264+
: 0, // 0 if flat_position_ids is not defined
265265
cos.stride(0),
266266
sin.stride(0),
267267
seq_len, // std::optional<long int>
@@ -276,6 +276,6 @@ std::tuple<at::Tensor, at::Tensor> rotary_embedding(const at::Tensor& q,
276276
// Reshape back to original shapes
277277
q_embed = q_embed.view(q_sizes.vec());
278278
k_embed = k_embed.view(k_sizes.vec());
279-
return {q_embed, k_embed};
279+
return std::make_tuple(q_embed, k_embed);
280280
}
281281
} // namespace flag_gems

0 commit comments

Comments
 (0)