@@ -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