diff --git a/miles/backends/sglang_utils/sglang_engine.py b/miles/backends/sglang_utils/sglang_engine.py index 1f4acafe63f..1002c8ca3de 100644 --- a/miles/backends/sglang_utils/sglang_engine.py +++ b/miles/backends/sglang_utils/sglang_engine.py @@ -11,6 +11,7 @@ LORA_ADAPTER_NAME, engine_loads_adapter_from_disk, is_multi_lora_enabled, + lora_adapter_pinned, lora_base_cpu_backup_enabled, lora_rollout_enabled, ) @@ -155,14 +156,21 @@ def _compute_server_args( ) elif lora_rollout_enabled(args): kwargs["enable_lora"] = True - kwargs["max_loras_per_batch"] = 1 + # SGLang's DP-attention LoRA only serves pinned adapters, and a pinned adapter + # needs a second slot: the anti-starvation check keeps one for the base model. + kwargs["max_loras_per_batch"] = 2 if lora_adapter_pinned(args) else 1 kwargs["max_lora_rank"] = max(getattr(args, "lora_rank", 0), 1) kwargs["lora_target_modules"] = ( ["all"] if args.lora_adapter_targets == "all-linear" else args.lora_adapter_targets ) if engine_loads_adapter_from_disk(args): - kwargs["lora_paths"] = [f"{LORA_ADAPTER_NAME}={args.lora_adapter_path}"] + if lora_adapter_pinned(args): + kwargs["lora_paths"] = [ + {"lora_name": LORA_ADAPTER_NAME, "lora_path": args.lora_adapter_path, "pinned": True} + ] + else: + kwargs["lora_paths"] = [f"{LORA_ADAPTER_NAME}={args.lora_adapter_path}"] elif args.lora_adapter_path is not None: logger.info("Skipping startup lora_paths: the trainer pushes the adapter in the first weight sync") else: diff --git a/miles/backends/training_utils/weight_update/updater.py b/miles/backends/training_utils/weight_update/updater.py index eca96319c23..94281b4c945 100644 --- a/miles/backends/training_utils/weight_update/updater.py +++ b/miles/backends/training_utils/weight_update/updater.py @@ -28,7 +28,7 @@ ) from miles.backends.training_utils.weight_update.utils import record_lora_checksums from miles.utils.distributed_utils import get_gloo_group -from miles.utils.lora.utils import LORA_ADAPTER_NAME +from miles.utils.lora.utils import LORA_ADAPTER_NAME, lora_adapter_pinned from miles.utils.timer import timer logger = logging.getLogger(__name__) @@ -160,5 +160,11 @@ def _register_new_lora_adapters(self, rollout_engines, adapters: list[tuple[str, config = self._lora_sync_config if adapter is not None: config = config | {"r": adapter.rank, "lora_alpha": adapter.alpha} - register_lora_adapter(rollout_engines, lora_name=lora_name, lora_config=config) + register_lora_adapter( + rollout_engines, + lora_name=lora_name, + lora_config=config, + # SGLang's DP-attention LoRA needs the adapter in the same slot on every DP rank. + pinned=lora_adapter_pinned(self.args), + ) self._registered_adapters.add(lora_name) diff --git a/miles/utils/lora/utils.py b/miles/utils/lora/utils.py index 34962c2de19..10d1f104c07 100644 --- a/miles/utils/lora/utils.py +++ b/miles/utils/lora/utils.py @@ -31,6 +31,11 @@ def lora_rollout_enabled(args: Namespace) -> bool: return is_lora_enabled(args) and not getattr(args, "lora_train_only", False) +def lora_adapter_pinned(args: Namespace) -> bool: + """SGLang's DP-attention LoRA keeps an adapter in the same slot on every DP rank, so it serves only pinned adapters.""" + return getattr(args, "sglang_enable_dp_attention", False) + + def engine_loads_adapter_from_disk(args: Namespace) -> bool: """Only when no trainer will push the adapter; otherwise the first weight sync carries it.""" return args.lora_adapter_path is not None and (args.debug_rollout_only or args.debug_skip_weight_update) diff --git a/tests/fast/backends/sglang_utils/test_compute_server_args.py b/tests/fast/backends/sglang_utils/test_compute_server_args.py index 0c6c6694a69..a7d62025219 100644 --- a/tests/fast/backends/sglang_utils/test_compute_server_args.py +++ b/tests/fast/backends/sglang_utils/test_compute_server_args.py @@ -18,6 +18,7 @@ def make_args(**overrides: object) -> SimpleNamespace: rollout_num_gpus_per_engine=1, offload_rollout=False, sglang_dp_size=1, + sglang_enable_dp_attention=False, sglang_pp_size=1, sglang_ep_size=1, sglang_mem_fraction_static=0.7, @@ -166,3 +167,30 @@ def test_without_a_trainer_push_the_engine_loads_the_adapter_itself(self, flag): server_args = compute(make_args(lora_rank=8, lora_adapter_path="/fake/adapter", **{flag: True})) assert server_args["lora_paths"] == ["miles_lora=/fake/adapter"] + + +class TestDpAttentionLoRA: + """SGLang's DP-attention LoRA serves only pinned adapters, and a pinned adapter needs a slot besides the base model's.""" + + def test_dp_attention_reserves_a_slot_beside_the_base_model(self): + server_args = compute(make_args(lora_rank=8, sglang_enable_dp_attention=True, sglang_dp_size=8)) + + assert server_args["max_loras_per_batch"] == 2 + + def test_without_dp_attention_the_engine_keeps_a_single_slot(self): + server_args = compute(make_args(lora_rank=8)) + + assert server_args["max_loras_per_batch"] == 1 + + def test_dp_attention_pins_the_adapter_the_engine_loads_itself(self): + server_args = compute( + make_args( + lora_rank=8, + lora_adapter_path="/fake/adapter", + debug_rollout_only=True, + sglang_enable_dp_attention=True, + sglang_dp_size=8, + ) + ) + + assert server_args["lora_paths"] == [{"lora_name": "miles_lora", "lora_path": "/fake/adapter", "pinned": True}] diff --git a/tests/fast/backends/training_utils/weight_update/test_lora_update_weight.py b/tests/fast/backends/training_utils/weight_update/test_lora_update_weight.py index 2438df747f6..9d7b67f043c 100644 --- a/tests/fast/backends/training_utils/weight_update/test_lora_update_weight.py +++ b/tests/fast/backends/training_utils/weight_update/test_lora_update_weight.py @@ -64,12 +64,12 @@ def test_base_names_do_not_contain_lora(self): class TestWeightUpdaterLoraConfig: """The updater requires a lora_sync_config exactly when LoRA is active.""" - def _make_updater(self, *, is_lora, lora_sync_config): + def _make_updater(self, *, is_lora, lora_sync_config, args=None): protocol = MagicMock() protocol.supports_lora = True with patch(f"{_UPDATER_MODULE}.get_weight_transfer_protocol", return_value=protocol): return WeightUpdater( - Namespace(), + args or Namespace(), [MagicMock()], weights_getter=lambda: {}, model_name="qwen", @@ -91,3 +91,18 @@ def test_lora_config_stored(self): def test_no_lora_no_config(self): updater = self._make_updater(is_lora=False, lora_sync_config=None) assert updater._lora_sync_config is None + + @pytest.mark.parametrize("dp_attention", [False, True], ids=["tp", "dp-attention"]) + def test_registration_pins_the_adapter_only_under_dp_attention(self, dp_attention): + """SGLang's DP-attention LoRA rejects unpinned adapters; other layouts keep them evictable.""" + updater = self._make_updater( + is_lora=True, + lora_sync_config={"peft_type": "LORA", "r": 8}, + args=Namespace(sglang_enable_dp_attention=dp_attention), + ) + engines = [MagicMock()] + with patch(f"{_UPDATER_MODULE}.register_lora_adapter") as register: + updater._register_new_lora_adapters(engines, [("miles_lora", None)]) + register.assert_called_once_with( + engines, lora_name="miles_lora", lora_config={"peft_type": "LORA", "r": 8}, pinned=dp_attention + )