From 59da32bca09fce31679f4b4a2f7add36e10a5053 Mon Sep 17 00:00:00 2001 From: cherry77-cloud <234872336+taking-lying-flat@users.noreply.github.com> Date: Fri, 7 Aug 2026 06:13:32 +0800 Subject: [PATCH] fix flash-decode RoPE caching --- src/mcore_bridge/model/gpt_model.py | 9 +++++---- 1 file changed, 5 insertions(+), 4 deletions(-) diff --git a/src/mcore_bridge/model/gpt_model.py b/src/mcore_bridge/model/gpt_model.py index 2158b90..7a6f4fd 100644 --- a/src/mcore_bridge/model/gpt_model.py +++ b/src/mcore_bridge/model/gpt_model.py @@ -216,10 +216,11 @@ def _get_rotary_pos_emb(self, decoder_input, position_ids, packed_seq_params, in assert (inference_context.is_static_batching() ), 'GPTModel currently only supports static inference batching.' # Flash decoding uses precomputed cos and sin for RoPE - rotary_pos_cos, rotary_pos_sin = self.rotary_pos_emb_cache.setdefault( - inference_context.max_sequence_length, - self.rotary_pos_emb.get_cos_sin(inference_context.max_sequence_length), - ) + max_sequence_length = inference_context.max_sequence_length + if max_sequence_length not in self.rotary_pos_emb_cache: + self.rotary_pos_emb_cache[max_sequence_length] = self.rotary_pos_emb.get_cos_sin( + max_sequence_length) + rotary_pos_cos, rotary_pos_sin = self.rotary_pos_emb_cache[max_sequence_length] else: rotary_seq_len = RotaryEmbedding.get_rotary_seq_len(self, inference_context, self.decoder, decoder_input, self.config, packed_seq_params)