Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 10 additions & 2 deletions miles/backends/sglang_utils/sglang_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Expand Down Expand Up @@ -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:
Expand Down
10 changes: 8 additions & 2 deletions miles/backends/training_utils/weight_update/updater.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__)
Expand Down Expand Up @@ -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)
5 changes: 5 additions & 0 deletions miles/utils/lora/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
28 changes: 28 additions & 0 deletions tests/fast/backends/sglang_utils/test_compute_server_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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}]
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand All @@ -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
)
Loading