recover patch files

This commit is contained in:
Yusheng Su
2025-07-01 07:29:09 +00:00
parent a127b10405
commit 990fc633c7
3 changed files with 1043 additions and 1 deletions
+136
View File
@@ -0,0 +1,136 @@
diff --git a/megatron/core/models/gpt/gpt_layer_specs.py b/megatron/core/models/gpt/gpt_layer_specs.py
index 07fbd64f..cc8b8ec7 100755
--- a/megatron/core/models/gpt/gpt_layer_specs.py
+++ b/megatron/core/models/gpt/gpt_layer_specs.py
@@ -74,6 +74,8 @@ def get_gpt_layer_with_transformer_engine_spec(
qk_l2_norm: Optional[bool] = False,
use_te_op_fuser: Optional[bool] = False,
use_kitchen: bool = False,
+ post_self_attn_layernorm: bool = False,
+ post_mlp_layernorm: bool = False,
) -> ModuleSpec:
"""Use this spec to use lower-level Transformer Engine modules (required for fp8 training).
@@ -173,9 +175,11 @@ def get_gpt_layer_with_transformer_engine_spec(
),
),
self_attn_bda=get_bias_dropout_add,
+ post_self_attn_layernorm=TENorm if post_self_attn_layernorm else IdentityOp,
pre_mlp_layernorm=backend.layer_norm() if num_experts else IdentityOp,
mlp=mlp,
mlp_bda=get_bias_dropout_add,
+ post_mlp_layernorm=TENorm if post_mlp_layernorm else IdentityOp,
sharded_state_dict_keys_map={
'mlp.0.weight': 'mlp.linear_fc1.layer_norm_weight',
'mlp.0.bias': 'mlp.linear_fc1.layer_norm_bias',
diff --git a/megatron/core/transformer/transformer_config.py b/megatron/core/transformer/transformer_config.py
index a256d70f..0d436eb0 100644
--- a/megatron/core/transformer/transformer_config.py
+++ b/megatron/core/transformer/transformer_config.py
@@ -167,6 +167,9 @@ class TransformerConfig(ModelParallelConfig):
qk_layernorm: bool = False
"""Whether to apply `normalization` type of normalization to the query and key embeddings."""
+ post_self_attn_layernorm: bool = False
+ post_mlp_layernorm: bool = False
+
test_mode: bool = False
"""Whether to run real-time tests."""
diff --git a/megatron/core/transformer/transformer_layer.py b/megatron/core/transformer/transformer_layer.py
index 51c90f5a..05c34955 100644
--- a/megatron/core/transformer/transformer_layer.py
+++ b/megatron/core/transformer/transformer_layer.py
@@ -228,6 +228,7 @@ class TransformerLayerSubmodules:
input_layernorm: Union[ModuleSpec, type] = IdentityOp
self_attention: Union[ModuleSpec, type] = IdentityOp
self_attn_bda: Union[ModuleSpec, type] = IdentityFuncOp
+ post_self_attn_layernorm: Union[ModuleSpec, type] = IdentityOp
pre_cross_attn_layernorm: Union[ModuleSpec, type] = IdentityOp
cross_attention: Union[ModuleSpec, type] = IdentityOp
@@ -236,6 +237,7 @@ class TransformerLayerSubmodules:
pre_mlp_layernorm: Union[ModuleSpec, type] = IdentityOp
mlp: Union[ModuleSpec, type] = IdentityOp
mlp_bda: Union[ModuleSpec, type] = IdentityFuncOp
+ post_mlp_layernorm: Union[ModuleSpec, type] = IdentityOp
# Mapping for sharded tensor keys to be applied in `sharded_state_dict` method
sharded_state_dict_keys_map: Dict[str, str] = field(default_factory=dict)
@@ -338,6 +340,14 @@ class TransformerLayer(MegatronModule, BaseTransformerLayer):
# [Module 3: BiasDropoutFusion]
self.self_attn_bda = build_module(submodules.self_attn_bda)
+ self.post_self_attn_layernorm = build_module(
+ submodules.post_self_attn_layernorm,
+ config=self.config,
+ hidden_size=self.config.hidden_size,
+ eps=self.config.layernorm_epsilon,
+ )
+
+
# [Module 4: Post SelfAttention] Optional Layernorm after self-attn
self.pre_cross_attn_layernorm = build_module(
submodules.pre_cross_attn_layernorm,
@@ -395,6 +405,13 @@ class TransformerLayer(MegatronModule, BaseTransformerLayer):
# [Module 9: BiasDropoutFusion]
self.mlp_bda = build_module(submodules.mlp_bda)
+ self.post_mlp_layernorm = build_module(
+ submodules.post_mlp_layernorm,
+ config=self.config,
+ hidden_size=self.config.hidden_size,
+ eps=self.config.layernorm_epsilon
+ )
+
self.recompute_input_layernorm = False
self.recompute_pre_mlp_layernorm = False
self.recompute_mlp = False
@@ -518,6 +535,11 @@ class TransformerLayer(MegatronModule, BaseTransformerLayer):
attention_output_with_bias[0]
)
+ attention_output, attention_output_bias = attention_output_with_bias
+ attention_output = self.post_self_attn_layernorm(attention_output)
+ attention_output_with_bias = (attention_output, attention_output_bias)
+
+
# TODO: could we move `bias_dropout_add_exec_handler` itself
# inside the module provided in the `bias_dropout_add_spec` module?
nvtx_range_push(suffix="self_attn_bda")
@@ -618,6 +640,10 @@ class TransformerLayer(MegatronModule, BaseTransformerLayer):
else:
mlp_output_with_bias = self.mlp(pre_mlp_layernorm_output)
+ mlp_output, mlp_output_bias = mlp_output_with_bias
+ mlp_output = self.post_mlp_layernorm(mlp_output)
+ mlp_output_with_bias = (mlp_output, mlp_output_bias)
+
if self.recompute_pre_mlp_layernorm:
# discard the output of the pre-mlp layernorm and register the recompute
# as a gradient hook of mlp_output_with_bias[0]
diff --git a/megatron/training/arguments.py b/megatron/training/arguments.py
index 7b1b35b1..c0c510e9 100644
--- a/megatron/training/arguments.py
+++ b/megatron/training/arguments.py
@@ -1158,6 +1158,9 @@ def core_transformer_config_from_args(args, config_class=None):
if args.is_hybrid_model:
kw_args['is_hybrid_model'] = args.is_hybrid_model
+ kw_args['post_self_attn_layernorm'] = args.post_self_attn_layernorm
+ kw_args['post_mlp_layernorm'] = args.post_mlp_layernorm
+
# handle quantization config
# NOTE: Kitchen arguments are only added to the namespace when
# Kitchen library is available.
@@ -1439,6 +1442,10 @@ def _add_network_size_args(parser):
action='store_true',
help='If set, use original BERT residula connection '
'ordering.')
+ group.add_argument('--post-self-attn-layernorm', action='store_true',
+ help='If set, use post self attention layernorm.')
+ group.add_argument('--post-mlp-layernorm', action='store_true',
+ help='If set, use post MLP layernorm.')
group.add_argument('--openai-gelu', action='store_true',
help='Use OpenAIs GeLU implementation. This option'
'should not be used unless for backward compatibility'
+906
View File
@@ -0,0 +1,906 @@
diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py
index 0b641d34..bc7cba90 100644
--- a/python/sglang/srt/configs/model_config.py
+++ b/python/sglang/srt/configs/model_config.py
@@ -390,15 +390,15 @@ class ModelConfig:
if quant_cfg is not None:
quant_method = quant_cfg.get("quant_method", "").lower()
- # Detect which checkpoint is it
- for _, method in QUANTIZATION_METHODS.items():
- quantization_override = method.override_quantization_method(
- quant_cfg, self.quantization
- )
- if quantization_override:
- quant_method = quantization_override
- self.quantization = quantization_override
- break
+ # # Detect which checkpoint is it
+ # for _, method in QUANTIZATION_METHODS.items():
+ # quantization_override = method.override_quantization_method(
+ # quant_cfg, self.quantization
+ # )
+ # if quantization_override:
+ # quant_method = quantization_override
+ # self.quantization = quantization_override
+ # break
# Verify quantization configurations.
if self.quantization is None:
diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py
index 75bccc9d..0ca97353 100644
--- a/python/sglang/srt/entrypoints/engine.py
+++ b/python/sglang/srt/entrypoints/engine.py
@@ -403,12 +403,15 @@ class Engine(EngineBase):
self.tokenizer_manager.init_weights_update_group(obj, None)
)
- def update_weights_from_distributed(self, name: str, dtype, shape):
+ def update_weights_from_distributed(
+ self, names: str, dtypes, shapes, group_name: str
+ ):
"""Update weights from distributed source."""
obj = UpdateWeightsFromDistributedReqInput(
- name=name,
- dtype=dtype,
- shape=shape,
+ names=names,
+ dtypes=dtypes,
+ shapes=shapes,
+ group_name=group_name,
)
loop = asyncio.get_event_loop()
return loop.run_until_complete(
diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py
index 06c67343..672b294d 100644
--- a/python/sglang/srt/entrypoints/http_server.py
+++ b/python/sglang/srt/entrypoints/http_server.py
@@ -605,6 +605,20 @@ async def separate_reasoning_request(obj: SeparateReasoningReqInput, request: Re
return ORJSONResponse(content=response_data, status_code=200)
+@app.post("/pause_generation")
+async def pause_generation(request: Request):
+ """Pause generation."""
+ await _global_state.tokenizer_manager.pause_generation()
+ return ORJSONResponse(content="Generation paused.", status_code=200)
+
+
+@app.post("/continue_generation")
+async def continue_generation(request: Request):
+ """Continue generation."""
+ await _global_state.tokenizer_manager.continue_generation()
+ return ORJSONResponse(content="Generation continued.", status_code=200)
+
+
##### OpenAI-compatible API endpoints #####
@@ -806,104 +820,105 @@ def _wait_and_warmup(
image_token_text: str,
launch_callback: Optional[Callable[[], None]] = None,
):
- headers = {}
- url = server_args.url()
- if server_args.api_key:
- headers["Authorization"] = f"Bearer {server_args.api_key}"
+ if not server_args.skip_warmup:
+ headers = {}
+ url = server_args.url()
+ if server_args.api_key:
+ headers["Authorization"] = f"Bearer {server_args.api_key}"
+
+ # Wait until the server is launched
+ success = False
+ for _ in range(120):
+ time.sleep(1)
+ try:
+ res = requests.get(url + "/get_model_info", timeout=5, headers=headers)
+ assert res.status_code == 200, f"{res=}, {res.text=}"
+ success = True
+ break
+ except (AssertionError, requests.exceptions.RequestException):
+ last_traceback = get_exception_traceback()
+ pass
+
+ if not success:
+ if pipe_finish_writer is not None:
+ pipe_finish_writer.send(last_traceback)
+ logger.error(f"Initialization failed. warmup error: {last_traceback}")
+ kill_process_tree(os.getpid())
+ return
+
+ model_info = res.json()
+
+ # Send a warmup request
+ request_name = "/generate" if model_info["is_generation"] else "/encode"
+ max_new_tokens = 8 if model_info["is_generation"] else 1
+ json_data = {
+ "sampling_params": {
+ "temperature": 0,
+ "max_new_tokens": max_new_tokens,
+ },
+ }
+ if server_args.skip_tokenizer_init:
+ json_data["input_ids"] = [[10, 11, 12] for _ in range(server_args.dp_size)]
+ # TODO Workaround the bug that embedding errors for list of size 1
+ if server_args.dp_size == 1:
+ json_data["input_ids"] = json_data["input_ids"][0]
+ else:
+ json_data["text"] = ["The capital city of France is"] * server_args.dp_size
+ # TODO Workaround the bug that embedding errors for list of size 1
+ if server_args.dp_size == 1:
+ json_data["text"] = json_data["text"][0]
+
+ # Debug dumping
+ if server_args.debug_tensor_dump_input_file:
+ json_data.pop("text", None)
+ json_data["input_ids"] = np.load(
+ server_args.debug_tensor_dump_input_file
+ ).tolist()
+ json_data["sampling_params"]["max_new_tokens"] = 0
- # Wait until the server is launched
- success = False
- for _ in range(120):
- time.sleep(1)
try:
- res = requests.get(url + "/get_model_info", timeout=5, headers=headers)
- assert res.status_code == 200, f"{res=}, {res.text=}"
- success = True
- break
- except (AssertionError, requests.exceptions.RequestException):
+ if server_args.disaggregation_mode == "null":
+ res = requests.post(
+ url + request_name,
+ json=json_data,
+ headers=headers,
+ timeout=600,
+ )
+ assert res.status_code == 200, f"{res}"
+ else:
+ logger.info(f"Start of prefill warmup ...")
+ json_data = {
+ "sampling_params": {
+ "temperature": 0.0,
+ "max_new_tokens": 8,
+ "ignore_eos": True,
+ },
+ "bootstrap_host": [FakeBootstrapHost] * server_args.dp_size,
+ # This is a hack to ensure fake transfer is enabled during prefill warmup
+ # ensure each dp rank has a unique bootstrap_room during prefill warmup
+ "bootstrap_room": [
+ i * (2**63 // server_args.dp_size) + (i % server_args.tp_size)
+ for i in range(server_args.dp_size)
+ ],
+ "input_ids": [[0, 1, 2, 3]] * server_args.dp_size,
+ }
+ res = requests.post(
+ url + request_name,
+ json=json_data,
+ headers=headers,
+ timeout=1800, # because of deep gemm precache is very long if not precache.
+ )
+ logger.info(
+ f"End of prefill warmup with status {res.status_code}, resp: {res.json()}"
+ )
+
+ except Exception:
last_traceback = get_exception_traceback()
- pass
-
- if not success:
- if pipe_finish_writer is not None:
- pipe_finish_writer.send(last_traceback)
- logger.error(f"Initialization failed. warmup error: {last_traceback}")
- kill_process_tree(os.getpid())
- return
-
- model_info = res.json()
-
- # Send a warmup request
- request_name = "/generate" if model_info["is_generation"] else "/encode"
- max_new_tokens = 8 if model_info["is_generation"] else 1
- json_data = {
- "sampling_params": {
- "temperature": 0,
- "max_new_tokens": max_new_tokens,
- },
- }
- if server_args.skip_tokenizer_init:
- json_data["input_ids"] = [[10, 11, 12] for _ in range(server_args.dp_size)]
- # TODO Workaround the bug that embedding errors for list of size 1
- if server_args.dp_size == 1:
- json_data["input_ids"] = json_data["input_ids"][0]
- else:
- json_data["text"] = ["The capital city of France is"] * server_args.dp_size
- # TODO Workaround the bug that embedding errors for list of size 1
- if server_args.dp_size == 1:
- json_data["text"] = json_data["text"][0]
-
- # Debug dumping
- if server_args.debug_tensor_dump_input_file:
- json_data.pop("text", None)
- json_data["input_ids"] = np.load(
- server_args.debug_tensor_dump_input_file
- ).tolist()
- json_data["sampling_params"]["max_new_tokens"] = 0
-
- try:
- if server_args.disaggregation_mode == "null":
- res = requests.post(
- url + request_name,
- json=json_data,
- headers=headers,
- timeout=600,
- )
- assert res.status_code == 200, f"{res}"
- else:
- logger.info(f"Start of prefill warmup ...")
- json_data = {
- "sampling_params": {
- "temperature": 0.0,
- "max_new_tokens": 8,
- "ignore_eos": True,
- },
- "bootstrap_host": [FakeBootstrapHost] * server_args.dp_size,
- # This is a hack to ensure fake transfer is enabled during prefill warmup
- # ensure each dp rank has a unique bootstrap_room during prefill warmup
- "bootstrap_room": [
- i * (2**63 // server_args.dp_size) + (i % server_args.tp_size)
- for i in range(server_args.dp_size)
- ],
- "input_ids": [[0, 1, 2, 3]] * server_args.dp_size,
- }
- res = requests.post(
- url + request_name,
- json=json_data,
- headers=headers,
- timeout=1800, # because of deep gemm precache is very long if not precache.
- )
- logger.info(
- f"End of prefill warmup with status {res.status_code}, resp: {res.json()}"
- )
-
- except Exception:
- last_traceback = get_exception_traceback()
- if pipe_finish_writer is not None:
- pipe_finish_writer.send(last_traceback)
- logger.error(f"Initialization failed. warmup error: {last_traceback}")
- kill_process_tree(os.getpid())
- return
+ if pipe_finish_writer is not None:
+ pipe_finish_writer.send(last_traceback)
+ logger.error(f"Initialization failed. warmup error: {last_traceback}")
+ kill_process_tree(os.getpid())
+ return
# Debug print
# logger.info(f"{res.json()=}")
diff --git a/python/sglang/srt/layers/quantization/fp8.py b/python/sglang/srt/layers/quantization/fp8.py
index c779f1f1..a3b1db7c 100644
--- a/python/sglang/srt/layers/quantization/fp8.py
+++ b/python/sglang/srt/layers/quantization/fp8.py
@@ -325,10 +325,10 @@ class Fp8LinearMethod(LinearMethodBase):
layer.input_scale = None
else:
weight, weight_scale = layer.weight.data, layer.weight_scale_inv.data
- layer.weight = torch.nn.Parameter(weight, requires_grad=False)
- layer.weight_scale_inv = torch.nn.Parameter(
- weight_scale, requires_grad=False
- )
+ # layer.weight = torch.nn.Parameter(weight, requires_grad=False)
+ # layer.weight_scale_inv = torch.nn.Parameter(
+ # weight_scale, requires_grad=False
+ # )
return
layer.weight = torch.nn.Parameter(layer.weight.data, requires_grad=False)
diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py
index 40c220c4..153be20a 100644
--- a/python/sglang/srt/managers/io_struct.py
+++ b/python/sglang/srt/managers/io_struct.py
@@ -727,9 +727,11 @@ class UpdateWeightFromDiskReqOutput:
@dataclass
class UpdateWeightsFromDistributedReqInput:
- name: str
- dtype: str
- shape: List[int]
+ names: List[str]
+ dtypes: List[str]
+ shapes: List[List[int]]
+ # The group name
+ group_name: str = "weight_update_group"
@dataclass
diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py
index a44515ab..532802df 100644
--- a/python/sglang/srt/managers/scheduler.py
+++ b/python/sglang/srt/managers/scheduler.py
@@ -2036,6 +2036,9 @@ class Scheduler(
# Delete requests in the waiting queue
to_del = []
for i, req in enumerate(self.waiting_queue):
+ if recv_req.rid == "":
+ to_del.append(i)
+ continue
if req.rid.startswith(recv_req.rid):
to_del.append(i)
@@ -2065,6 +2068,9 @@ class Scheduler(
reqs = self.running_batch.reqs + self.cur_batch.reqs
for req in reqs:
+ if recv_req.rid == "" and not req.finished():
+ req.to_abort = True
+ continue
if req.rid.startswith(recv_req.rid) and not req.finished():
# Abort method 3: set `to_abort=True`
# The request will still run one decode forward pass.
@@ -2096,7 +2102,7 @@ class Scheduler(
) -> Tuple[bool, str]:
"""Update the online model parameter."""
success, message = self.tp_worker.update_weights_from_distributed(recv_req)
- if success:
+ if success and False:
flash_cache_success = self.flush_cache()
assert flash_cache_success, "Cache flush failed after updating weights"
else:
@@ -2105,7 +2111,11 @@ class Scheduler(
def update_weights_from_tensor(self, recv_req: UpdateWeightsFromTensorReqInput):
"""Update the online model parameter from tensors."""
- success, message = self.tp_worker.update_weights_from_tensor(recv_req)
+ if self.draft_worker is not None:
+ success, message = self.draft_worker.update_weights_from_tensor(recv_req)
+ else:
+ success, message = self.tp_worker.update_weights_from_tensor(recv_req)
+
# TODO extract common code b/t update_weights_from_distributed and update_weights_from_tensor later
if success:
if recv_req.flush_cache:
diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py
index d71bbdf0..f3968afb 100644
--- a/python/sglang/srt/managers/tokenizer_manager.py
+++ b/python/sglang/srt/managers/tokenizer_manager.py
@@ -199,6 +199,8 @@ class TokenizerManager:
self.is_image_gen = self.model_config.is_image_gen
self.context_len = self.model_config.context_len
self.image_token_id = self.model_config.image_token_id
+ self._updating = False
+ self._cond = asyncio.Condition()
if self.model_config.is_multimodal:
import_processors()
@@ -407,6 +409,8 @@ class TokenizerManager:
request: Optional[fastapi.Request] = None,
):
created_time = time.time()
+ async with self._cond:
+ await self._cond.wait_for(lambda: not self._updating)
self.auto_create_handle_loop()
@@ -789,7 +793,7 @@ class TokenizerManager:
return (await self.flush_cache_communicator(FlushCacheReqInput()))[0]
def abort_request(self, rid: str):
- if rid not in self.rid_to_state:
+ if rid != "" and rid not in self.rid_to_state:
return
req = AbortReq(rid)
self.send_to_scheduler.send_pyobj(req)
@@ -844,6 +848,16 @@ class TokenizerManager:
self.auto_create_handle_loop()
await self.expert_distribution_communicator(ExpertDistributionReq.DUMP_RECORD)
+ async def pause_generation(self):
+ async with self._cond:
+ self._updating = True
+ self.abort_request("")
+
+ async def continue_generation(self):
+ async with self._cond:
+ self._updating = False
+ self._cond.notify_all()
+
async def update_weights_from_disk(
self,
obj: UpdateWeightFromDiskReqInput,
@@ -1423,7 +1437,21 @@ class TokenizerManager:
asyncio.create_task(asyncio.to_thread(background_task))
def _handle_abort_req(self, recv_obj):
- self.rid_to_state.pop(recv_obj.rid, None)
+ state = self.rid_to_state[recv_obj.rid]
+ state.finished = True
+ state.out_list.append({
+ "text": "",
+ "meta_info": {
+ "id": recv_obj.rid,
+ "finish_reason": {
+ "type": "abort",
+ "message": "Abort before prefill",
+ },
+ "prompt_tokens": 1,
+ "completion_tokens": 0,
+ },
+ })
+ state.event.set()
def _handle_open_session_req_output(self, recv_obj):
self.session_futures[recv_obj.session_id].set_result(
diff --git a/python/sglang/srt/managers/tp_worker.py b/python/sglang/srt/managers/tp_worker.py
index 786a34a1..4fe46338 100644
--- a/python/sglang/srt/managers/tp_worker.py
+++ b/python/sglang/srt/managers/tp_worker.py
@@ -38,6 +38,7 @@ from sglang.srt.managers.schedule_batch import ModelWorkerBatch, global_server_a
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool, TokenToKVPoolAllocator
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
from sglang.srt.model_executor.model_runner import ModelRunner
+from sglang.srt.patch_torch import monkey_patch_torch_reductions
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import MultiprocessingSerializer, broadcast_pyobj, set_random_seed
@@ -247,11 +248,12 @@ class TpModelWorker:
self, recv_req: UpdateWeightsFromDistributedReqInput
):
success, message = self.model_runner.update_weights_from_distributed(
- recv_req.name, recv_req.dtype, recv_req.shape
+ recv_req.names, recv_req.dtypes, recv_req.shapes, recv_req.group_name
)
return success, message
def update_weights_from_tensor(self, recv_req: UpdateWeightsFromTensorReqInput):
+ monkey_patch_torch_reductions()
success, message = self.model_runner.update_weights_from_tensor(
named_tensors=MultiprocessingSerializer.deserialize(
recv_req.serialized_named_tensors[self.tp_rank]
diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py
index 995dedd0..b040f4b0 100644
--- a/python/sglang/srt/model_executor/model_runner.py
+++ b/python/sglang/srt/model_executor/model_runner.py
@@ -215,6 +215,7 @@ class ModelRunner:
self.support_pp = (
"pp_proxy_tensors" in inspect.signature(self.model.forward).parameters
)
+ self._model_update_group = {}
def initialize(self, min_per_gpu_memory: float):
server_args = self.server_args
@@ -704,7 +705,7 @@ class ModelRunner:
)
try:
- self._model_update_group = init_custom_process_group(
+ self._model_update_group[group_name] = init_custom_process_group(
backend=backend,
init_method=f"tcp://{master_address}:{master_port}",
world_size=world_size,
@@ -717,7 +718,7 @@ class ModelRunner:
logger.error(message)
return False, message
- def update_weights_from_distributed(self, name, dtype, shape):
+ def update_weights_from_distributed(self, names, dtypes, shapes, group_name):
"""
Update specific parameter in the model weights online
through `_model_update_group` process group.
@@ -727,19 +728,34 @@ class ModelRunner:
dtype: the data type of the parameter to be updated.
shape: the shape of the parameter to be updated.
"""
- target_dtype = (
- dtype if isinstance(dtype, torch.dtype) else getattr(torch, dtype)
- )
- assert (
- self._model_update_group is not None
- ), "model update group must be initialized"
+ assert group_name in self._model_update_group, (
+ f"Group {group_name} not in {list(self._model_update_group.keys)}. "
+ "Please call `init_weights_update_group` first."
+ )
try:
- weights = torch.empty(shape, dtype=target_dtype, device=self.device)
- torch.distributed.broadcast(weights, src=0, group=self._model_update_group)
- self.model.load_weights([(name, weights)])
- return True, f"Succeeded to update parameter {name} online."
+ weights = []
+ handles = []
+ for name, dtype, shape in zip(names, dtypes, shapes):
+ target_dtype = (
+ dtype if isinstance(dtype, torch.dtype) else getattr(torch, dtype)
+ )
+ weight = torch.empty(shape, dtype=target_dtype, device=self.device)
+ handles.append(
+ torch.distributed.broadcast(
+ weight,
+ src=0,
+ group=self._model_update_group[group_name],
+ async_op=True,
+ )
+ )
+ weights.append((name, weight))
+ for handle in handles:
+ handle.wait()
+
+ self.model.load_weights(weights)
+ return True, f"Succeeded to update parameter online."
except Exception as e:
error_msg = (
diff --git a/python/sglang/srt/openai_api/protocol.py b/python/sglang/srt/openai_api/protocol.py
index 351c1c56..180dd877 100644
--- a/python/sglang/srt/openai_api/protocol.py
+++ b/python/sglang/srt/openai_api/protocol.py
@@ -210,7 +210,7 @@ class CompletionResponseStreamChoice(BaseModel):
index: int
text: str
logprobs: Optional[LogProbs] = None
- finish_reason: Optional[Literal["stop", "length", "content_filter"]] = None
+ finish_reason: Optional[Literal["stop", "length", "content_filter", "abort"]] = None
matched_stop: Union[None, int, str] = None
@@ -444,7 +444,7 @@ class ChatCompletionResponseStreamChoice(BaseModel):
delta: DeltaMessage
logprobs: Optional[Union[LogProbs, ChoiceLogprobs]] = None
finish_reason: Optional[
- Literal["stop", "length", "tool_calls", "content_filter", "function_call"]
+ Literal["stop", "length", "tool_calls", "content_filter", "function_call", "abort"]
] = None
matched_stop: Union[None, int, str] = None
diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py
index ac04cdc7..d087745e 100644
--- a/python/sglang/srt/server_args.py
+++ b/python/sglang/srt/server_args.py
@@ -47,6 +47,7 @@ class ServerArgs:
tokenizer_path: Optional[str] = None
tokenizer_mode: str = "auto"
skip_tokenizer_init: bool = False
+ skip_warmup: bool = False
load_format: str = "auto"
trust_remote_code: bool = False
dtype: str = "auto"
@@ -122,6 +123,8 @@ class ServerArgs:
# Multi-node distributed serving
dist_init_addr: Optional[str] = None
+ nccl_port: Optional[int] = None
+ other_ports: Optional[List[int]] = None
nnodes: int = 1
node_rank: int = 0
@@ -538,6 +541,11 @@ class ServerArgs:
action="store_true",
help="If set, skip init tokenizer and pass input_ids in generate request.",
)
+ parser.add_argument(
+ "--skip-warmup",
+ action="store_true",
+ help="If set, skip warmup.",
+ )
parser.add_argument(
"--load-format",
type=str,
@@ -1492,6 +1500,9 @@ class ServerArgs:
help="Set multimodal attention backend.",
)
+ parser.add_argument("--nccl-port", type=int, default=None)
+ parser.add_argument("--other-ports", type=int, nargs="+", default=None)
+
@classmethod
def from_cli_args(cls, args: argparse.Namespace):
args.tp_size = args.tensor_parallel_size
@@ -1580,14 +1591,17 @@ class PortArgs:
@staticmethod
def init_new(server_args, dp_rank: Optional[int] = None) -> "PortArgs":
- port = server_args.port + random.randint(100, 1000)
- while True:
- if is_port_available(port):
- break
- if port < 60000:
- port += 42
- else:
- port -= 43
+ if server_args.nccl_port is None:
+ port = server_args.port + random.randint(100, 1000)
+ while True:
+ if is_port_available(port):
+ break
+ if port < 60000:
+ port += 42
+ else:
+ port -= 43
+ else:
+ port = server_args.nccl_port
if not server_args.enable_dp_attention:
# Normal case, use IPC within a single node
@@ -1613,6 +1627,27 @@ class PortArgs:
), "please provide --dist-init-addr as host:port of head node"
dist_init_host, dist_init_port = dist_init_addr
+
+ if server_args.other_ports is not None:
+ assert port not in server_args.other_ports
+ tokenizer_port = server_args.other_ports[0]
+ detokenizer_port = server_args.other_ports[1]
+ rpc_ipc_port = server_args.other_ports[2]
+ if dp_rank is None:
+ scheduler_input_port = server_args.other_ports[
+ 3
+ ] # TokenizerManager to DataParallelController
+ else:
+ scheduler_input_port = server_args.other_ports[3 + 1 + dp_rank]
+
+ return PortArgs(
+ tokenizer_ipc_name=f"tcp://{dist_init_host}:{tokenizer_port}",
+ scheduler_input_ipc_name=f"tcp://{dist_init_host}:{scheduler_input_port}",
+ detokenizer_ipc_name=f"tcp://{dist_init_host}:{detokenizer_port}",
+ nccl_port=port,
+ rpc_ipc_name=f"tcp://{dist_init_host}:{rpc_ipc_port}",
+ )
+
port_base = int(dist_init_port) + 1
if dp_rank is None:
scheduler_input_port = (
diff --git a/python/sglang/srt/speculative/eagle_worker.py b/python/sglang/srt/speculative/eagle_worker.py
index bc0b50f3..4dbd1cfa 100644
--- a/python/sglang/srt/speculative/eagle_worker.py
+++ b/python/sglang/srt/speculative/eagle_worker.py
@@ -11,6 +11,10 @@ from sglang.srt.distributed import GroupCoordinator, patch_tensor_parallel_group
from sglang.srt.layers.dp_attention import disable_dp_size
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.layers.sampler import get_token_ids_logprobs, get_top_logprobs
+from sglang.srt.managers.io_struct import (
+ UpdateWeightsFromDistributedReqInput,
+ UpdateWeightsFromTensorReqInput,
+)
from sglang.srt.managers.schedule_batch import (
ScheduleBatch,
get_last_loc,
@@ -22,6 +26,7 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardBatch,
ForwardMode,
)
+from sglang.srt.patch_torch import monkey_patch_torch_reductions
from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.build_eagle_tree import build_tree_kernel_efficient
from sglang.srt.speculative.eagle_draft_cuda_graph_runner import (
@@ -39,7 +44,13 @@ from sglang.srt.speculative.eagle_utils import (
select_top_k_tokens,
)
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
-from sglang.srt.utils import empty_context, fast_topk, get_available_gpu_memory, is_cuda
+from sglang.srt.utils import (
+ MultiprocessingSerializer,
+ empty_context,
+ fast_topk,
+ get_available_gpu_memory,
+ is_cuda,
+)
if is_cuda():
from sgl_kernel import segment_packbits
@@ -772,6 +783,29 @@ class EAGLEWorker(TpModelWorker):
logger.error("Detected errors during sampling! NaN in the logits.")
raise ValueError("Detected errors during sampling! NaN in the logits.")
+ def update_weights_from_distributed(
+ self, recv_req: UpdateWeightsFromDistributedReqInput
+ ):
+ raise NotImplementedError(
+ "EAGLEWorker does not support update_weights_from_distributed."
+ )
+
+ def update_weights_from_tensor(self, recv_req: UpdateWeightsFromTensorReqInput):
+ monkey_patch_torch_reductions()
+ named_tensors = MultiprocessingSerializer.deserialize(
+ recv_req.serialized_named_tensors[self.tp_rank]
+ )
+ success, message = self.model_runner.update_weights_from_tensor(
+ named_tensors=named_tensors,
+ load_format=recv_req.load_format,
+ )
+ success, message = self.target_worker.model_runner.update_weights_from_tensor(
+ named_tensors=named_tensors,
+ load_format=recv_req.load_format,
+ )
+
+ return success, message
+
def load_token_map(token_map_path: str) -> List[int]:
if not os.path.exists(token_map_path):
diff --git a/python/sglang/srt/utils.py b/python/sglang/srt/utils.py
index a669b430..9ab99ba4 100644
--- a/python/sglang/srt/utils.py
+++ b/python/sglang/srt/utils.py
@@ -1787,8 +1787,16 @@ def get_ip() -> str:
def get_open_port() -> int:
port = os.getenv("SGLANG_PORT")
- if port is not None:
- port = int(port)
+ used_port = os.getenv("SGLANG_USED_PORT")
+ if used_port is not None:
+ used_port = list(map(int, used_port.split(",")))
+
+ if port is not None or used_port is not None:
+ if used_port is not None:
+ port = max(used_port) + random.randint(1, 1000)
+ else:
+ port = int(port)
+
while True:
try:
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
diff --git a/sgl-router/py_src/sglang_router/launch_router.py b/sgl-router/py_src/sglang_router/launch_router.py
index 4f036a25..3e0cf9dd 100644
--- a/sgl-router/py_src/sglang_router/launch_router.py
+++ b/sgl-router/py_src/sglang_router/launch_router.py
@@ -43,6 +43,7 @@ class RouterArgs:
max_payload_size: int = 4 * 1024 * 1024 # 4MB
verbose: bool = False
log_dir: Optional[str] = None
+ log_level: Optional[str] = None
# Service discovery configuration
service_discovery: bool = False
selector: Dict[str, str] = dataclasses.field(default_factory=dict)
@@ -157,6 +158,13 @@ class RouterArgs:
default=None,
help="Directory to store log files. If not specified, logs are only output to console.",
)
+ parser.add_argument(
+ f"--{prefix}log-level",
+ type=str,
+ default=None,
+ choices=["debug", "info", "warning", "error", "critical"],
+ help="Set the logging level. If not specified, defaults to INFO.",
+ )
parser.add_argument(
f"--{prefix}service-discovery",
action="store_true",
@@ -225,6 +233,7 @@ class RouterArgs:
max_payload_size=getattr(args, f"{prefix}max_payload_size"),
verbose=getattr(args, f"{prefix}verbose", False),
log_dir=getattr(args, f"{prefix}log_dir", None),
+ log_level=getattr(args, f"{prefix}log_level", None),
service_discovery=getattr(args, f"{prefix}service_discovery", False),
selector=cls._parse_selector(getattr(args, f"{prefix}selector", None)),
service_discovery_port=getattr(args, f"{prefix}service_discovery_port"),
@@ -292,6 +301,7 @@ def launch_router(args: argparse.Namespace) -> Optional[Router]:
max_payload_size=router_args.max_payload_size,
verbose=router_args.verbose,
log_dir=router_args.log_dir,
+ log_level=router_args.log_level,
service_discovery=router_args.service_discovery,
selector=router_args.selector,
service_discovery_port=router_args.service_discovery_port,
diff --git a/sgl-router/py_src/sglang_router/router.py b/sgl-router/py_src/sglang_router/router.py
index c189cd58..08fde994 100644
--- a/sgl-router/py_src/sglang_router/router.py
+++ b/sgl-router/py_src/sglang_router/router.py
@@ -32,6 +32,7 @@ class Router:
max_tree_size: Maximum size of the approximation tree for cache-aware routing. Default: 2^24
verbose: Enable verbose logging. Default: False
log_dir: Directory to store log files. If None, logs are only output to console. Default: None
+ log_level: Logging level. Options: 'debug', 'info', 'warning', 'error', 'critical'.
service_discovery: Enable Kubernetes service discovery. When enabled, the router will
automatically discover worker pods based on the selector. Default: False
selector: Dictionary mapping of label keys to values for Kubernetes pod selection.
@@ -60,6 +61,7 @@ class Router:
max_payload_size: int = 4 * 1024 * 1024, # 4MB
verbose: bool = False,
log_dir: Optional[str] = None,
+ log_level: Optional[str] = None,
service_discovery: bool = False,
selector: Dict[str, str] = None,
service_discovery_port: int = 80,
@@ -85,6 +87,7 @@ class Router:
max_payload_size=max_payload_size,
verbose=verbose,
log_dir=log_dir,
+ log_level=log_level,
service_discovery=service_discovery,
selector=selector,
service_discovery_port=service_discovery_port,
diff --git a/sgl-router/py_test/test_launch_router.py b/sgl-router/py_test/test_launch_router.py
index c6f0444f..990f37a0 100644
--- a/sgl-router/py_test/test_launch_router.py
+++ b/sgl-router/py_test/test_launch_router.py
@@ -38,6 +38,7 @@ class TestLaunchRouter(unittest.TestCase):
max_payload_size=4 * 1024 * 1024, # 4MB
verbose=False,
log_dir=None,
+ log_level=None,
service_discovery=False,
selector=None,
service_discovery_port=80,
diff --git a/sgl-router/src/lib.rs b/sgl-router/src/lib.rs
index 4915d3c5..71a411a1 100644
--- a/sgl-router/src/lib.rs
+++ b/sgl-router/src/lib.rs
@@ -33,6 +33,7 @@ struct Router {
max_payload_size: usize,
verbose: bool,
log_dir: Option<String>,
+ log_level: Option<String>,
service_discovery: bool,
selector: HashMap<String, String>,
service_discovery_port: u16,
@@ -59,6 +60,7 @@ impl Router {
max_payload_size = 4 * 1024 * 1024,
verbose = false,
log_dir = None,
+ log_level = None,
service_discovery = false,
selector = HashMap::new(),
service_discovery_port = 80,
@@ -81,6 +83,7 @@ impl Router {
max_payload_size: usize,
verbose: bool,
log_dir: Option<String>,
+ log_level: Option<String>,
service_discovery: bool,
selector: HashMap<String, String>,
service_discovery_port: u16,
@@ -103,6 +106,7 @@ impl Router {
max_payload_size,
verbose,
log_dir,
+ log_level,
service_discovery,
selector,
service_discovery_port,
@@ -164,6 +168,7 @@ impl Router {
verbose: self.verbose,
max_payload_size: self.max_payload_size,
log_dir: self.log_dir.clone(),
+ log_level: self.log_level.clone(),
service_discovery_config,
prometheus_config,
})
diff --git a/sgl-router/src/server.rs b/sgl-router/src/server.rs
index 0d6cf691..7ebdc442 100644
--- a/sgl-router/src/server.rs
+++ b/sgl-router/src/server.rs
@@ -161,6 +161,7 @@ pub struct ServerConfig {
pub verbose: bool,
pub max_payload_size: usize,
pub log_dir: Option<String>,
+ pub log_level: Option<String>,
pub service_discovery_config: Option<ServiceDiscoveryConfig>,
pub prometheus_config: Option<PrometheusConfig>,
}
@@ -174,7 +175,21 @@ pub async fn startup(config: ServerConfig) -> std::io::Result<()> {
level: if config.verbose {
Level::DEBUG
} else {
- Level::INFO
+ config.log_level
+ .as_deref()
+ .and_then(|s| {
+ match s.to_uppercase().parse::<Level>() {
+ Ok(l) => Some(l),
+ Err(_) => {
+ warn!(
+ "Invalid log level string: '{}'. Defaulting to INFO (as verbose is false).",
+ s
+ );
+ None
+ }
+ }
+ })
+ .unwrap_or(Level::INFO)
},
json_format: false,
log_dir: config.log_dir.clone(),
+1 -1
View File
@@ -183,7 +183,7 @@ MEGATRON_LM_PATH=$(pip list | grep megatron-core | awk '{print $NF}')
ray job submit --address="http://127.0.0.1:8265" \
--runtime-env-json='{
"env_vars": {
"PYTHONPATH": "'$MEGATRON_LM_PATH'",
"PYTHONPATH": "/workspace/Megatron-LM-amd_version/",
"CUDA_DEVICE_MAX_CONNECTIONS": "1"
}
}' \