From 6e7ee7bad430b5a24969cb54bebe2f2ebe94ca05 Mon Sep 17 00:00:00 2001 From: Rahul Chalamala <22563365+rchalamala@users.noreply.github.com> Date: Thu, 17 Sep 2026 08:07:20 -0700 Subject: [PATCH] mm: MultimodalInputs.merge invalidates the M-RoPE delta decode cache and adopts the other's delta (#87) When an earlier session turn had no M-RoPE delta, merge() dropped the new turn's delta, and it never cleared mrope_position_delta_repeated_cache, so decode positions after a multimodal session continuation came from a stale or missing delta. Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- python/sglang/srt/managers/schedule_batch.py | 3 ++ .../test_mm_inputs_merge_mrope_delta.py | 43 +++++++++++++++++++ 2 files changed, 46 insertions(+) create mode 100644 test/registered/unit/managers/test_mm_inputs_merge_mrope_delta.py diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 52c72a0e24b2..507b13bddafe 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -888,6 +888,9 @@ def merge(self, other: MultimodalInputs): self.mrope_position_delta = torch.cat( [self.mrope_position_delta, other.mrope_position_delta], dim=0 ) + elif other.mrope_position_delta is not None: + self.mrope_position_delta = other.mrope_position_delta + self.mrope_position_delta_repeated_cache = None for key, val in other.__dict__.items(): if "_id" in key: diff --git a/test/registered/unit/managers/test_mm_inputs_merge_mrope_delta.py b/test/registered/unit/managers/test_mm_inputs_merge_mrope_delta.py new file mode 100644 index 000000000000..518147cc4a32 --- /dev/null +++ b/test/registered/unit/managers/test_mm_inputs_merge_mrope_delta.py @@ -0,0 +1,43 @@ +import unittest + +import torch + +from sglang.srt.managers.schedule_batch import ( + Modality, + MultimodalDataItem, + MultimodalInputs, +) +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=5, suite="base-a-test-cpu") + +IM_TOKEN_ID = 7 + + +class TestMultimodalInputsMergeMropeDelta(CustomTestCase): + def test_merge_invalidates_mrope_delta_cache(self): + def _mm(delta): + item = MultimodalDataItem( + modality=Modality.IMAGE, hash=7, pad_value=7, offsets=[(0, 0)] + ) + return MultimodalInputs( + mm_items=[item], + im_token_id=IM_TOKEN_ID, + mrope_position_delta=torch.tensor([[delta]]), + ) + + base = _mm(3) + base.mrope_position_delta_repeated_cache = torch.zeros(3, 1, dtype=torch.long) + base.merge(_mm(5)) + self.assertIsNone(base.mrope_position_delta_repeated_cache) + self.assertEqual(base.mrope_position_delta.flatten().tolist(), [3, 5]) + + no_delta = _mm(1) + no_delta.mrope_position_delta = None + no_delta.merge(_mm(9)) + self.assertEqual(no_delta.mrope_position_delta.flatten().tolist(), [9]) + + +if __name__ == "__main__": + unittest.main()