mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
recover patch files
This commit is contained in:
@@ -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'
|
||||
@@ -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(),
|
||||
@@ -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"
|
||||
}
|
||||
}' \
|
||||
|
||||
Reference in New Issue
Block a user