Skip to content
Original file line number Diff line number Diff line change
Expand Up @@ -139,7 +139,8 @@ def _update_expert_bucket_weights(
def _pause_and_prepare_engines(self) -> None:
"""Pause rollout engines, flush cache, and run pre-process if needed."""
if dist.get_rank() == 0:
ray.get([engine.pause_generation.remote() for engine in self.rollout_engines])
mode = self.args.pause_generation_mode
ray.get([engine.pause_generation.remote(mode=mode) for engine in self.rollout_engines])
ray.get([engine.flush_cache.remote() for engine in self.rollout_engines])

# int4/fp4 pre_process
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -176,7 +176,8 @@ def update_weights(self) -> None:

rank = dist.get_rank()
if rank == 0:
ray.get([engine.pause_generation.remote() for engine in self.rollout_engines])
mode = self.args.pause_generation_mode
ray.get([engine.pause_generation.remote(mode=mode) for engine in self.rollout_engines])
ray.get([engine.flush_cache.remote() for engine in self.rollout_engines])
if self.quantization_config and self.quantization_config["quant_method"] in ["compressed-tensors"]:
post_process_weights(
Expand Down
7 changes: 5 additions & 2 deletions miles/backends/sglang_utils/sglang_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -501,8 +501,11 @@ def update_weights_from_distributed(
payload,
)

def pause_generation(self):
response = requests.post(f"http://{self.server_host}:{self.server_port}/pause_generation", json={})
def pause_generation(self, mode: str = "retract"):
response = requests.post(
f"http://{self.server_host}:{self.server_port}/pause_generation",
json={"mode": mode},
)
response.raise_for_status()
return response

Expand Down
13 changes: 13 additions & 0 deletions miles/utils/arguments.py
Original file line number Diff line number Diff line change
Expand Up @@ -452,6 +452,19 @@ def add_rollout_arguments(parser):
default=1,
help="Interval for updating the weights",
)
parser.add_argument(
"--pause-generation-mode",
type=str,
choices=["abort", "retract", "in_place"],
default="retract",
help=(
"How SGLang pauses in-flight requests during weight updates. "
"'abort' immediately terminates all requests (previous default). "
"'retract' moves running requests back to the waiting queue and "
"recomputes KV cache after update. "
"'in_place' freezes requests and resumes with existing KV cache."
),
)
parser.add_argument(
"--keep-old-actor",
action="store_true",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -56,6 +56,7 @@ def _make_args(**overrides):
update_weight_buffer_size=1 << 30,
actor_num_nodes=1,
actor_num_gpus_per_node=1,
pause_generation_mode="retract",
)
defaults.update(overrides)
return Namespace(**defaults)
Expand Down
Loading